Compare commits
18
Commits
8447f88f61
...
6f09b1d2ee
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6f09b1d2ee | ||
|
|
b2230fefd8 | ||
|
|
654e6eb0d1 | ||
|
|
a317a4756b | ||
|
|
9b7e6c205f | ||
|
|
602b5ce216 | ||
|
|
8152760b5f | ||
|
|
8c052c99ee | ||
|
|
2667b8116d | ||
|
|
6dffb0305a | ||
|
|
49a9c6b3d2 | ||
|
|
cdf9145ecf | ||
|
|
85f0461b3b | ||
|
|
9f0e9195f7 | ||
|
|
88751d0b08 | ||
|
|
d0e5d910de | ||
|
|
a03504a280 | ||
|
|
d033b2ef0f |
@@ -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
|
||||||
|
|||||||
+1
-1
@@ -95,7 +95,7 @@ type: short description (~50 chars)
|
|||||||
|
|
||||||
## 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).
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|||||||
@@ -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, and ROUGE evaluation tools |
|
||||||
|
| **Extensibility** | Factory and registry architecture for models, datasets, training strategies, callbacks, kernels, and protocol components |
|
||||||
|
|
||||||
### Getting Started
|
### Getting Started
|
||||||
|
|
||||||
@@ -252,7 +256,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).
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|||||||
@@ -97,6 +97,7 @@ class AutoRegressiveLMConfig(BaseModelConfig):
|
|||||||
norm_topk_prob: bool = True
|
norm_topk_prob: bool = True
|
||||||
decoder_sparse_step: int = 1
|
decoder_sparse_step: int = 1
|
||||||
mlp_only_layers: Optional[list[int]] = None
|
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:
|
||||||
|
|||||||
@@ -18,7 +18,9 @@ SDPA is handled by the attention backend, not the wrapper functions.
|
|||||||
from astrai.extension.attention_backend import (
|
from astrai.extension.attention_backend import (
|
||||||
ATTN_BACKEND,
|
ATTN_BACKEND,
|
||||||
AttentionBackend,
|
AttentionBackend,
|
||||||
|
AttentionBackendFactory,
|
||||||
CudaBackend,
|
CudaBackend,
|
||||||
|
FlashAttnBackend,
|
||||||
TorchNativeBackend,
|
TorchNativeBackend,
|
||||||
attention,
|
attention,
|
||||||
attn_backend,
|
attn_backend,
|
||||||
@@ -36,8 +38,10 @@ from astrai.extension.rotary_backend import apply_rotary_emb
|
|||||||
__all__ = [
|
__all__ = [
|
||||||
"ATTN_BACKEND",
|
"ATTN_BACKEND",
|
||||||
"AttentionBackend",
|
"AttentionBackend",
|
||||||
|
"AttentionBackendFactory",
|
||||||
"CudaBackend",
|
"CudaBackend",
|
||||||
"TorchNativeBackend",
|
"TorchNativeBackend",
|
||||||
|
"FlashAttnBackend",
|
||||||
"TensorLayout",
|
"TensorLayout",
|
||||||
"attention",
|
"attention",
|
||||||
"attn_backend",
|
"attn_backend",
|
||||||
|
|||||||
@@ -30,6 +30,8 @@ Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
|||||||
|
|
||||||
import contextvars
|
import contextvars
|
||||||
import enum
|
import enum
|
||||||
|
import importlib
|
||||||
|
import threading
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from typing import Optional, Union
|
from typing import Optional, Union
|
||||||
@@ -42,18 +44,94 @@ from astrai.extension.attention_ops import (
|
|||||||
attn_paged_decode,
|
attn_paged_decode,
|
||||||
attn_paged_prefill,
|
attn_paged_prefill,
|
||||||
)
|
)
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
from astrai.inference.core.cache import KVCache
|
from astrai.inference.core.cache import KVCache
|
||||||
|
|
||||||
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
|
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
|
||||||
"attn_backend"
|
"attn_backend"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
_lock = threading.Lock()
|
||||||
|
_flash_available: Optional[bool] = None
|
||||||
|
|
||||||
|
|
||||||
|
def flash_attn_available() -> bool:
|
||||||
|
"""Return ``True`` if the optional ``flash-attn`` package is usable.
|
||||||
|
|
||||||
|
``flash-attn`` is not a hard dependency (declared only as an optional
|
||||||
|
extra and imported lazily), so this is checked at first use and cached.
|
||||||
|
The check is stronger than "import works": it also gates on the GPU
|
||||||
|
compute capability for the installed major version and smoke-tests a
|
||||||
|
real tiny kernel call, because wheels that import fine can still fail
|
||||||
|
at the first actual invocation (wrong arch build, torch mismatch, or a
|
||||||
|
missing ``flash_attn_func`` entry point). It never raises.
|
||||||
|
"""
|
||||||
|
global _flash_available
|
||||||
|
if _flash_available is None:
|
||||||
|
with _lock:
|
||||||
|
if _flash_available is None:
|
||||||
|
_flash_available = _flash_attn_check()
|
||||||
|
return _flash_available
|
||||||
|
|
||||||
|
|
||||||
|
_flash_attn_module = None
|
||||||
|
_flash_attn_import_tried = False
|
||||||
|
|
||||||
|
|
||||||
|
def _get_flash_attn():
|
||||||
|
"""Lazily import and cache the optional ``flash_attn`` module.
|
||||||
|
|
||||||
|
Uses ``importlib.import_module`` so no static import binds the name when
|
||||||
|
the package is absent. Returns the module object, or ``None`` if the
|
||||||
|
package is not installed or cannot be imported. Never raises.
|
||||||
|
"""
|
||||||
|
global _flash_attn_module, _flash_attn_import_tried
|
||||||
|
if not _flash_attn_import_tried:
|
||||||
|
_flash_attn_import_tried = True
|
||||||
|
try:
|
||||||
|
_flash_attn_module = importlib.import_module("flash_attn")
|
||||||
|
except Exception:
|
||||||
|
_flash_attn_module = None
|
||||||
|
return _flash_attn_module
|
||||||
|
|
||||||
|
|
||||||
|
def _flash_attn_check() -> bool:
|
||||||
|
if not torch.cuda.is_available():
|
||||||
|
return False
|
||||||
|
fa = _get_flash_attn()
|
||||||
|
if fa is None:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# version + compute-capability gate:
|
||||||
|
# FlashAttention-2 kernels need sm_70+; FlashAttention-3 (tcgen05,
|
||||||
|
# sm_90/sm_100) needs sm_90+.
|
||||||
|
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
|
||||||
|
|
||||||
|
# smoke-test the real kernel: a wheel that imports but was built for a
|
||||||
|
# different arch/torch fails here instead of at the first real forward.
|
||||||
|
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):
|
class ATTN_BACKEND(enum.Enum):
|
||||||
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
|
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
|
||||||
|
|
||||||
TORCH_NATIVE = "torch_native"
|
TORCH_NATIVE = "torch_native"
|
||||||
CUDA = "cuda"
|
CUDA = "cuda"
|
||||||
|
FLASH = "flash"
|
||||||
|
|
||||||
|
|
||||||
def get_backend() -> "AttentionBackend":
|
def get_backend() -> "AttentionBackend":
|
||||||
@@ -69,11 +147,11 @@ def get_backend() -> "AttentionBackend":
|
|||||||
|
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def attn_backend(backend: Union[ATTN_BACKEND, "AttentionBackend", type]):
|
def attn_backend(backend: Union[str, ATTN_BACKEND, "AttentionBackend", type]):
|
||||||
"""Context manager to select an attention backend.
|
"""Context manager to select an attention backend.
|
||||||
|
|
||||||
Mirrors ``torch.nn.attention.sdpa_kernel``. Accepts an
|
Mirrors ``torch.nn.attention.sdpa_kernel``. Accepts an
|
||||||
``ATTN_BACKEND`` enum value, a backend class, or a backend instance.
|
registered name, ``ATTN_BACKEND`` enum value, backend class, or instance.
|
||||||
|
|
||||||
Examples::
|
Examples::
|
||||||
|
|
||||||
@@ -85,14 +163,17 @@ def attn_backend(backend: Union[ATTN_BACKEND, "AttentionBackend", type]):
|
|||||||
...
|
...
|
||||||
"""
|
"""
|
||||||
if isinstance(backend, ATTN_BACKEND):
|
if isinstance(backend, ATTN_BACKEND):
|
||||||
instance = _BACKEND_REGISTRY[backend]()
|
instance = AttentionBackendFactory.create(backend.value)
|
||||||
|
elif isinstance(backend, str):
|
||||||
|
instance = AttentionBackendFactory.create(backend)
|
||||||
elif isinstance(backend, type) and issubclass(backend, AttentionBackend):
|
elif isinstance(backend, type) and issubclass(backend, AttentionBackend):
|
||||||
instance = backend()
|
instance = backend()
|
||||||
elif isinstance(backend, AttentionBackend):
|
elif isinstance(backend, AttentionBackend):
|
||||||
instance = backend
|
instance = backend
|
||||||
else:
|
else:
|
||||||
raise TypeError(
|
raise TypeError(
|
||||||
f"expected ATTN_BACKEND, AttentionBackend type, or instance, "
|
f"expected a registered name, ATTN_BACKEND, AttentionBackend type, "
|
||||||
|
f"or instance, "
|
||||||
f"got {type(backend).__name__}"
|
f"got {type(backend).__name__}"
|
||||||
)
|
)
|
||||||
token = _current_backend.set(instance)
|
token = _current_backend.set(instance)
|
||||||
@@ -224,6 +305,11 @@ class AttentionBackend(ABC):
|
|||||||
"""Multi-token prefill or training forward."""
|
"""Multi-token prefill or training forward."""
|
||||||
|
|
||||||
|
|
||||||
|
class AttentionBackendFactory(BaseFactory[AttentionBackend]):
|
||||||
|
"""Factory for registered attention backends."""
|
||||||
|
|
||||||
|
|
||||||
|
@AttentionBackendFactory.register(ATTN_BACKEND.TORCH_NATIVE.value)
|
||||||
class TorchNativeBackend(AttentionBackend):
|
class TorchNativeBackend(AttentionBackend):
|
||||||
"""Reference backend using torch SDPA with indirect KV cache indexing.
|
"""Reference backend using torch SDPA with indirect KV cache indexing.
|
||||||
|
|
||||||
@@ -294,11 +380,13 @@ class TorchNativeBackend(AttentionBackend):
|
|||||||
k = repeat_kv(k, n_rep)
|
k = repeat_kv(k, n_rep)
|
||||||
v = repeat_kv(v, n_rep)
|
v = repeat_kv(v, n_rep)
|
||||||
|
|
||||||
q = q.permute(0, 2, 1, 3)
|
out = F.scaled_dot_product_attention(
|
||||||
k = k.permute(0, 2, 1, 3)
|
q.permute(0, 2, 1, 3),
|
||||||
v = v.permute(0, 2, 1, 3)
|
k.permute(0, 2, 1, 3),
|
||||||
|
v.permute(0, 2, 1, 3),
|
||||||
out = F.scaled_dot_product_attention(q, k, v, attn_mask, is_causal=is_causal)
|
attn_mask,
|
||||||
|
is_causal=is_causal,
|
||||||
|
)
|
||||||
out = out.permute(0, 2, 1, 3).contiguous().flatten(2)
|
out = out.permute(0, 2, 1, 3).contiguous().flatten(2)
|
||||||
return out
|
return out
|
||||||
|
|
||||||
@@ -306,6 +394,7 @@ class TorchNativeBackend(AttentionBackend):
|
|||||||
_default_backend = TorchNativeBackend()
|
_default_backend = TorchNativeBackend()
|
||||||
|
|
||||||
|
|
||||||
|
@AttentionBackendFactory.register(ATTN_BACKEND.CUDA.value)
|
||||||
class CudaBackend(AttentionBackend):
|
class CudaBackend(AttentionBackend):
|
||||||
"""CUDA kernel backend with direct KV cache access.
|
"""CUDA kernel backend with direct KV cache access.
|
||||||
|
|
||||||
@@ -376,7 +465,7 @@ class CudaBackend(AttentionBackend):
|
|||||||
q_len = q.size(1)
|
q_len = q.size(1)
|
||||||
|
|
||||||
kv_indptr = kv_cache.kv_indptr
|
kv_indptr = kv_cache.kv_indptr
|
||||||
qo_indptr = torch.arange(b + 1, dtype=torch.int32, device=q.device) * q_len
|
qo_indptr = kv_cache.qo_indptr
|
||||||
|
|
||||||
q_flat = q.reshape(b * q_len, q.size(2), q.size(3))
|
q_flat = q.reshape(b * q_len, q.size(2), q.size(3))
|
||||||
|
|
||||||
@@ -395,7 +484,87 @@ class CudaBackend(AttentionBackend):
|
|||||||
return out.reshape(b, q_len, q.size(2), q.size(3)).flatten(2)
|
return out.reshape(b, q_len, q.size(2), q.size(3)).flatten(2)
|
||||||
|
|
||||||
|
|
||||||
_BACKEND_REGISTRY: dict[ATTN_BACKEND, type[AttentionBackend]] = {
|
@AttentionBackendFactory.register(ATTN_BACKEND.FLASH.value)
|
||||||
ATTN_BACKEND.TORCH_NATIVE: TorchNativeBackend,
|
class FlashAttnBackend(AttentionBackend):
|
||||||
ATTN_BACKEND.CUDA: CudaBackend,
|
"""FlashAttention (FA2/FA3) backend via the optional ``flash-attn`` package.
|
||||||
}
|
|
||||||
|
Uses the general ``flash_attn_func`` entry point for both prefill and
|
||||||
|
single-token decode, mirroring ``TorchNativeBackend``'s KV-cache gather.
|
||||||
|
This backend only does flash attention — inputs ``flash-attn`` cannot
|
||||||
|
express (missing package, custom attention mask, fp32, unsupported
|
||||||
|
head_dim) raise a clear error instead of silently falling back to torch.
|
||||||
|
|
||||||
|
For a torch fallback, select ``TorchNativeBackend`` instead.
|
||||||
|
"""
|
||||||
|
|
||||||
|
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
|
||||||
|
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
||||||
|
if q.size(1) == 1 and attn_mask is not None and attn_mask.dim() == 4:
|
||||||
|
pos_mask = attn_mask[:, 0, 0]
|
||||||
|
else:
|
||||||
|
pos_mask = (
|
||||||
|
torch.arange(max_len, device=q.device)[None, :]
|
||||||
|
< kv_cache.seq_lens[:, None]
|
||||||
|
)
|
||||||
|
indices = torch.where(pos_mask, indices, torch.zeros_like(indices))
|
||||||
|
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)
|
||||||
|
|
||||||
|
if attn_mask is not None and not is_causal:
|
||||||
|
raise ValueError(
|
||||||
|
"FlashAttnBackend does not support a custom attention mask; "
|
||||||
|
"use a causal mask or select TorchNativeBackend."
|
||||||
|
)
|
||||||
|
fa = _get_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().flatten(2)
|
||||||
|
|||||||
@@ -35,7 +35,7 @@ from astrai.inference.core import (
|
|||||||
KVCache,
|
KVCache,
|
||||||
KVStorage,
|
KVStorage,
|
||||||
PagePool,
|
PagePool,
|
||||||
PrefixCache,
|
RadixCache,
|
||||||
ReqToTokenPool,
|
ReqToTokenPool,
|
||||||
Task,
|
Task,
|
||||||
TaskManager,
|
TaskManager,
|
||||||
@@ -66,7 +66,7 @@ __all__ = [
|
|||||||
"KVCache",
|
"KVCache",
|
||||||
"KVStorage",
|
"KVStorage",
|
||||||
"PagePool",
|
"PagePool",
|
||||||
"PrefixCache",
|
"RadixCache",
|
||||||
"ReqToTokenPool",
|
"ReqToTokenPool",
|
||||||
"page_hash",
|
"page_hash",
|
||||||
"sample",
|
"sample",
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ from astrai.inference.core.cache import (
|
|||||||
KVCache,
|
KVCache,
|
||||||
KVStorage,
|
KVStorage,
|
||||||
PagePool,
|
PagePool,
|
||||||
PrefixCache,
|
RadixCache,
|
||||||
ReqToTokenPool,
|
ReqToTokenPool,
|
||||||
page_hash,
|
page_hash,
|
||||||
)
|
)
|
||||||
@@ -18,7 +18,7 @@ __all__ = [
|
|||||||
"KVCache",
|
"KVCache",
|
||||||
"KVStorage",
|
"KVStorage",
|
||||||
"PagePool",
|
"PagePool",
|
||||||
"PrefixCache",
|
"RadixCache",
|
||||||
"ReqToTokenPool",
|
"ReqToTokenPool",
|
||||||
"page_hash",
|
"page_hash",
|
||||||
"Executor",
|
"Executor",
|
||||||
|
|||||||
+183
-50
@@ -4,7 +4,7 @@ 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 2 — ReqToTokenPool: index table [req_idx, pos] → physical token slot
|
||||||
Layer 3 — Allocator: slot/page allocation with ref-counting and LRU
|
Layer 3 — Allocator: slot/page allocation with ref-counting and LRU
|
||||||
|
|
||||||
PagePool orchestrates all three plus PrefixCache (content addressing).
|
PagePool orchestrates all three plus RadixCache (prefix addressing).
|
||||||
KVCache is a pure dataclass passed to the model for direct buffer access.
|
KVCache is a pure dataclass passed to the model for direct buffer access.
|
||||||
|
|
||||||
Two modes:
|
Two modes:
|
||||||
@@ -20,11 +20,15 @@ from typing import Callable, Dict, List, Optional
|
|||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.inference.core.workspace import InferenceWorkspace
|
||||||
|
|
||||||
def page_hash(token_ids: List[int], page_idx: int, page_size: int) -> int:
|
|
||||||
|
def page_hash(
|
||||||
|
token_ids: List[int], page_idx: int, page_size: int, parent_hash: int = 0
|
||||||
|
) -> int:
|
||||||
start = page_idx * page_size
|
start = page_idx * page_size
|
||||||
end = min(start + page_size, len(token_ids))
|
end = min(start + page_size, len(token_ids))
|
||||||
h = 0
|
h = parent_hash
|
||||||
for i in range(start, end):
|
for i in range(start, end):
|
||||||
h = (h * 31 + token_ids[i]) & 0xFFFFFFFFFFFFFFFF
|
h = (h * 31 + token_ids[i]) & 0xFFFFFFFFFFFFFFFF
|
||||||
return h
|
return h
|
||||||
@@ -81,45 +85,96 @@ class Allocator:
|
|||||||
self._lru.move_to_end(idx)
|
self._lru.move_to_end(idx)
|
||||||
|
|
||||||
|
|
||||||
class PrefixCache:
|
class RadixNode:
|
||||||
"""Hash-based prefix matching: maps page hashes to physical page indices."""
|
"""A page-aligned edge in the CPU-side prefix radix."""
|
||||||
|
|
||||||
|
__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):
|
def __init__(self, page_size: int):
|
||||||
self._page_size = page_size
|
self._page_size = page_size
|
||||||
|
self._root = RadixNode()
|
||||||
|
self._page_to_node: Dict[int, RadixNode] = {}
|
||||||
|
# Retained as an introspection-compatible map; matching never relies on
|
||||||
|
# this lossy value.
|
||||||
self._page_to_hash: Dict[int, int] = {}
|
self._page_to_hash: Dict[int, int] = {}
|
||||||
self._hash_to_page: Dict[int, int] = {}
|
|
||||||
self._lock = threading.Lock()
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
def evict(self, idx: int):
|
def evict(self, idx: int):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
h = self._page_to_hash.pop(idx, None)
|
node = self._page_to_node.pop(idx, None)
|
||||||
if h is not None:
|
self._page_to_hash.pop(idx, None)
|
||||||
self._hash_to_page.pop(h, 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:
|
def has_page(self, idx: int) -> bool:
|
||||||
with self._lock:
|
with self._lock:
|
||||||
return idx in self._page_to_hash
|
return idx in self._page_to_node
|
||||||
|
|
||||||
def lookup(self, token_ids: List[int]) -> List[int]:
|
def lookup(self, token_ids: List[int]) -> List[int]:
|
||||||
with self._lock:
|
with self._lock:
|
||||||
full_pages = len(token_ids) // self._page_size
|
full_pages = len(token_ids) // self._page_size
|
||||||
hits: List[int] = []
|
hits: List[int] = []
|
||||||
|
node = self._root
|
||||||
for i in range(full_pages):
|
for i in range(full_pages):
|
||||||
h = page_hash(token_ids, i, self._page_size)
|
start = i * self._page_size
|
||||||
p = self._hash_to_page.get(h)
|
page_tokens = tuple(token_ids[start : start + self._page_size])
|
||||||
if p is None:
|
child = node.children.get(page_tokens)
|
||||||
|
if child is None or child.page_idx is None:
|
||||||
break
|
break
|
||||||
hits.append(p)
|
hits.append(child.page_idx)
|
||||||
|
node = child
|
||||||
return hits
|
return hits
|
||||||
|
|
||||||
def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int):
|
def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
h = page_hash(token_ids, logical_page_idx, self._page_size)
|
full_pages = len(token_ids) // self._page_size
|
||||||
old_h = self._page_to_hash.pop(page_idx, None)
|
if logical_page_idx >= full_pages:
|
||||||
if old_h is not None:
|
return
|
||||||
self._hash_to_page.pop(old_h, None)
|
old = self._page_to_node.pop(page_idx, None)
|
||||||
self._page_to_hash[page_idx] = h
|
self._page_to_hash.pop(page_idx, None)
|
||||||
self._hash_to_page[h] = page_idx
|
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)
|
||||||
|
self._page_to_hash.pop(replaced, None)
|
||||||
|
node.page_idx = page_idx
|
||||||
|
self._page_to_node[page_idx] = node
|
||||||
|
self._page_to_hash[page_idx] = page_hash(
|
||||||
|
token_ids, logical_page_idx, self._page_size
|
||||||
|
)
|
||||||
|
|
||||||
|
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 ReqToTokenPool:
|
class ReqToTokenPool:
|
||||||
@@ -215,12 +270,13 @@ class KVCache:
|
|||||||
out_cache_loc: Tensor
|
out_cache_loc: Tensor
|
||||||
max_len: int = 0
|
max_len: int = 0
|
||||||
kv_indptr: Optional[Tensor] = None
|
kv_indptr: Optional[Tensor] = None
|
||||||
|
qo_indptr: Optional[Tensor] = None
|
||||||
|
|
||||||
|
|
||||||
class PagePool:
|
class PagePool:
|
||||||
"""Top-level KV cache manager.
|
"""Top-level KV cache manager.
|
||||||
|
|
||||||
Combines KVStorage + ReqToTokenPool + Allocator + PrefixCache.
|
Combines KVStorage + ReqToTokenPool + Allocator + RadixCache.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
n_layers: Number of transformer layers.
|
n_layers: Number of transformer layers.
|
||||||
@@ -272,11 +328,11 @@ class PagePool:
|
|||||||
i * max_seq_len, (i + 1) * max_seq_len, device=device
|
i * max_seq_len, (i + 1) * max_seq_len, device=device
|
||||||
)
|
)
|
||||||
self._alloc: Optional[Allocator] = None
|
self._alloc: Optional[Allocator] = None
|
||||||
self._prefix: Optional[PrefixCache] = None
|
self._prefix: Optional[RadixCache] = None
|
||||||
else:
|
else:
|
||||||
n_pages = self.n_tokens // page_size
|
n_pages = self.n_tokens // page_size
|
||||||
self._alloc = Allocator(n_pages)
|
self._alloc = Allocator(n_pages)
|
||||||
self._prefix = PrefixCache(page_size) if page_size > 1 else None
|
self._prefix = RadixCache(page_size) if page_size > 1 else None
|
||||||
if self._prefix is not None:
|
if self._prefix is not None:
|
||||||
self._alloc.on_evict = self._prefix.evict
|
self._alloc.on_evict = self._prefix.evict
|
||||||
|
|
||||||
@@ -287,6 +343,14 @@ class PagePool:
|
|||||||
self._task_pages: Dict[str, List[int]] = {}
|
self._task_pages: Dict[str, List[int]] = {}
|
||||||
self._lock = threading.Lock()
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
# Steady-state decode validation state: the ordered task set and its
|
||||||
|
# Python seq_lens mirror. When the same set advances every sequence
|
||||||
|
# by exactly one token per step, bind_tasks updates the stable
|
||||||
|
# buffers in-place (+=1 / +=inc) instead of re-cumsumming. Any
|
||||||
|
# task-set change is a miss and rebuilds.
|
||||||
|
self._bind_sig: Optional[tuple] = None
|
||||||
|
self._bind_seq_lens: Optional[List[int]] = None
|
||||||
|
|
||||||
# ---- task lifecycle ----
|
# ---- task lifecycle ----
|
||||||
|
|
||||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
||||||
@@ -371,33 +435,39 @@ class PagePool:
|
|||||||
|
|
||||||
def task_extend(self, task_id: str, pos: int) -> bool:
|
def task_extend(self, task_id: str, pos: int) -> bool:
|
||||||
req_idx = self._task_req.get(task_id)
|
req_idx = self._task_req.get(task_id)
|
||||||
if req_idx is None:
|
if req_idx is None or pos >= self.max_seq_len:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
if self.contiguous:
|
# Paged mode must also claim a physical slot for the new token;
|
||||||
return pos < self.max_seq_len
|
# contiguous mode's block is pre-allocated so this is a no-op.
|
||||||
|
if not self.contiguous and not self._extend_slot(task_id, req_idx, pos):
|
||||||
|
return False
|
||||||
|
|
||||||
|
self._task_len[req_idx] = pos + 1
|
||||||
|
return True
|
||||||
|
|
||||||
|
def _extend_slot(self, task_id: str, req_idx: int, pos: int) -> bool:
|
||||||
|
"""Allocate the physical slot for one extended token (paged mode)."""
|
||||||
if self.page_size == 1:
|
if self.page_size == 1:
|
||||||
slots = self._alloc_tokens(1)
|
slots = self._alloc_tokens(1)
|
||||||
if slots is None:
|
if slots is None:
|
||||||
return False
|
return False
|
||||||
self._task_slots.setdefault(task_id, []).extend(slots)
|
self._task_slots.setdefault(task_id, []).extend(slots)
|
||||||
self._req_pool.req_to_token[req_idx, pos] = slots[0]
|
self._req_pool.req_to_token[req_idx, pos] = slots[0]
|
||||||
else:
|
return True
|
||||||
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
|
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
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def task_cached(self, task_id: str) -> int:
|
def task_cached(self, task_id: str) -> int:
|
||||||
@@ -413,32 +483,94 @@ class PagePool:
|
|||||||
for i in range(start_logical_page, min(full_pages, len(pages))):
|
for i in range(start_logical_page, min(full_pages, len(pages))):
|
||||||
self._prefix.record(pages[i], prompt_ids, i)
|
self._prefix.record(pages[i], prompt_ids, i)
|
||||||
|
|
||||||
|
def task_cacheable_ids(
|
||||||
|
self, task_id: str, prompt_ids: List[int], output_ids: List[int]
|
||||||
|
):
|
||||||
|
"""Return the sequence whose KV entries are already materialized.
|
||||||
|
|
||||||
|
The first sampled output is produced by prompt prefill, and the last
|
||||||
|
sampled output has not been decoded into KV yet. Therefore the cache
|
||||||
|
can safely retain the prompt plus every output except the last one.
|
||||||
|
"""
|
||||||
|
return list(prompt_ids) + list(output_ids[:-1])
|
||||||
|
|
||||||
# ---- bind for forward ----
|
# ---- bind for forward ----
|
||||||
|
|
||||||
def bind_tasks(
|
def bind_tasks(
|
||||||
self,
|
self,
|
||||||
task_ids: List[str],
|
task_ids: List[str],
|
||||||
seq_lens: List[int],
|
workspace: InferenceWorkspace,
|
||||||
device: torch.device,
|
device: Optional[torch.device] = None,
|
||||||
start_pos: Optional[int] = None,
|
start_pos: Optional[int] = None,
|
||||||
) -> KVCache:
|
) -> KVCache:
|
||||||
|
if device is None:
|
||||||
|
device = workspace.device
|
||||||
req_indices = [self._task_req[tid] for tid in task_ids]
|
req_indices = [self._task_req[tid] for tid in task_ids]
|
||||||
req_pool_indices = torch.tensor(req_indices, dtype=torch.long, device=device)
|
# Per-request lengths come from the pool's own tracking (task_alloc
|
||||||
seq_lens_t = torch.tensor(seq_lens, dtype=torch.long, device=device)
|
# sets len(prompt_ids); task_extend sets pos+1), so callers need not
|
||||||
|
# pass them.
|
||||||
|
seq_lens = [self._task_len[req_idx] for req_idx in req_indices]
|
||||||
|
b = len(task_ids)
|
||||||
|
sig = tuple(task_ids)
|
||||||
|
|
||||||
|
# Write into the caller's workspace buffers (fixed addresses, sized
|
||||||
|
# to max_batch/max_seq at init) — the sole owner of the per-step
|
||||||
|
# KV bind tensors.
|
||||||
|
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
|
||||||
|
|
||||||
|
incremental = (
|
||||||
|
start_pos is None
|
||||||
|
and self._bind_sig is not None
|
||||||
|
and self._bind_sig == sig
|
||||||
|
and self._bind_seq_lens is not None
|
||||||
|
and len(self._bind_seq_lens) == b
|
||||||
|
and all(s == p + 1 for s, p in zip(seq_lens, self._bind_seq_lens))
|
||||||
|
)
|
||||||
|
if incremental:
|
||||||
|
# Steady-state decode: advance the stable buffers in-place.
|
||||||
|
# Normal-mode buffers keep ``+=`` legal regardless of whether
|
||||||
|
# this runs inside ``torch.inference_mode()``.
|
||||||
|
sl_buf[:b] += 1
|
||||||
|
kvp_buf[: b + 1] += inc_buf[: b + 1]
|
||||||
|
req_pool_indices = rpi_buf[:b]
|
||||||
|
seq_lens_t = sl_buf[:b]
|
||||||
|
kv_indptr = kvp_buf[: b + 1]
|
||||||
|
else:
|
||||||
|
# Cold path: fill the stable buffers from fresh host tensors.
|
||||||
|
rpi_buf[:b].copy_(
|
||||||
|
torch.tensor(req_indices, dtype=torch.long, 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]
|
||||||
|
self._bind_sig = sig
|
||||||
|
self._bind_seq_lens = list(seq_lens)
|
||||||
|
|
||||||
if start_pos is not None:
|
if start_pos is not None:
|
||||||
seq_len = seq_lens[0]
|
seq_len = seq_lens[0]
|
||||||
out_cache_loc = self._req_pool.req_to_token[
|
out_cache_loc = self._req_pool.req_to_token[
|
||||||
req_pool_indices, start_pos:seq_len
|
req_pool_indices, start_pos:seq_len
|
||||||
]
|
]
|
||||||
|
# Ragged query segmentation for the prefill kernel, computed once
|
||||||
|
# (was rebuilt per layer in CudaBackend.fwd_prefill).
|
||||||
|
q_len = seq_len - start_pos
|
||||||
|
workspace.qo_indptr[: b + 1].copy_(
|
||||||
|
torch.arange(b + 1, dtype=torch.int32, device=device) * q_len
|
||||||
|
)
|
||||||
|
qo_indptr = workspace.qo_indptr[: b + 1]
|
||||||
else:
|
else:
|
||||||
write_pos = seq_lens_t - 1
|
write_pos = seq_lens_t - 1
|
||||||
out_cache_loc = self._req_pool.req_to_token[
|
loc = self._req_pool.req_to_token[req_pool_indices, write_pos].unsqueeze(-1)
|
||||||
req_pool_indices, write_pos
|
ocl_buf[:b].copy_(loc)
|
||||||
].unsqueeze(-1)
|
out_cache_loc = ocl_buf[:b]
|
||||||
|
qo_indptr = None
|
||||||
kv_indptr = torch.zeros(len(seq_lens) + 1, dtype=torch.int32, device=device)
|
|
||||||
kv_indptr[1:] = seq_lens_t.cumsum(0).to(torch.int32)
|
|
||||||
|
|
||||||
return KVCache(
|
return KVCache(
|
||||||
k_buffer=self._storage.k_buffer,
|
k_buffer=self._storage.k_buffer,
|
||||||
@@ -449,6 +581,7 @@ class PagePool:
|
|||||||
out_cache_loc=out_cache_loc,
|
out_cache_loc=out_cache_loc,
|
||||||
max_len=max(seq_lens),
|
max_len=max(seq_lens),
|
||||||
kv_indptr=kv_indptr,
|
kv_indptr=kv_indptr,
|
||||||
|
qo_indptr=qo_indptr,
|
||||||
)
|
)
|
||||||
|
|
||||||
# ---- internals ----
|
# ---- internals ----
|
||||||
|
|||||||
@@ -1,10 +1,13 @@
|
|||||||
import logging
|
import logging
|
||||||
|
from dataclasses import dataclass
|
||||||
from typing import List, Optional
|
from typing import List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.inference.core.cache import PagePool
|
from astrai.inference.core.cache import PagePool
|
||||||
from astrai.inference.core.task import Task
|
from astrai.inference.core.task import Task
|
||||||
|
from astrai.inference.core.workspace import InferenceWorkspace
|
||||||
from astrai.inference.sample import sample
|
from astrai.inference.sample import sample
|
||||||
from astrai.model.automodel import AutoModel
|
from astrai.model.automodel import AutoModel
|
||||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||||
@@ -12,6 +15,42 @@ from astrai.tokenize.tokenizer import AutoTokenizer
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@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())
|
||||||
|
|
||||||
|
|
||||||
|
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()),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class Executor:
|
class Executor:
|
||||||
"""Model forward passes for prefill and decode phases."""
|
"""Model forward passes for prefill and decode phases."""
|
||||||
|
|
||||||
@@ -29,9 +68,82 @@ class Executor:
|
|||||||
self.device = device or next(model.parameters()).device
|
self.device = device or next(model.parameters()).device
|
||||||
self.dtype = dtype or next(model.parameters()).dtype
|
self.dtype = dtype or next(model.parameters()).dtype
|
||||||
|
|
||||||
def execute_prefill(self, tasks: List[Task], prompt_len: int, start_pos: int = 0):
|
# Per-step decode cache for the steady-state case where the same
|
||||||
|
# ordered task set decodes one token per step. Sampling params are
|
||||||
|
# constant across steps; position_ids grows by exactly 1. Single-slot:
|
||||||
|
# any task-set change is a cache miss.
|
||||||
|
self._decode_cache: Optional[tuple] = 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.
|
||||||
|
self._workspace = InferenceWorkspace(
|
||||||
|
max_batch_size=kv_cache.max_batch_size,
|
||||||
|
max_seq_len=kv_cache.max_seq_len,
|
||||||
|
device=self.device,
|
||||||
|
dtype=self.dtype,
|
||||||
|
)
|
||||||
|
|
||||||
|
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:
|
if start_pos >= prompt_len:
|
||||||
return
|
return []
|
||||||
|
|
||||||
tasks = sorted(tasks, key=lambda t: t.task_id)
|
tasks = sorted(tasks, key=lambda t: t.task_id)
|
||||||
batch_sz = len(tasks)
|
batch_sz = len(tasks)
|
||||||
@@ -53,14 +165,19 @@ class Executor:
|
|||||||
)
|
)
|
||||||
|
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
self.model(
|
outputs = self.model(
|
||||||
input_ids,
|
input_ids,
|
||||||
input_mask=input_mask,
|
input_mask=input_mask,
|
||||||
position_ids=position_ids,
|
position_ids=position_ids,
|
||||||
kv_cache=self.kv_cache.bind_tasks(
|
kv_cache=self.kv_cache.bind_tasks(
|
||||||
task_ids, [prompt_len] * batch_sz, self.device, start_pos=start_pos
|
task_ids,
|
||||||
|
self._workspace,
|
||||||
|
start_pos=start_pos,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
logits = outputs["logits"][:, -1, :]
|
||||||
|
|
||||||
|
return tasks, self._sample_logits(logits, tasks, return_logprobs)
|
||||||
|
|
||||||
def execute_decode(
|
def execute_decode(
|
||||||
self, tasks: List[Task], return_logprobs: bool = False
|
self, tasks: List[Task], return_logprobs: bool = False
|
||||||
@@ -82,93 +199,43 @@ class Executor:
|
|||||||
if not tasks:
|
if not tasks:
|
||||||
return []
|
return []
|
||||||
|
|
||||||
input_ids = torch.tensor(
|
input_ids = self._workspace.fill_input_ids(
|
||||||
[t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks],
|
[t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks]
|
||||||
dtype=torch.long,
|
).unsqueeze(1)
|
||||||
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]
|
task_ids = [t.task_id for t in tasks]
|
||||||
|
|
||||||
temperatures = torch.tensor([t.temperature for t in tasks], device=self.device)
|
sig = tuple(task_ids)
|
||||||
top_ks = torch.tensor([t.top_k for t in tasks], device=self.device)
|
cur_positions = [t.next_pos for t in tasks]
|
||||||
top_ps = torch.tensor([t.top_p for t in tasks], device=self.device)
|
cached = self._decode_cache
|
||||||
freq_penalties = torch.tensor(
|
if (
|
||||||
[t.frequency_penalty for t in tasks], device=self.device
|
cached is not None
|
||||||
)
|
and cached[0] == sig
|
||||||
|
and cur_positions == [p + 1 for p in cached[1]]
|
||||||
has_freq = bool((freq_penalties != 0).any())
|
):
|
||||||
if has_freq:
|
_, _, info, position_ids = cached
|
||||||
history_lists = []
|
position_ids += 1
|
||||||
history_lens = []
|
self._decode_cache = (sig, cur_positions, info, position_ids)
|
||||||
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:
|
else:
|
||||||
padded_ids = None
|
info = _build_sampling_batch_info(tasks, self.device)
|
||||||
padded_mask = None
|
position_ids = torch.tensor(
|
||||||
|
cur_positions, dtype=torch.long, device=self.device
|
||||||
|
)
|
||||||
|
self._decode_cache = (sig, cur_positions, info, position_ids)
|
||||||
|
|
||||||
|
total_len = max(t.next_pos for t in tasks) + 1
|
||||||
|
input_mask = self._workspace.decode_mask(position_ids, total_len)
|
||||||
|
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
outputs = self.model(
|
outputs = self.model(
|
||||||
input_ids.unsqueeze(1),
|
input_ids,
|
||||||
input_mask=input_mask,
|
input_mask=input_mask,
|
||||||
kv_cache=self.kv_cache.bind_tasks(
|
kv_cache=self.kv_cache.bind_tasks(
|
||||||
task_ids,
|
task_ids,
|
||||||
[t.next_pos + 1 for t in tasks],
|
self._workspace,
|
||||||
self.device,
|
|
||||||
),
|
),
|
||||||
position_ids=position_ids.unsqueeze(1),
|
position_ids=position_ids.unsqueeze(1),
|
||||||
)
|
)
|
||||||
logits = outputs["logits"][:, -1, :]
|
logits = outputs["logits"][:, -1, :]
|
||||||
|
|
||||||
if return_logprobs:
|
return self._sample_logits(logits, tasks, return_logprobs, info=info)
|
||||||
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()
|
|
||||||
|
|||||||
@@ -83,6 +83,80 @@ class InferenceScheduler:
|
|||||||
def get_stats(self) -> Dict[str, Any]:
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
return self._task_mgr.get_stats()
|
return self._task_mgr.get_stats()
|
||||||
|
|
||||||
|
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``.
|
||||||
|
"""
|
||||||
|
cache = self._cache
|
||||||
|
|
||||||
|
to_prefill = [t for t in tasks if t.output_tokens == 0 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], List[Task]] = {}
|
||||||
|
for t in to_prefill:
|
||||||
|
start_pos = min(cache.task_cached(t.task_id), len(t.prompt_ids) - 1)
|
||||||
|
groups.setdefault((len(t.prompt_ids), start_pos), []).append(t)
|
||||||
|
|
||||||
|
for (prompt_len, start_pos), group in groups.items():
|
||||||
|
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
|
||||||
|
prefilled_ids.add(t.task_id)
|
||||||
|
produced.append(t)
|
||||||
|
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
|
||||||
|
)
|
||||||
|
|
||||||
|
decoded: List[Task] = []
|
||||||
|
aborted: List[Task] = []
|
||||||
|
for t in tasks:
|
||||||
|
if t.task_id in prefilled_ids:
|
||||||
|
continue
|
||||||
|
if cache.task_extend(t.task_id, t.next_pos):
|
||||||
|
decoded.append(t)
|
||||||
|
else:
|
||||||
|
t.status = TaskStatus.ABORTED
|
||||||
|
aborted.append(t)
|
||||||
|
|
||||||
|
if decoded:
|
||||||
|
step_out = self._executor.execute_decode(
|
||||||
|
decoded, return_logprobs=return_logprobs
|
||||||
|
)
|
||||||
|
for t, out in zip(decoded, step_out):
|
||||||
|
t.output_ids.append(out[0] if return_logprobs else out)
|
||||||
|
t.output_tokens += 1
|
||||||
|
produced.append(t)
|
||||||
|
|
||||||
|
return produced, aborted
|
||||||
|
|
||||||
def _run_generation_loop(self):
|
def _run_generation_loop(self):
|
||||||
stop_ids = self._task_mgr.tokenizer.stop_ids
|
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||||
cache = self._cache
|
cache = self._cache
|
||||||
@@ -90,6 +164,13 @@ class InferenceScheduler:
|
|||||||
while not self._stop_event.is_set():
|
while not self._stop_event.is_set():
|
||||||
finished = self._task_mgr.remove_finished_tasks(stop_ids)
|
finished = self._task_mgr.remove_finished_tasks(stop_ids)
|
||||||
for task in finished:
|
for task in finished:
|
||||||
|
if task.status == TaskStatus.FINISHED:
|
||||||
|
cache.task_record_hashes(
|
||||||
|
task.task_id,
|
||||||
|
cache.task_cacheable_ids(
|
||||||
|
task.task_id, task.prompt_ids, task.output_ids
|
||||||
|
),
|
||||||
|
)
|
||||||
cache.task_free(task.task_id)
|
cache.task_free(task.task_id)
|
||||||
|
|
||||||
active = self._task_mgr.get_active_tasks()
|
active = self._task_mgr.get_active_tasks()
|
||||||
@@ -111,61 +192,21 @@ class InferenceScheduler:
|
|||||||
|
|
||||||
active = self._task_mgr.get_active_tasks()
|
active = self._task_mgr.get_active_tasks()
|
||||||
|
|
||||||
to_prefill = [
|
decoded, aborted = self._step(active)
|
||||||
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 aborted:
|
||||||
for t in to_prefill:
|
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||||
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():
|
for t in decoded:
|
||||||
self._executor.execute_prefill(group, prompt_len, start_pos)
|
new_text = t.decode_new_token(self._task_mgr.tokenizer)
|
||||||
start_logical_page = start_pos // getattr(
|
if new_text:
|
||||||
cache, "page_size", 64
|
self._task_mgr.invoke_callback(t.task_id, new_text)
|
||||||
)
|
if t.is_finished(stop_ids):
|
||||||
for t in group:
|
remaining = t.flush_remaining(self._task_mgr.tokenizer)
|
||||||
cache.task_record_hashes(
|
if remaining:
|
||||||
t.task_id, t.prompt_ids, start_logical_page
|
self._task_mgr.invoke_callback(t.task_id, remaining)
|
||||||
)
|
|
||||||
|
|
||||||
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)
|
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:
|
except Exception as e:
|
||||||
self._stop_event.set()
|
self._stop_event.set()
|
||||||
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
||||||
@@ -265,36 +306,10 @@ class InferenceScheduler:
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
live = [t for t in tasks if t is not None]
|
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:
|
while live:
|
||||||
valid: List[Task] = []
|
decoded, _ = self._step(live, return_logprobs=return_logprobs)
|
||||||
for t in sorted(live, key=lambda x: x.task_id):
|
live = [t for t in decoded if not t.is_finished(stop_ids)]
|
||||||
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:
|
finally:
|
||||||
for t in tasks:
|
for t in tasks:
|
||||||
if t is not None:
|
if t is not None:
|
||||||
|
|||||||
@@ -105,7 +105,8 @@ class Task:
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def next_pos(self) -> int:
|
def next_pos(self) -> int:
|
||||||
return self.input_tokens + len(self.output_ids)
|
# The first output is sampled from prefill and enters KV on the next step.
|
||||||
|
return self.input_tokens + max(0, len(self.output_ids) - 1)
|
||||||
|
|
||||||
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:
|
||||||
|
|||||||
@@ -0,0 +1,111 @@
|
|||||||
|
"""Pre-allocated buffers for the inference decode hot path.
|
||||||
|
|
||||||
|
Mirrors SGLang's pre-allocated input buffers (``input_buffers.py``): tensors
|
||||||
|
are sized once to the server's maximum dimensions and sliced to the live
|
||||||
|
batch each step, so the per-token decode loop never calls
|
||||||
|
``torch.empty``/``torch.zeros``/``torch.arange`` for the hot shapes. Fills
|
||||||
|
go through ``out=`` variants (``torch.ge``) which write into the stable
|
||||||
|
buffers instead of allocating fresh results.
|
||||||
|
|
||||||
|
All buffers are allocated eagerly at init (nothing is lazy), so the
|
||||||
|
workspace is CUDA-graph-capture friendly: the decode step reads/writes
|
||||||
|
fixed-address tensors with no allocation during capture.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
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.
|
||||||
|
|
||||||
|
No re-allocation while the server's bounds are respected.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
max_batch_size: int,
|
||||||
|
max_seq_len: int,
|
||||||
|
device: torch.device,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
):
|
||||||
|
self.max_batch_size = max_batch_size
|
||||||
|
self.max_seq_len = max_seq_len
|
||||||
|
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.long, device=device
|
||||||
|
)
|
||||||
|
self.seq_lens = torch.empty((max_batch_size,), dtype=torch.long, device=device)
|
||||||
|
self.kv_indptr = torch.empty(
|
||||||
|
(max_batch_size + 1,), dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
self.qo_indptr = torch.empty(
|
||||||
|
(max_batch_size + 1,), dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
self.inc = torch.arange(max_batch_size + 1, dtype=torch.int32, device=device)
|
||||||
|
self.out_cache_loc = torch.empty(
|
||||||
|
(max_batch_size, 1), dtype=torch.long, 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
|
||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -6,13 +6,14 @@ from torch import Tensor
|
|||||||
|
|
||||||
from astrai.inference.core.cache import KVCache
|
from astrai.inference.core.cache import KVCache
|
||||||
from astrai.model.components.attention import AttnFactory
|
from astrai.model.components.attention import AttnFactory
|
||||||
from astrai.model.components.mlp import FFNFactory
|
from astrai.model.components.mlp import FFNFactory, RouterStats
|
||||||
from astrai.model.components.norm import RMSNorm
|
from astrai.model.components.norm import RMSNorm
|
||||||
|
|
||||||
|
|
||||||
class DecoderOutput(TypedDict):
|
class DecoderOutput(TypedDict):
|
||||||
hidden_states: Tensor
|
hidden_states: Tensor
|
||||||
aux_loss: Optional[Tensor]
|
aux_loss: Optional[Tensor]
|
||||||
|
router_stats: Optional[RouterStats]
|
||||||
|
|
||||||
|
|
||||||
class DecoderBlock(nn.Module):
|
class DecoderBlock(nn.Module):
|
||||||
@@ -66,4 +67,8 @@ class DecoderBlock(nn.Module):
|
|||||||
mlp_output = self.mlp(normalized)
|
mlp_output = self.mlp(normalized)
|
||||||
x = mlp_output["hidden_states"] + x
|
x = mlp_output["hidden_states"] + x
|
||||||
|
|
||||||
return {"hidden_states": x, "aux_loss": mlp_output["aux_loss"]}
|
return {
|
||||||
|
"hidden_states": x,
|
||||||
|
"aux_loss": mlp_output["aux_loss"],
|
||||||
|
"router_stats": mlp_output.get("router_stats"),
|
||||||
|
}
|
||||||
|
|||||||
@@ -13,14 +13,26 @@ 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):
|
class FFNOutput(TypedDict):
|
||||||
hidden_states: Tensor
|
hidden_states: Tensor
|
||||||
aux_loss: Optional[Tensor]
|
aux_loss: Optional[Tensor]
|
||||||
|
router_stats: Optional[RouterStats]
|
||||||
|
|
||||||
|
|
||||||
class RoutedOutput(TypedDict):
|
class RoutedOutput(TypedDict):
|
||||||
hidden_states: Tensor
|
hidden_states: Tensor
|
||||||
aux_loss: Optional[Tensor]
|
aux_loss: Optional[Tensor]
|
||||||
|
router_stats: Optional[RouterStats]
|
||||||
|
|
||||||
|
|
||||||
@FFNFactory.register("mlp")
|
@FFNFactory.register("mlp")
|
||||||
@@ -34,7 +46,7 @@ class MLP(nn.Module):
|
|||||||
def forward(self, x: Tensor) -> FFNOutput:
|
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 {"hidden_states": out, "aux_loss": None}
|
return {"hidden_states": out, "aux_loss": None, "router_stats": None}
|
||||||
|
|
||||||
|
|
||||||
@FFNFactory.register("moe")
|
@FFNFactory.register("moe")
|
||||||
@@ -95,7 +107,11 @@ class DeepSeekMoE(nn.Module):
|
|||||||
routed_output = self._routed_forward(x_flat, include_aux_loss)
|
routed_output = self._routed_forward(x_flat, include_aux_loss)
|
||||||
|
|
||||||
out = (shared_out + routed_output["hidden_states"]).view(bsz, seq_len, dim)
|
out = (shared_out + routed_output["hidden_states"]).view(bsz, seq_len, dim)
|
||||||
return {"hidden_states": out, "aux_loss": routed_output["aux_loss"]}
|
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:
|
||||||
@@ -108,34 +124,54 @@ class DeepSeekMoE(nn.Module):
|
|||||||
def _routed_forward(self, x: Tensor, include_aux_loss: bool) -> RoutedOutput:
|
def _routed_forward(self, x: Tensor, include_aux_loss: bool) -> RoutedOutput:
|
||||||
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)
|
||||||
if self.norm_topk_prob:
|
if self.norm_topk_prob:
|
||||||
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
||||||
|
|
||||||
aux_loss = None
|
aux_loss = None
|
||||||
|
router_stats = None
|
||||||
if include_aux_loss:
|
if include_aux_loss:
|
||||||
expert_load = F.one_hot(
|
expert_load = F.one_hot(topk_indices, num_classes=E).float()
|
||||||
topk_indices, num_classes=self.n_routed_experts
|
|
||||||
).float()
|
|
||||||
expert_load = expert_load.mean(dim=(0, 1))
|
expert_load = expert_load.mean(dim=(0, 1))
|
||||||
router_prob = router_probs.float().mean(dim=0)
|
router_prob = router_probs.float().mean(dim=0)
|
||||||
aux_loss = self.n_routed_experts * (expert_load * router_prob).sum()
|
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 = self.routed_experts[expert_idx]
|
expert_output = self.routed_experts[expert_idx](flat_tokens[start:end])[
|
||||||
expert_input = x[token_idx]
|
"hidden_states"
|
||||||
expert_output = expert(expert_input)["hidden_states"]
|
]
|
||||||
|
output.index_add_(
|
||||||
|
0,
|
||||||
|
order[start:end] // K,
|
||||||
|
expert_output * flat_weights[start:end],
|
||||||
|
)
|
||||||
|
start = end
|
||||||
|
|
||||||
weights = topk_weights[token_idx, k_idx].unsqueeze(-1)
|
return {
|
||||||
output.index_add_(0, token_idx, expert_output * weights)
|
"hidden_states": output,
|
||||||
|
"aux_loss": aux_loss,
|
||||||
return {"hidden_states": output, "aux_loss": aux_loss}
|
"router_stats": router_stats,
|
||||||
|
}
|
||||||
|
|||||||
@@ -114,6 +114,7 @@ class AutoRegressiveLM(AutoModel):
|
|||||||
use_sdpa_causal_mask = attn_mask is None
|
use_sdpa_causal_mask = attn_mask is None
|
||||||
|
|
||||||
aux_losses = []
|
aux_losses = []
|
||||||
|
router_stats_list = []
|
||||||
for layer in self.layers:
|
for layer in self.layers:
|
||||||
layer_output = layer(
|
layer_output = layer(
|
||||||
x,
|
x,
|
||||||
@@ -123,8 +124,10 @@ class AutoRegressiveLM(AutoModel):
|
|||||||
use_sdpa_causal_mask,
|
use_sdpa_causal_mask,
|
||||||
)
|
)
|
||||||
x = layer_output["hidden_states"]
|
x = layer_output["hidden_states"]
|
||||||
if layer_output["aux_loss"] is not None:
|
stats = layer_output.get("router_stats")
|
||||||
|
if stats is not None:
|
||||||
aux_losses.append(layer_output["aux_loss"])
|
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)
|
||||||
@@ -132,4 +135,5 @@ class AutoRegressiveLM(AutoModel):
|
|||||||
output = {"logits": logits, "hidden_states": hidden_states}
|
output = {"logits": logits, "hidden_states": hidden_states}
|
||||||
if aux_losses:
|
if aux_losses:
|
||||||
output["aux_loss"] = torch.stack(aux_losses).mean()
|
output["aux_loss"] = torch.stack(aux_losses).mean()
|
||||||
|
output["router_stats"] = router_stats_list
|
||||||
return output
|
return output
|
||||||
|
|||||||
@@ -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)
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -88,3 +88,23 @@ 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_aux_loss(ctx):
|
||||||
|
return ctx.strategy._moe_metrics.get("aux_loss")
|
||||||
|
|
||||||
|
|
||||||
|
def ctx_get_router_entropy(ctx):
|
||||||
|
return ctx.strategy._moe_metrics.get("router_entropy")
|
||||||
|
|
||||||
|
|
||||||
|
def ctx_get_dead_expert_fraction(ctx):
|
||||||
|
return ctx.strategy._moe_metrics.get("dead_expert_fraction")
|
||||||
|
|
||||||
|
|
||||||
|
def ctx_get_load_imbalance_mean(ctx):
|
||||||
|
return ctx.strategy._moe_metrics.get("load_imbalance_mean")
|
||||||
|
|
||||||
|
|
||||||
|
def ctx_get_load_imbalance_max(ctx):
|
||||||
|
return ctx.strategy._moe_metrics.get("load_imbalance_max")
|
||||||
|
|||||||
+108
-5
@@ -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, abstractmethod
|
||||||
from typing import Callable, Dict, Optional, TypedDict, 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,6 +9,7 @@ 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
|
||||||
|
|
||||||
@@ -21,6 +22,7 @@ class LossOutput(TypedDict):
|
|||||||
class LogprobsOutput(TypedDict):
|
class LogprobsOutput(TypedDict):
|
||||||
logprobs: Tensor
|
logprobs: Tensor
|
||||||
aux_loss: Optional[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]:
|
||||||
@@ -75,7 +77,11 @@ def get_logprobs(
|
|||||||
logprobs = (token_logprobs * shifted_loss_mask).sum(dim=-1)
|
logprobs = (token_logprobs * shifted_loss_mask).sum(dim=-1)
|
||||||
else:
|
else:
|
||||||
logprobs = token_logprobs * shifted_loss_mask
|
logprobs = token_logprobs * shifted_loss_mask
|
||||||
return {"logprobs": logprobs, "aux_loss": outputs.get("aux_loss")}
|
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:
|
||||||
@@ -94,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.
|
||||||
|
|
||||||
@@ -115,6 +183,7 @@ class BaseStrategy(ABC):
|
|||||||
self.device = device
|
self.device = device
|
||||||
self.executor = kwargs.pop("executor", None)
|
self.executor = kwargs.pop("executor", None)
|
||||||
self.moe_aux_loss_coef = kwargs.pop("moe_aux_loss_coef", 0.01)
|
self.moe_aux_loss_coef = kwargs.pop("moe_aux_loss_coef", 0.01)
|
||||||
|
self._moe_metrics: Dict[str, float] = {}
|
||||||
self.extra_kwargs = kwargs
|
self.extra_kwargs = kwargs
|
||||||
self._rollout_runner = None
|
self._rollout_runner = None
|
||||||
|
|
||||||
@@ -138,6 +207,7 @@ class BaseStrategy(ABC):
|
|||||||
task_loss: Tensor,
|
task_loss: Tensor,
|
||||||
metrics: Dict[str, Tensor],
|
metrics: Dict[str, Tensor],
|
||||||
aux_loss: Optional[Tensor] = None,
|
aux_loss: Optional[Tensor] = None,
|
||||||
|
router_stats: Optional[List[RouterStats]] = None,
|
||||||
) -> LossOutput:
|
) -> LossOutput:
|
||||||
total_loss = task_loss
|
total_loss = task_loss
|
||||||
if aux_loss is not None:
|
if aux_loss is not None:
|
||||||
@@ -145,6 +215,7 @@ class BaseStrategy(ABC):
|
|||||||
total_loss = total_loss + weighted_aux_loss
|
total_loss = total_loss + weighted_aux_loss
|
||||||
metrics["moe_aux_loss"] = aux_loss
|
metrics["moe_aux_loss"] = aux_loss
|
||||||
metrics["moe_aux_loss_weighted"] = weighted_aux_loss
|
metrics["moe_aux_loss_weighted"] = weighted_aux_loss
|
||||||
|
self._refresh_moe_diagnostics(aux_loss, router_stats)
|
||||||
metrics["loss"] = total_loss
|
metrics["loss"] = total_loss
|
||||||
return {
|
return {
|
||||||
"loss": total_loss,
|
"loss": total_loss,
|
||||||
@@ -188,6 +259,20 @@ 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:
|
||||||
@@ -230,6 +315,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__(
|
||||||
@@ -257,7 +343,12 @@ class SEQStrategy(BaseStrategy):
|
|||||||
label_smoothing=self.label_smoothing,
|
label_smoothing=self.label_smoothing,
|
||||||
)
|
)
|
||||||
|
|
||||||
return self._loss_output(loss, {"task_loss": loss}, outputs.get("aux_loss"))
|
return self._loss_output(
|
||||||
|
loss,
|
||||||
|
{"task_loss": loss},
|
||||||
|
outputs.get("aux_loss"),
|
||||||
|
outputs.get("router_stats"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@StrategyFactory.register("sft")
|
@StrategyFactory.register("sft")
|
||||||
@@ -265,6 +356,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__(
|
||||||
@@ -304,7 +396,12 @@ class SFTStrategy(BaseStrategy):
|
|||||||
label_smoothing=self.label_smoothing,
|
label_smoothing=self.label_smoothing,
|
||||||
)
|
)
|
||||||
|
|
||||||
return self._loss_output(loss, {"task_loss": loss}, outputs.get("aux_loss"))
|
return self._loss_output(
|
||||||
|
loss,
|
||||||
|
{"task_loss": loss},
|
||||||
|
outputs.get("aux_loss"),
|
||||||
|
outputs.get("router_stats"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@StrategyFactory.register("dpo")
|
@StrategyFactory.register("dpo")
|
||||||
@@ -379,7 +476,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 self._loss_output(dpo_loss, {"dpo_loss": dpo_loss}, aux_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
|
||||||
@@ -549,6 +651,7 @@ class GRPOStrategy(BaseStrategy):
|
|||||||
task_loss,
|
task_loss,
|
||||||
{"policy_loss": policy_loss, "kl_loss": kl_penalty},
|
{"policy_loss": policy_loss, "kl_loss": kl_penalty},
|
||||||
aux_loss,
|
aux_loss,
|
||||||
|
policy_output.get("router_stats"),
|
||||||
)
|
)
|
||||||
|
|
||||||
def supports_online(self) -> bool:
|
def supports_online(self) -> bool:
|
||||||
|
|||||||
@@ -17,10 +17,15 @@ from astrai.parallel import only_on_rank
|
|||||||
from astrai.parallel.setup import get_current_device
|
from astrai.parallel.setup import get_current_device
|
||||||
from astrai.serialization import Checkpoint
|
from astrai.serialization import Checkpoint
|
||||||
from astrai.trainer.metric_util import (
|
from astrai.trainer.metric_util import (
|
||||||
|
ctx_get_dead_expert_fraction,
|
||||||
ctx_get_grad_norm,
|
ctx_get_grad_norm,
|
||||||
ctx_get_grad_snr,
|
ctx_get_grad_snr,
|
||||||
|
ctx_get_load_imbalance_max,
|
||||||
|
ctx_get_load_imbalance_mean,
|
||||||
ctx_get_loss,
|
ctx_get_loss,
|
||||||
ctx_get_lr,
|
ctx_get_lr,
|
||||||
|
ctx_get_moe_aux_loss,
|
||||||
|
ctx_get_router_entropy,
|
||||||
ctx_get_val_loss,
|
ctx_get_val_loss,
|
||||||
)
|
)
|
||||||
from astrai.trainer.train_context import TrainContext
|
from astrai.trainer.train_context import TrainContext
|
||||||
@@ -257,6 +262,11 @@ 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": ctx_get_moe_aux_loss,
|
||||||
|
"router_entropy": ctx_get_router_entropy,
|
||||||
|
"dead_expert_fraction": ctx_get_dead_expert_fraction,
|
||||||
|
"load_imbalance_mean": ctx_get_load_imbalance_mean,
|
||||||
|
"load_imbalance_max": ctx_get_load_imbalance_max,
|
||||||
}
|
}
|
||||||
|
|
||||||
def _metrics(self, context: TrainContext, names):
|
def _metrics(self, context: TrainContext, names):
|
||||||
|
|||||||
@@ -0,0 +1,74 @@
|
|||||||
|
cmake_minimum_required(VERSION 3.18)
|
||||||
|
project(astrai_kernels LANGUAGES CUDA CXX)
|
||||||
|
|
||||||
|
set(CMAKE_CXX_STANDARD 17)
|
||||||
|
set(CMAKE_CXX_STANDARD_REQUIRED ON)
|
||||||
|
set(CMAKE_CUDA_STANDARD 17)
|
||||||
|
|
||||||
|
find_package(CUDAToolkit REQUIRED)
|
||||||
|
|
||||||
|
if(NOT DEFINED TORCH_HOME)
|
||||||
|
set(TORCH_HOME "$ENV{TORCH_HOME}")
|
||||||
|
endif()
|
||||||
|
if(NOT TORCH_HOME)
|
||||||
|
message(FATAL_ERROR "TORCH_HOME must point at the torch install dir (site-packages/torch)")
|
||||||
|
endif()
|
||||||
|
|
||||||
|
if(NOT DEFINED PYTHON_INCLUDE_DIR)
|
||||||
|
set(PYTHON_INCLUDE_DIR "/usr/include/python${PYTHON_VERSION_MAJOR}.${PYTHON_VERSION_MINOR}")
|
||||||
|
endif()
|
||||||
|
|
||||||
|
if(NOT DEFINED ASTRAI_CUDA_ARCH)
|
||||||
|
if(DEFINED ENV{ASTRAI_CUDA_ARCH})
|
||||||
|
set(ASTRAI_CUDA_ARCH "$ENV{ASTRAI_CUDA_ARCH}")
|
||||||
|
else()
|
||||||
|
set(ASTRAI_CUDA_ARCH 80)
|
||||||
|
endif()
|
||||||
|
endif()
|
||||||
|
|
||||||
|
set(TORCH_LIB_DIR "${TORCH_HOME}/lib")
|
||||||
|
set(CUDA_LIB_DIR "/usr/local/cuda/lib64")
|
||||||
|
|
||||||
|
set(CXX_FLAGS -O3 -funroll-loops)
|
||||||
|
set(NVCC_FLAGS -O3
|
||||||
|
--expt-relaxed-constexpr
|
||||||
|
--use_fast_math
|
||||||
|
"--ptxas-options=-O3,-v"
|
||||||
|
--extra-device-vectorization
|
||||||
|
--threads=16)
|
||||||
|
|
||||||
|
set(TORCH_LIBS
|
||||||
|
"${TORCH_LIB_DIR}/libtorch_python.so"
|
||||||
|
"${TORCH_LIB_DIR}/libtorch_cuda.so"
|
||||||
|
"${TORCH_LIB_DIR}/libc10_cuda.so"
|
||||||
|
"${TORCH_LIB_DIR}/libtorch_cpu.so"
|
||||||
|
"${TORCH_LIB_DIR}/libtorch.so"
|
||||||
|
"${TORCH_LIB_DIR}/libc10.so"
|
||||||
|
CUDA::cudart)
|
||||||
|
|
||||||
|
set(CMAKE_CUDA_ARCHITECTURES "${ASTRAI_CUDA_ARCH}")
|
||||||
|
|
||||||
|
set(KERNELS attn_decode attn_prefill attn_paged_decode attn_paged_prefill rotary_emb)
|
||||||
|
|
||||||
|
foreach(name ${KERNELS})
|
||||||
|
add_library(${name} MODULE "${CMAKE_CURRENT_SOURCE_DIR}/kernels/${name}.cu")
|
||||||
|
|
||||||
|
target_compile_definitions(${name} PRIVATE TORCH_EXTENSION_NAME=${name})
|
||||||
|
|
||||||
|
target_include_directories(${name} PRIVATE
|
||||||
|
"${TORCH_HOME}/include"
|
||||||
|
"${TORCH_HOME}/include/torch/csrc/api/include"
|
||||||
|
"${PYTHON_INCLUDE_DIR}")
|
||||||
|
|
||||||
|
target_link_libraries(${name} PRIVATE ${TORCH_LIBS})
|
||||||
|
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,76 +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("attn_paged_prefill")
|
|
||||||
register("rotary_emb")
|
|
||||||
+16
-49
@@ -9,13 +9,19 @@ enum TensorLayout : int {
|
|||||||
};
|
};
|
||||||
|
|
||||||
|
|
||||||
|
// Unified attention params covering BOTH addressing modes:
|
||||||
|
// - Contiguous K/V: dense [batch, kv_head, kv_len, head_dim] tensors (k/v).
|
||||||
|
// - Paged (SGLang-style): flat pool [size, kv_head, head_dim] + req_to_token.
|
||||||
|
// Each kernel selects the addressing via a KVSource policy (see
|
||||||
|
// attn_kv_source.cuh); a given call only touches the fields of one mode, so
|
||||||
|
// this is a POD shared by both paths rather than two parallel structs that
|
||||||
|
// drift out of sync.
|
||||||
template<typename T, typename AT = float>
|
template<typename T, typename AT = float>
|
||||||
struct AttentionParams {
|
struct AttentionParams {
|
||||||
|
// ---- shared across all paths ----
|
||||||
int batch;
|
int batch;
|
||||||
int q_head;
|
int q_head;
|
||||||
int kv_head;
|
int kv_head;
|
||||||
int q_len;
|
|
||||||
int kv_len;
|
|
||||||
int head_dim;
|
int head_dim;
|
||||||
int use_mask;
|
int use_mask;
|
||||||
int causal_offset; // -1 = non-causal; >=0 = absolute position of first Q token
|
int causal_offset; // -1 = non-causal; >=0 = absolute position of first Q token
|
||||||
@@ -24,54 +30,27 @@ struct AttentionParams {
|
|||||||
|
|
||||||
// Q strides (element offsets for each dim — layout-agnostic)
|
// Q strides (element offsets for each dim — layout-agnostic)
|
||||||
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
|
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
|
||||||
// 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],
|
// 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)
|
// or 4D [batch, n_heads, q_len, kv_len] (head dim broadcasts when stride=0)
|
||||||
int mask_b_stride; // batch stride
|
int mask_b_stride; // batch stride
|
||||||
int mask_h_stride; // head stride (0 = broadcast across heads)
|
int mask_h_stride; // head stride (0 = broadcast across heads)
|
||||||
int mask_q_stride; // q stride (0 = all q rows share)
|
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;
|
const bool* __restrict__ mask;
|
||||||
|
|
||||||
|
const T* __restrict__ q;
|
||||||
T* __restrict__ o;
|
T* __restrict__ o;
|
||||||
AT* __restrict__ o_part;
|
AT* __restrict__ o_part;
|
||||||
AT* __restrict__ ml_part;
|
AT* __restrict__ ml_part;
|
||||||
};
|
|
||||||
|
|
||||||
// ---- PagedAttentionParams ----
|
// ---- contiguous K/V mode ----
|
||||||
// SGLang-style indirect params over a shared KV pool.
|
int q_len;
|
||||||
// k_cache/v_cache: [size, kv_head, head_dim] (bare buffers, no gather).
|
int kv_len;
|
||||||
// req_to_token: [num_reqs, max_context_len] token -> slot.
|
int kv_stride_b, kv_stride_h, kv_stride_l, kv_stride_d;
|
||||||
// req_pool_indices:[batch] rows of the current batch into req_to_token.
|
const T* __restrict__ k;
|
||||||
// kv_indptr: [batch+1] prefix sum of per-request seq_lens (device).
|
const T* __restrict__ v;
|
||||||
// qo_indptr: [batch+1] prefix sum of per-request q_len (prefill) or
|
|
||||||
// nullptr for decode (q_len == 1 everywhere).
|
|
||||||
template<typename T, typename AT = float>
|
|
||||||
struct PagedAttentionParams {
|
|
||||||
int batch;
|
|
||||||
int q_head;
|
|
||||||
int kv_head;
|
|
||||||
int head_dim;
|
|
||||||
int num_splits;
|
|
||||||
int use_mask;
|
|
||||||
int causal_offset; // -1 = non-causal; >=0 = causal (per-request offset
|
|
||||||
// computed inside kernel from kv_indptr/qo_indptr)
|
|
||||||
float scale;
|
|
||||||
|
|
||||||
// Q: [total_q, q_head, head_dim] (3D flattened — no batch dim).
|
// ---- paged (SGLang flat pool) mode ----
|
||||||
// For decode total_q == batch (q_len=1 per request).
|
|
||||||
// For prefill total_q == qo_indptr[batch].
|
|
||||||
int q_stride_l, q_stride_h, q_stride_d;
|
|
||||||
|
|
||||||
// Q: [total_q, q_head, head_dim]
|
|
||||||
const T* __restrict__ q;
|
|
||||||
|
|
||||||
// Flat KV pool: [size, kv_head, head_dim]
|
|
||||||
const T* __restrict__ k_cache;
|
const T* __restrict__ k_cache;
|
||||||
const T* __restrict__ v_cache;
|
const T* __restrict__ v_cache;
|
||||||
|
|
||||||
@@ -84,16 +63,4 @@ struct PagedAttentionParams {
|
|||||||
int max_seq_len; // max per-request seq_len (host-side, for split computation)
|
int max_seq_len; // max per-request seq_len (host-side, for split computation)
|
||||||
int total_q; // total Q tokens across all requests (host-side, for grid)
|
int total_q; // total Q tokens across all requests (host-side, for grid)
|
||||||
int max_q_len; // max per-request q_len (host-side, for prefill grid)
|
int max_q_len; // max per-request q_len (host-side, for prefill grid)
|
||||||
|
|
||||||
// Mask: [batch, max_seq_len] (decode) or [batch, 1, q_len, kv_len]
|
|
||||||
// (prefill, optional). mask_h_stride/mask_q_stride are 0 when those
|
|
||||||
// dims are size 1 (broadcast).
|
|
||||||
int mask_b_stride;
|
|
||||||
int mask_h_stride;
|
|
||||||
int mask_q_stride;
|
|
||||||
const bool* __restrict__ mask;
|
|
||||||
|
|
||||||
T* __restrict__ o;
|
|
||||||
AT* __restrict__ o_part;
|
|
||||||
AT* __restrict__ ml_part;
|
|
||||||
};
|
};
|
||||||
|
|||||||
@@ -2,10 +2,16 @@
|
|||||||
#include <cuda_bf16.h>
|
#include <cuda_bf16.h>
|
||||||
#include <float.h>
|
#include <float.h>
|
||||||
#include "attn_common.h"
|
#include "attn_common.h"
|
||||||
|
#include "attn_kv_source.cuh"
|
||||||
#include "attn_warp_utils.cuh"
|
#include "attn_warp_utils.cuh"
|
||||||
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,15 +21,16 @@ __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_stride_d;
|
||||||
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[q_off + i * p.q_stride_d]);
|
||||||
|
|
||||||
// KV: [batch, kv_head, kv_len, head_dim] — stride-based base
|
|
||||||
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
|
||||||
int mask_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
|
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};
|
||||||
@@ -31,24 +38,25 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
|||||||
extern __shared__ __align__(16) bf16 k_smem[];
|
extern __shared__ __align__(16) bf16 k_smem[];
|
||||||
|
|
||||||
// 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 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::kv_addr(p, kctx, kc, d_dim, true);
|
||||||
k_smem[i] = p.k[g_off];
|
k_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
|
||||||
}
|
}
|
||||||
__syncthreads();
|
__syncthreads();
|
||||||
|
|
||||||
@@ -65,7 +73,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 +82,15 @@ __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
|
// V read via KV policy; when masked (beta == 0) or the slot is
|
||||||
+ lane * hd_per_thread * p.kv_stride_d;
|
// empty the term vanishes, so no extra branches are needed.
|
||||||
for (int i = 0; i < hd_per_thread; i++)
|
for (int i = 0; i < hd_per_thread; i++) {
|
||||||
acc_reg[i] = fmaf(acc_reg[i], alpha,
|
KVAddr a = KV::kv_addr(p, kctx, kv_idx, lane * hd_per_thread + i, true);
|
||||||
__bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta);
|
float vv = a.valid
|
||||||
|
? __bfloat162float(*reinterpret_cast<const bf16*>(a.v))
|
||||||
|
: 0.0f;
|
||||||
|
acc_reg[i] = fmaf(acc_reg[i], alpha, vv * beta);
|
||||||
|
}
|
||||||
m = new_m;
|
m = new_m;
|
||||||
}
|
}
|
||||||
__syncthreads();
|
__syncthreads();
|
||||||
@@ -98,6 +110,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 +140,6 @@ __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_stride_d;
|
||||||
p.o[o_off] = __float2bfloat16(acc * inv);
|
p.o[o_off] = __float2bfloat16(acc * inv);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,19 +2,22 @@
|
|||||||
#include <cfloat>
|
#include <cfloat>
|
||||||
#include <cuda_bf16.h>
|
#include <cuda_bf16.h>
|
||||||
#include "attn_common.h"
|
#include "attn_common.h"
|
||||||
|
#include "attn_kv_source.cuh"
|
||||||
#include "attn_mma_utils.cuh"
|
#include "attn_mma_utils.cuh"
|
||||||
#include "attn_warp_utils.cuh"
|
#include "attn_warp_utils.cuh"
|
||||||
|
|
||||||
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing.
|
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing, unified
|
||||||
// Decode has q_len == 1, so we pack G = q_head/kv_head query heads into the
|
// across contiguous and paged (SGLang flat-pool) K/V via the KV template
|
||||||
// M=16 rows of mma.sync.m16n8k16, turning G independent GEMVs into a single
|
// parameter. Decode has q_len == 1, so we pack G = q_head/kv_head query
|
||||||
// GEMM that reuses each loaded K/V tile across all G heads.
|
// 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,13 +34,16 @@ __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;
|
||||||
@@ -51,13 +57,12 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
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,11 +72,11 @@ __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;
|
||||||
|
KVAddr a = KV::kv_addr(p, kctx, kc, d, valid);
|
||||||
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;
|
cp_async_16_pred(&dK[off], a.k, a.valid);
|
||||||
cp_async_16_pred(&dK[off], &p.k[g_off], valid);
|
cp_async_16_pred(&dV[off], a.v, a.valid);
|
||||||
cp_async_16_pred(&dV[off], &p.v[g_off], valid);
|
|
||||||
}
|
}
|
||||||
cp_async_commit();
|
cp_async_commit();
|
||||||
};
|
};
|
||||||
@@ -96,8 +101,10 @@ __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, 0, 0,
|
||||||
|
|||||||
+107
-125
@@ -1,19 +1,22 @@
|
|||||||
#pragma once
|
#pragma once
|
||||||
// Shared attention dispatchers — used by both production .cu and test .cu.
|
// Shared attention dispatchers — used by both production .cu and test .cu.
|
||||||
// No torch dependency; pure CUDA.
|
// No torch dependency; pure CUDA.
|
||||||
|
//
|
||||||
|
// The paged and contiguous kernels are unified by the KVSource policy
|
||||||
|
// (ContigKV / PagedKV from attn_kv_source.cuh), so each launcher struct
|
||||||
|
// below is templated on KV and the paged dispatch is just the same launcher
|
||||||
|
// instantiated with PagedKV. Only the grid/split math differs, and that is
|
||||||
|
// covered by KV::host_q_len / KV::host_kv_len.
|
||||||
|
|
||||||
#include <cuda_runtime.h>
|
#include <cuda_runtime.h>
|
||||||
#include <algorithm>
|
#include <algorithm>
|
||||||
#include "attn_warp_utils.cuh"
|
#include "attn_warp_utils.cuh"
|
||||||
|
#include "attn_kv_source.cuh"
|
||||||
#include "attn_prefill_split_q.cuh"
|
#include "attn_prefill_split_q.cuh"
|
||||||
#include "attn_decode_split_kv.cuh"
|
#include "attn_decode_split_kv.cuh"
|
||||||
#include "attn_paged_decode_split_kv.cuh"
|
|
||||||
#include "attn_paged_prefill_split_q.cuh"
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
#ifndef ASTRAI_NO_MMA
|
||||||
#include "attn_prefill_split_q_mma.cuh"
|
#include "attn_prefill_split_q_mma.cuh"
|
||||||
#include "attn_decode_split_kv_mma.cuh"
|
#include "attn_decode_split_kv_mma.cuh"
|
||||||
#include "attn_paged_decode_split_kv_mma.cuh"
|
|
||||||
#include "attn_paged_prefill_split_q_mma.cuh"
|
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
// Split-KV: compute number of splits to fill all SMs for small-batch decode.
|
// Split-KV: compute number of splits to fill all SMs for small-batch decode.
|
||||||
@@ -39,7 +42,7 @@ inline int compute_num_splits(int base_blocks, int tiles_total,
|
|||||||
// template <int HEAD_DIM, bool IsCausal, bool HasMask>; HEAD_DIM is forwarded
|
// template <int HEAD_DIM, bool IsCausal, bool HasMask>; HEAD_DIM is forwarded
|
||||||
// as the first template argument so callers only spell it once.
|
// as the first template argument so callers only spell it once.
|
||||||
//
|
//
|
||||||
// Usage: DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_mma, HEAD_DIM, p, group_size);
|
// 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, ...) \
|
#define DISPATCH_CAUSAL_MASK(is_causal, has_mask, FN, HEAD_DIM, ...) \
|
||||||
do { \
|
do { \
|
||||||
if (is_causal) { \
|
if (is_causal) { \
|
||||||
@@ -52,28 +55,39 @@ inline int compute_num_splits(int base_blocks, int tiles_total,
|
|||||||
} while (0)
|
} while (0)
|
||||||
|
|
||||||
// ======================================================================
|
// ======================================================================
|
||||||
// Prefill
|
// Prefill launchers (KV selects ContigKV or PagedKV addressing)
|
||||||
// ======================================================================
|
// ======================================================================
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
#ifndef ASTRAI_NO_MMA
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
template <typename KV>
|
||||||
static inline void launch_prefill_mma(AttentionParams<bf16>& p, cudaStream_t stream) {
|
struct PrefillLauncherMMA {
|
||||||
constexpr int WARPS = 4;
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||||
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
|
constexpr int WARPS = 4;
|
||||||
dim3 grid((p.q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS), p.q_head, p.batch);
|
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
||||||
dim3 block(Traits::NUM_THREADS);
|
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
|
||||||
attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block, 0, stream>>>(p);
|
int q_len = KV::host_q_len(p);
|
||||||
}
|
dim3 grid((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, KV, IsCausal, HasMask>
|
||||||
|
<<<grid, block, 0, stream>>>(p);
|
||||||
|
}
|
||||||
|
};
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
template <typename KV>
|
||||||
static inline void launch_prefill_scalar(AttentionParams<bf16>& p, cudaStream_t stream) {
|
struct PrefillLauncherScalar {
|
||||||
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
|
static void launch(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||||
dim3 block(G, ROWS);
|
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
||||||
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask><<<grid, block, 0, stream>>>(p);
|
int q_len = KV::host_q_len(p);
|
||||||
}
|
dim3 grid((q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
|
||||||
|
dim3 block(G, ROWS);
|
||||||
|
attn_prefill_split_q_kernel_t<HEAD_DIM, KV, G, ROWS, P_BC, IsCausal, HasMask>
|
||||||
|
<<<grid, block, 0, stream>>>(p);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
template <int HEAD_DIM>
|
template <int HEAD_DIM>
|
||||||
static inline void dispatch_prefill(AttentionParams<bf16>& p, cudaStream_t stream) {
|
static inline void dispatch_prefill(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||||
@@ -81,14 +95,34 @@ static inline void dispatch_prefill(AttentionParams<bf16>& p, cudaStream_t strea
|
|||||||
bool has_mask = (p.use_mask && p.mask);
|
bool has_mask = (p.use_mask && p.mask);
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
#ifndef ASTRAI_NO_MMA
|
||||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_mma, HEAD_DIM, p, stream);
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||||
|
PrefillLauncherMMA<ContigKV>::template launch,
|
||||||
|
HEAD_DIM, p, stream);
|
||||||
#else
|
#else
|
||||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_scalar, HEAD_DIM, p, stream);
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||||
|
PrefillLauncherScalar<ContigKV>::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
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||||
|
PrefillLauncherMMA<PagedKV>::template launch,
|
||||||
|
HEAD_DIM, p, stream);
|
||||||
|
#else
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||||
|
PrefillLauncherScalar<PagedKV>::template launch,
|
||||||
|
HEAD_DIM, p, stream);
|
||||||
#endif
|
#endif
|
||||||
}
|
}
|
||||||
|
|
||||||
// ======================================================================
|
// ======================================================================
|
||||||
// Decode
|
// Decode launchers (KV selects ContigKV or PagedKV addressing)
|
||||||
// ======================================================================
|
// ======================================================================
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
#ifndef ASTRAI_NO_MMA
|
||||||
@@ -96,31 +130,41 @@ static inline void dispatch_prefill(AttentionParams<bf16>& p, cudaStream_t strea
|
|||||||
// For D=256, BC=16 also reduces register pressure (fewer Sacc/PV frags),
|
// For D=256, BC=16 also reduces register pressure (fewer Sacc/PV frags),
|
||||||
// enabling STAGES=2 (double-buffer) within the 32KB smem budget — eliminates
|
// enabling STAGES=2 (double-buffer) within the 32KB smem budget — eliminates
|
||||||
// the 176-byte spill that STAGES=1+BC=32 suffered.
|
// the 176-byte spill that STAGES=1+BC=32 suffered.
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
template <typename KV>
|
||||||
static inline void launch_decode_mma(AttentionParams<bf16>& p, int group_size, cudaStream_t stream) {
|
struct DecodeLauncherMMA {
|
||||||
int G = p.q_head / p.kv_head;
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
constexpr int MAX_G = 16;
|
static void launch(AttentionParams<bf16>& p, int group_size, cudaStream_t stream) {
|
||||||
int num_passes = (G + MAX_G - 1) / MAX_G;
|
int G = p.q_head / p.kv_head;
|
||||||
constexpr int BC = 16;
|
constexpr int MAX_G = 16;
|
||||||
int tiles_total = (p.kv_len + BC - 1) / BC;
|
int num_passes = (G + MAX_G - 1) / MAX_G;
|
||||||
p.num_splits = compute_num_splits(p.batch * p.kv_head * num_passes, tiles_total, 2);
|
constexpr int BC = 16;
|
||||||
constexpr int STAGES = 2;
|
int kv_len = KV::host_kv_len(p);
|
||||||
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
int tiles_total = (kv_len + BC - 1) / BC;
|
||||||
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
p.num_splits = compute_num_splits(p.batch * p.kv_head * num_passes, tiles_total, 2);
|
||||||
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32, 0, stream>>>(p);
|
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
|
#endif
|
||||||
|
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
template <typename KV>
|
||||||
static inline void launch_decode_scalar(AttentionParams<bf16>& p, int group_size, cudaStream_t stream) {
|
struct DecodeLauncherScalar {
|
||||||
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
static void launch(AttentionParams<bf16>& p, int group_size, cudaStream_t stream) {
|
||||||
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
int kv_len = KV::host_kv_len(p);
|
||||||
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
|
int chunks_total = (kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||||
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||||
dim3 block(32, g);
|
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
||||||
attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem, stream>>>(p);
|
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, KV, IsCausal, HasMask>
|
||||||
|
<<<grid, block, smem, stream>>>(p);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
template <int HEAD_DIM>
|
template <int HEAD_DIM>
|
||||||
static inline void dispatch_decode(AttentionParams<bf16>& p, cudaStream_t stream) {
|
static inline void dispatch_decode(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||||
@@ -129,95 +173,33 @@ static inline void dispatch_decode(AttentionParams<bf16>& p, cudaStream_t stream
|
|||||||
int group_size = p.q_head / p.kv_head;
|
int group_size = p.q_head / p.kv_head;
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
#ifndef ASTRAI_NO_MMA
|
||||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_mma, HEAD_DIM, p, group_size, stream);
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||||
|
DecodeLauncherMMA<ContigKV>::template launch,
|
||||||
|
HEAD_DIM, p, group_size, stream);
|
||||||
#else
|
#else
|
||||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_scalar, HEAD_DIM, p, group_size, stream);
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||||
|
DecodeLauncherScalar<ContigKV>::template launch,
|
||||||
|
HEAD_DIM, p, group_size, stream);
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
|
attn_decode_combine_kernel<ContigKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
|
||||||
}
|
|
||||||
|
|
||||||
// ======================================================================
|
|
||||||
// Paged Decode (SGLang-style: flat pool + req_to_token + kv_indptr)
|
|
||||||
// ======================================================================
|
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
||||||
static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, cudaStream_t stream) {
|
|
||||||
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.max_seq_len + BC - 1) / BC;
|
|
||||||
p.num_splits = compute_num_splits(p.batch * p.kv_head * num_passes, tiles_total, 2);
|
|
||||||
constexpr int STAGES = 2;
|
|
||||||
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
|
||||||
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
|
||||||
paged_attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask> <<<grid, 32, 0, stream>>>(p);
|
|
||||||
}
|
|
||||||
#endif
|
|
||||||
|
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
||||||
static inline void launch_paged_decode_scalar(PagedAttentionParams<bf16>& p, int group_size, cudaStream_t stream) {
|
|
||||||
int chunks_total = (p.max_seq_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);
|
|
||||||
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, stream>>>(p);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
template <int HEAD_DIM>
|
template <int HEAD_DIM>
|
||||||
static inline void dispatch_paged_decode(PagedAttentionParams<bf16>& p, cudaStream_t stream) {
|
static inline void dispatch_paged_decode(AttentionParams<bf16>& p, cudaStream_t stream) {
|
||||||
bool is_causal = (p.causal_offset >= 0);
|
bool is_causal = (p.causal_offset >= 0);
|
||||||
bool has_mask = (p.use_mask && p.mask);
|
bool has_mask = (p.use_mask && p.mask);
|
||||||
int group_size = p.q_head / p.kv_head;
|
int group_size = p.q_head / p.kv_head;
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
#ifndef ASTRAI_NO_MMA
|
||||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_mma, HEAD_DIM, p, stream);
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||||
|
DecodeLauncherMMA<PagedKV>::template launch,
|
||||||
|
HEAD_DIM, p, group_size, stream);
|
||||||
#else
|
#else
|
||||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_scalar, HEAD_DIM, p, group_size, stream);
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask,
|
||||||
|
DecodeLauncherScalar<PagedKV>::template launch,
|
||||||
|
HEAD_DIM, p, group_size, stream);
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
|
attn_decode_combine_kernel<PagedKV><<<p.batch * p.q_head, p.head_dim, 0, stream>>>(p);
|
||||||
}
|
|
||||||
|
|
||||||
// ======================================================================
|
|
||||||
// Paged Prefill (SGLang-style: flat pool + ragged batch)
|
|
||||||
// ======================================================================
|
|
||||||
|
|
||||||
#ifndef ASTRAI_NO_MMA
|
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
||||||
static inline void launch_paged_prefill_mma(PagedAttentionParams<bf16>& p, cudaStream_t stream) {
|
|
||||||
constexpr int WARPS = 4;
|
|
||||||
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
|
||||||
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
|
|
||||||
int max_q_tiles = (p.max_q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS);
|
|
||||||
dim3 grid(max_q_tiles, p.q_head, p.batch);
|
|
||||||
dim3 block(Traits::NUM_THREADS);
|
|
||||||
paged_attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block, 0, stream>>>(p);
|
|
||||||
}
|
|
||||||
#endif
|
|
||||||
|
|
||||||
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
|
||||||
static inline void launch_paged_prefill_scalar(PagedAttentionParams<bf16>& p, cudaStream_t stream) {
|
|
||||||
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
|
||||||
int max_q_tiles = (p.max_q_len + ROWS - 1) / ROWS;
|
|
||||||
dim3 grid(max_q_tiles, p.q_head, p.batch);
|
|
||||||
dim3 block(G, ROWS);
|
|
||||||
paged_attn_prefill_split_q_kernel<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask>
|
|
||||||
<<<grid, block, 0, stream>>>(p);
|
|
||||||
}
|
|
||||||
|
|
||||||
template <int HEAD_DIM>
|
|
||||||
static inline void dispatch_paged_prefill(PagedAttentionParams<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, launch_paged_prefill_mma, HEAD_DIM, p, stream);
|
|
||||||
#else
|
|
||||||
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_prefill_scalar, HEAD_DIM, p, stream);
|
|
||||||
#endif
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -149,7 +149,7 @@ inline void attn_pack_paged_decode_params(
|
|||||||
c10::optional<torch::Tensor> mask,
|
c10::optional<torch::Tensor> mask,
|
||||||
int64_t causal_offset,
|
int64_t causal_offset,
|
||||||
double scale,
|
double scale,
|
||||||
PagedAttentionParams<T>& p
|
AttentionParams<T>& p
|
||||||
) {
|
) {
|
||||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
|
||||||
@@ -229,7 +229,7 @@ inline void attn_pack_paged_prefill_params(
|
|||||||
int64_t max_q_len,
|
int64_t max_q_len,
|
||||||
int64_t causal_offset,
|
int64_t causal_offset,
|
||||||
double scale,
|
double scale,
|
||||||
PagedAttentionParams<T>& p
|
AttentionParams<T>& p
|
||||||
) {
|
) {
|
||||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,159 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
|
||||||
|
// ============================================================================
|
||||||
|
// KVSource policies — the single dimension along which the paged and
|
||||||
|
// non-paged attention kernels differ. Each kernel is templated on one of
|
||||||
|
// these (ContigKV / PagedKV) and stays fully generic: the policy owns every
|
||||||
|
// place where "where does K/V live" and "what is this request's seq_len"
|
||||||
|
// are answered. All methods are __host__ __device__ so the same policy
|
||||||
|
// serves both the device kernels (addressing, seq_len) and the host-side
|
||||||
|
// launchers (grid / split computation).
|
||||||
|
//
|
||||||
|
// ContigKV: K/V are dense [batch, kv_head, kv_len, head_dim] tensors.
|
||||||
|
// Params fields used: k, v, kv_stride_*, kv_len, q_len,
|
||||||
|
// q_stride_b, 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_stride_l.
|
||||||
|
//
|
||||||
|
// Addressing state that is constant across a whole kernel invocation for one
|
||||||
|
// (batch, kv_head) pair is captured once by make_ctx<HEAD_DIM>() and passed
|
||||||
|
// to kv_addr, so the load loops never redo the hoistable base computation
|
||||||
|
// (e.g. the req_pool_indices global read) element-by-element.
|
||||||
|
// ============================================================================
|
||||||
|
|
||||||
|
// Every policy method is static + callable from both host and device code.
|
||||||
|
#define HOST_DEV_FORCEINLINE static __host__ __device__ __forceinline__
|
||||||
|
|
||||||
|
using bf16 = __nv_bfloat16;
|
||||||
|
|
||||||
|
// Hoisted per-(batch, kv_head) addressing context.
|
||||||
|
struct KVContext {
|
||||||
|
int kv_base; // contig: batch*kv_stride_b + kv_head*kv_stride_h
|
||||||
|
int64_t req_idx; // paged: req_pool_indices[batch]
|
||||||
|
int64_t rtt_stride; // paged: max_context_len
|
||||||
|
int64_t pool_stride; // paged: kv_head * HEAD_DIM
|
||||||
|
int64_t head_off; // paged: kv_head * HEAD_DIM
|
||||||
|
};
|
||||||
|
|
||||||
|
// Per-element K/V global addresses for one (kc, d) position of a K/V tile.
|
||||||
|
// The pointers are ALWAYS the computed addresses (never nullptr) — callers
|
||||||
|
// gate on `valid` (cp.async src_size=0, or a guarded scalar deref). `valid`
|
||||||
|
// starts as "within the request's seq_len"; the paged policy further degrades
|
||||||
|
// it when req_to_token maps the position to a negative slot (empty padding).
|
||||||
|
// This matches the original hand-rolled load loops, where the address was
|
||||||
|
// always formed and the predicate decided whether anything was read.
|
||||||
|
struct KVAddr {
|
||||||
|
const void* k;
|
||||||
|
const void* v;
|
||||||
|
bool valid;
|
||||||
|
};
|
||||||
|
|
||||||
|
// ---- Contiguous K/V ----
|
||||||
|
struct ContigKV {
|
||||||
|
static constexpr bool kPaged = false;
|
||||||
|
|
||||||
|
// host-side length hooks (grid + split computation in the launchers)
|
||||||
|
HOST_DEV_FORCEINLINE int host_q_len(const AttentionParams<bf16>& p) {
|
||||||
|
return p.q_len;
|
||||||
|
}
|
||||||
|
HOST_DEV_FORCEINLINE int host_kv_len(const AttentionParams<bf16>& p) {
|
||||||
|
return p.kv_len;
|
||||||
|
}
|
||||||
|
|
||||||
|
// prefill: element offset of the request's Q rows (kernel adds qrow*q_stride_l)
|
||||||
|
HOST_DEV_FORCEINLINE int q_base(
|
||||||
|
const AttentionParams<bf16>& p, int batch, int q_head) {
|
||||||
|
return batch * p.q_stride_b + q_head * p.q_stride_h;
|
||||||
|
}
|
||||||
|
// decode: same offset (q_len == 1, so there is no row stride component)
|
||||||
|
HOST_DEV_FORCEINLINE int q_decode_base(
|
||||||
|
const AttentionParams<bf16>& p, int batch, int q_head) {
|
||||||
|
return batch * p.q_stride_b + q_head * p.q_stride_h;
|
||||||
|
}
|
||||||
|
|
||||||
|
HOST_DEV_FORCEINLINE int kv_len(const AttentionParams<bf16>& p, int batch) {
|
||||||
|
return p.kv_len;
|
||||||
|
}
|
||||||
|
HOST_DEV_FORCEINLINE int q_len(const AttentionParams<bf16>& p, int batch) {
|
||||||
|
return p.q_len;
|
||||||
|
}
|
||||||
|
HOST_DEV_FORCEINLINE int causal_offset(const AttentionParams<bf16>& p, int batch) {
|
||||||
|
return p.causal_offset;
|
||||||
|
}
|
||||||
|
// decode: exclusive bound of the single query's attend range
|
||||||
|
HOST_DEV_FORCEINLINE int decode_attend_len(const AttentionParams<bf16>& p, int batch) {
|
||||||
|
return (p.kv_len < p.causal_offset + 1) ? p.kv_len : (p.causal_offset + 1);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
HOST_DEV_FORCEINLINE KVContext make_ctx(
|
||||||
|
const AttentionParams<bf16>& p, int batch, int kv_head) {
|
||||||
|
KVContext c = {};
|
||||||
|
c.kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||||
|
return c;
|
||||||
|
}
|
||||||
|
HOST_DEV_FORCEINLINE KVAddr kv_addr(
|
||||||
|
const AttentionParams<bf16>& p, const KVContext& c, int kc, int d, bool valid) {
|
||||||
|
const int g_off = c.kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
|
||||||
|
return {&p.k[g_off], &p.v[g_off], valid};
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
// ---- Paged (SGLang-style flat pool) K/V ----
|
||||||
|
struct PagedKV {
|
||||||
|
static constexpr bool kPaged = true;
|
||||||
|
|
||||||
|
HOST_DEV_FORCEINLINE int host_q_len(const AttentionParams<bf16>& p) {
|
||||||
|
return p.max_q_len;
|
||||||
|
}
|
||||||
|
HOST_DEV_FORCEINLINE int host_kv_len(const AttentionParams<bf16>& p) {
|
||||||
|
return p.max_seq_len;
|
||||||
|
}
|
||||||
|
|
||||||
|
// prefill: Q rows start at qo_indptr[batch] (ragged batch base)
|
||||||
|
HOST_DEV_FORCEINLINE int q_base(
|
||||||
|
const AttentionParams<bf16>& p, int batch, int q_head) {
|
||||||
|
return p.qo_indptr[batch] * p.q_stride_l + q_head * p.q_stride_h;
|
||||||
|
}
|
||||||
|
// decode: Q is [batch, q_head, head_dim], so batch is the outer row
|
||||||
|
HOST_DEV_FORCEINLINE int q_decode_base(
|
||||||
|
const AttentionParams<bf16>& p, int batch, int q_head) {
|
||||||
|
return batch * p.q_stride_l + q_head * p.q_stride_h;
|
||||||
|
}
|
||||||
|
|
||||||
|
HOST_DEV_FORCEINLINE int kv_len(const AttentionParams<bf16>& p, int batch) {
|
||||||
|
return p.kv_indptr[batch + 1] - p.kv_indptr[batch];
|
||||||
|
}
|
||||||
|
HOST_DEV_FORCEINLINE int q_len(const AttentionParams<bf16>& p, int batch) {
|
||||||
|
return p.qo_indptr[batch + 1] - p.qo_indptr[batch];
|
||||||
|
}
|
||||||
|
HOST_DEV_FORCEINLINE int causal_offset(const AttentionParams<bf16>& p, int batch) {
|
||||||
|
return kv_len(p, batch) - q_len(p, batch);
|
||||||
|
}
|
||||||
|
// decode: the query is the last token, so [0, seq_len) IS its causal range
|
||||||
|
HOST_DEV_FORCEINLINE int decode_attend_len(const AttentionParams<bf16>& p, int batch) {
|
||||||
|
return kv_len(p, batch);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
HOST_DEV_FORCEINLINE KVContext make_ctx(
|
||||||
|
const AttentionParams<bf16>& p, int batch, int kv_head) {
|
||||||
|
KVContext c = {};
|
||||||
|
c.req_idx = p.req_pool_indices[batch];
|
||||||
|
c.rtt_stride = (int64_t)p.max_context_len;
|
||||||
|
c.pool_stride = (int64_t)p.kv_head * HEAD_DIM;
|
||||||
|
c.head_off = (int64_t)kv_head * HEAD_DIM;
|
||||||
|
return c;
|
||||||
|
}
|
||||||
|
HOST_DEV_FORCEINLINE KVAddr kv_addr(
|
||||||
|
const AttentionParams<bf16>& p, const KVContext& c, int kc, int d, bool valid) {
|
||||||
|
const int64_t slot = valid ? p.req_to_token[c.req_idx * c.rtt_stride + kc] : 0;
|
||||||
|
const bool ok = valid && (slot >= 0);
|
||||||
|
const int64_t gmem_off = slot * c.pool_stride + c.head_off + d;
|
||||||
|
return {&p.k_cache[gmem_off], &p.v_cache[gmem_off], ok};
|
||||||
|
}
|
||||||
|
};
|
||||||
@@ -16,7 +16,7 @@ torch::Tensor attn_paged_decode(
|
|||||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
auto stream = at::cuda::getCurrentCUDAStream();
|
auto stream = at::cuda::getCurrentCUDAStream();
|
||||||
|
|
||||||
PagedAttentionParams<bf16> p;
|
AttentionParams<bf16> p;
|
||||||
attn_pack_paged_decode_params(q, k_cache, v_cache,
|
attn_pack_paged_decode_params(q, k_cache, v_cache,
|
||||||
req_to_token, req_pool_indices, kv_indptr,
|
req_to_token, req_pool_indices, kv_indptr,
|
||||||
max_seq_len, mask, causal_offset, scale, p);
|
max_seq_len, mask, causal_offset, scale, p);
|
||||||
|
|||||||
@@ -1,151 +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;
|
|
||||||
|
|
||||||
// Scalar paged decode (fallback for sm < 80, no tensor cores).
|
|
||||||
// Reads K/V from flat pool via req_to_token indexing.
|
|
||||||
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;
|
|
||||||
|
|
||||||
const int seq_len = p.kv_indptr[batch + 1] - p.kv_indptr[batch];
|
|
||||||
const int64_t req_idx = p.req_pool_indices[batch];
|
|
||||||
|
|
||||||
float q_reg[8];
|
|
||||||
int q_off = batch * p.q_stride_l + 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 = (seq_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;
|
|
||||||
const int64_t pool_stride = (int64_t)p.kv_head * p.head_dim;
|
|
||||||
const int64_t head_off = (int64_t)kv_head * p.head_dim;
|
|
||||||
const int64_t rtt_stride = (int64_t)p.max_context_len;
|
|
||||||
|
|
||||||
for (int ci = ch_begin; ci < ch_end; ci++) {
|
|
||||||
int chunk_start = ci * PDC_CHUNK;
|
|
||||||
int this_chunk = min(PDC_CHUNK, seq_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;
|
|
||||||
int64_t slot = p.req_to_token[req_idx * rtt_stride + pos];
|
|
||||||
if (slot >= 0) {
|
|
||||||
int64_t off = slot * pool_stride + head_off + 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;
|
|
||||||
}
|
|
||||||
// Decode: the query is the last token, so its valid range [0,
|
|
||||||
// seq_len) IS the causal range. IsCausal is accepted for dispatch
|
|
||||||
// uniformity but must not apply causal_offset masking here.
|
|
||||||
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;
|
|
||||||
int64_t slot = p.req_to_token[req_idx * rtt_stride + pos];
|
|
||||||
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 (slot >= 0) {
|
|
||||||
int64_t v_base = slot * pool_stride + head_off;
|
|
||||||
#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_l + q_head * p.q_stride_h + d * p.q_stride_d;
|
|
||||||
p.o[o_off] = __float2bfloat16(acc * inv);
|
|
||||||
}
|
|
||||||
@@ -1,178 +0,0 @@
|
|||||||
#pragma once
|
|
||||||
#include <cfloat>
|
|
||||||
#include <cuda_bf16.h>
|
|
||||||
#include "attn_common.h"
|
|
||||||
#include "attn_mma_utils.cuh"
|
|
||||||
#include "attn_warp_utils.cuh"
|
|
||||||
|
|
||||||
// SGLang-style split-KV tensor-core decode.
|
|
||||||
//
|
|
||||||
// Reads K/V directly from a flat pool [size, kv_head, head_dim] via
|
|
||||||
// req_to_token indexing — no gather, no page-table dimension.
|
|
||||||
// Each batch element has its own seq_len (from kv_indptr), eliminating
|
|
||||||
// padding waste: short sequences only process the tiles they own.
|
|
||||||
//
|
|
||||||
// For decode (q_len=1), causal masking is implicit — each request attends
|
|
||||||
// to [0, seq_len) which is exactly its valid range. The IsCausal flag
|
|
||||||
// is accepted for dispatch uniformity but does not change maxc.
|
|
||||||
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;
|
|
||||||
|
|
||||||
// Per-request seq_len from device-side kv_indptr — no padding.
|
|
||||||
const int seq_len = p.kv_indptr[batch + 1] - p.kv_indptr[batch];
|
|
||||||
const int64_t req_idx = p.req_pool_indices[batch];
|
|
||||||
|
|
||||||
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];
|
|
||||||
|
|
||||||
const int q_base = batch * p.q_stride_l + 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 = (seq_len + Traits::BC - 1) / Traits::BC;
|
|
||||||
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
|
|
||||||
const int ti_begin = split * tiles_per_split;
|
|
||||||
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
|
|
||||||
|
|
||||||
// Flat pool stride: [size, kv_head, head_dim] — contiguous.
|
|
||||||
const int64_t pool_stride = (int64_t)p.kv_head * Traits::HEAD_DIM;
|
|
||||||
const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM;
|
|
||||||
const int64_t rtt_stride = (int64_t)p.max_context_len;
|
|
||||||
|
|
||||||
// ---- Load tile lambda: SGLang addressing ----
|
|
||||||
// slot = req_to_token[req_idx * max_context_len + kc]
|
|
||||||
// gmem = k_cache[slot * pool_stride + head_off + d]
|
|
||||||
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 < seq_len);
|
|
||||||
if constexpr (HasMask) {
|
|
||||||
valid = valid && p.mask[batch * p.mask_b_stride + kc];
|
|
||||||
}
|
|
||||||
int64_t slot = valid ? p.req_to_token[req_idx * rtt_stride + kc] : 0;
|
|
||||||
valid = valid && (slot >= 0);
|
|
||||||
int64_t gmem_base = slot * pool_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();
|
|
||||||
};
|
|
||||||
|
|
||||||
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;
|
|
||||||
|
|
||||||
// For decode, maxc = seq_len regardless of IsCausal — the valid
|
|
||||||
// range [0, seq_len) IS the causal range (query is the last token).
|
|
||||||
mma_softmax_tile<Traits, HasMask>(kv0, seq_len, seq_len,
|
|
||||||
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 {
|
|
||||||
for (int i = 0; i < ntiles; i++)
|
|
||||||
load_tile(ti_begin + i, i);
|
|
||||||
cp_async_wait_group<0>();
|
|
||||||
__syncwarp();
|
|
||||||
for (int it = 0; it < ntiles; it++)
|
|
||||||
process_tile(it, it);
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---- write partials ----
|
|
||||||
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;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -17,7 +17,7 @@ torch::Tensor attn_paged_prefill(
|
|||||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
auto stream = at::cuda::getCurrentCUDAStream();
|
auto stream = at::cuda::getCurrentCUDAStream();
|
||||||
|
|
||||||
PagedAttentionParams<bf16> p;
|
AttentionParams<bf16> p;
|
||||||
attn_pack_paged_prefill_params(q, k_cache, v_cache,
|
attn_pack_paged_prefill_params(q, k_cache, v_cache,
|
||||||
req_to_token, req_pool_indices,
|
req_to_token, req_pool_indices,
|
||||||
kv_indptr, qo_indptr, mask,
|
kv_indptr, qo_indptr, mask,
|
||||||
|
|||||||
@@ -1,126 +0,0 @@
|
|||||||
#pragma once
|
|
||||||
#include <cuda_bf16.h>
|
|
||||||
#include <float.h>
|
|
||||||
#include "attn_common.h"
|
|
||||||
|
|
||||||
using bf16 = __nv_bfloat16;
|
|
||||||
|
|
||||||
// Scalar paged prefill (fallback for sm < 80, no tensor cores).
|
|
||||||
// Reads K/V from a flat pool via req_to_token, supports ragged batches
|
|
||||||
// via qo_indptr + kv_indptr. Mirrors the split-Q MMA kernel's indexing:
|
|
||||||
// grid (max_q_tiles, q_head, batch), block (G, ROWS).
|
|
||||||
//
|
|
||||||
// HasMask: 4D mask [batch, 1, q_len, kv_len] (True=keep), columns are
|
|
||||||
// request-local kv positions. q_head is the q-index (mask_h broadcast).
|
|
||||||
//
|
|
||||||
// group_reduce_sum<G> is provided by attn_prefill_split_q.cuh (already
|
|
||||||
// included via the dispatcher).
|
|
||||||
template <int HEAD_DIM, int G, int ROWS, int P_BC, bool IsCausal, bool HasMask>
|
|
||||||
__global__ void paged_attn_prefill_split_q_kernel(PagedAttentionParams<bf16> p) {
|
|
||||||
constexpr int DPT = HEAD_DIM / G;
|
|
||||||
|
|
||||||
const int q_tile = blockIdx.x;
|
|
||||||
const int q_head = blockIdx.y;
|
|
||||||
const int req_b = blockIdx.z;
|
|
||||||
const int gpos = threadIdx.x; // 0..G-1 (d-chunk)
|
|
||||||
const int row = threadIdx.y; // 0..ROWS-1 (q row within tile)
|
|
||||||
const int q_row = q_tile * ROWS + row;
|
|
||||||
|
|
||||||
const int seq_len = p.kv_indptr[req_b + 1] - p.kv_indptr[req_b];
|
|
||||||
const int q_len = p.qo_indptr[req_b + 1] - p.qo_indptr[req_b];
|
|
||||||
const int causal_off = seq_len - q_len;
|
|
||||||
const int64_t req_idx = p.req_pool_indices[req_b];
|
|
||||||
const int kv_head = q_head / (p.q_head / p.kv_head);
|
|
||||||
|
|
||||||
__shared__ __align__(16) bf16 sK[P_BC * HEAD_DIM];
|
|
||||||
__shared__ __align__(16) bf16 sV[P_BC * HEAD_DIM];
|
|
||||||
|
|
||||||
// Q base: absolute token = qo_indptr[req_b] + q_row.
|
|
||||||
float qreg[DPT];
|
|
||||||
if (q_row < q_len) {
|
|
||||||
int q_off = (p.qo_indptr[req_b] + q_row) * p.q_stride_l
|
|
||||||
+ q_head * p.q_stride_h + gpos * DPT * p.q_stride_d;
|
|
||||||
#pragma unroll
|
|
||||||
for (int i = 0; i < DPT; i++)
|
|
||||||
qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
|
|
||||||
}
|
|
||||||
|
|
||||||
float m = -FLT_MAX, l = 0.0f, acc[DPT];
|
|
||||||
#pragma unroll
|
|
||||||
for (int i = 0; i < DPT; i++) acc[i] = 0.0f;
|
|
||||||
|
|
||||||
const int64_t pool_stride = (int64_t)p.kv_head * p.head_dim;
|
|
||||||
const int64_t head_off = (int64_t)kv_head * p.head_dim;
|
|
||||||
const int64_t rtt_stride = (int64_t)p.max_context_len;
|
|
||||||
const int mask_base = req_b * p.mask_b_stride + q_head * p.mask_h_stride
|
|
||||||
+ q_row * p.mask_q_stride;
|
|
||||||
|
|
||||||
int tiles = (seq_len + P_BC - 1) / P_BC;
|
|
||||||
int tt = G * ROWS;
|
|
||||||
int lid = row * G + gpos;
|
|
||||||
|
|
||||||
// Each warp holds (32/G) q-rows; reduce only within this row's G lanes.
|
|
||||||
int lane_in_warp = lid & 31;
|
|
||||||
unsigned gmask = (G == 32) ? 0xFFFFFFFFu
|
|
||||||
: (((1u << G) - 1u) << (lane_in_warp & ~(G - 1)));
|
|
||||||
|
|
||||||
for (int ti = 0; ti < tiles; ti++) {
|
|
||||||
int kv0 = ti * P_BC;
|
|
||||||
int tlen = min(P_BC, seq_len - kv0);
|
|
||||||
|
|
||||||
// Load K/V tile into shared memory via req_to_token (request-local pos).
|
|
||||||
for (int i = lid; i < tlen * HEAD_DIM; i += tt) {
|
|
||||||
int s = i / HEAD_DIM, d_dim = i % HEAD_DIM;
|
|
||||||
int pos = kv0 + s;
|
|
||||||
int64_t slot = p.req_to_token[req_idx * rtt_stride + pos];
|
|
||||||
int64_t off = slot * pool_stride + head_off + d_dim;
|
|
||||||
sK[i] = (slot >= 0) ? p.k_cache[off] : __float2bfloat16(0.0f);
|
|
||||||
sV[i] = (slot >= 0) ? p.v_cache[off] : __float2bfloat16(0.0f);
|
|
||||||
}
|
|
||||||
__syncthreads();
|
|
||||||
|
|
||||||
int lim = tlen;
|
|
||||||
if constexpr (IsCausal) {
|
|
||||||
if (q_row < q_len) {
|
|
||||||
int ep = causal_off + q_row + 1;
|
|
||||||
if (kv0 >= ep)
|
|
||||||
lim = 0;
|
|
||||||
else if (kv0 + tlen > ep)
|
|
||||||
lim = ep - kv0;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
for (int s = 0; s < lim; s++) {
|
|
||||||
bool keep = true;
|
|
||||||
if constexpr (HasMask) {
|
|
||||||
if (q_row < q_len && !p.mask[mask_base + kv0 + s])
|
|
||||||
keep = false;
|
|
||||||
}
|
|
||||||
float w = 0.0f;
|
|
||||||
#pragma unroll
|
|
||||||
for (int i = 0; i < DPT; i++)
|
|
||||||
w += qreg[i] * __bfloat162float(sK[s * HEAD_DIM + gpos * DPT + i]);
|
|
||||||
w = group_reduce_sum<G>(w, gmask) * p.scale;
|
|
||||||
if (!keep) w = -FLT_MAX;
|
|
||||||
|
|
||||||
float nm = fmaxf(m, w);
|
|
||||||
float alpha = __expf(m - nm);
|
|
||||||
float beta = __expf(w - nm);
|
|
||||||
l = l * alpha + beta;
|
|
||||||
#pragma unroll
|
|
||||||
for (int i = 0; i < DPT; i++)
|
|
||||||
acc[i] = acc[i] * alpha
|
|
||||||
+ __bfloat162float(sV[s * HEAD_DIM + gpos * DPT + i]) * beta;
|
|
||||||
m = nm;
|
|
||||||
}
|
|
||||||
__syncthreads();
|
|
||||||
}
|
|
||||||
|
|
||||||
if (q_row >= q_len) return;
|
|
||||||
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
|
||||||
int o_off = (p.qo_indptr[req_b] + q_row) * p.q_stride_l
|
|
||||||
+ q_head * p.q_stride_h + gpos * DPT * p.q_stride_d;
|
|
||||||
#pragma unroll
|
|
||||||
for (int i = 0; i < DPT; i++)
|
|
||||||
p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * inv);
|
|
||||||
}
|
|
||||||
@@ -1,164 +0,0 @@
|
|||||||
#pragma once
|
|
||||||
#include <cfloat>
|
|
||||||
#include <cuda_bf16.h>
|
|
||||||
#include "attn_common.h"
|
|
||||||
#include "attn_mma_utils.cuh"
|
|
||||||
|
|
||||||
// SGLang-style split-Q tensor-core prefill.
|
|
||||||
//
|
|
||||||
// Reads K/V directly from a flat pool [size, kv_head, head_dim] via
|
|
||||||
// req_to_token — no gather, no temporary tensor. Supports ragged batches:
|
|
||||||
// each request has its own q_len and kv_len, addressed via qo_indptr and
|
|
||||||
// kv_indptr.
|
|
||||||
//
|
|
||||||
// Grid: (max_q_tiles, q_head, batch) — one batch element per blockIdx.z.
|
|
||||||
// Blocks beyond a request's q_len exit early after writing sentinel-free
|
|
||||||
// no-ops. This avoids the binary-search approach and guarantees every Q
|
|
||||||
// token is covered, even when q_len < BR*WARPS (e.g. decode-like prefill).
|
|
||||||
//
|
|
||||||
// Q layout: [total_q, q_head, head_dim] (3D, flattened across requests).
|
|
||||||
// O layout: same as Q.
|
|
||||||
//
|
|
||||||
// IsCausal is a compile-time bool. When true, each Q row qi (within its
|
|
||||||
// request) attends to [0, causal_offset_b + qi + 1) where
|
|
||||||
// causal_offset_b = kv_len_b - q_len_b (position of first Q token).
|
|
||||||
template <typename Traits, bool IsCausal, bool HasMask>
|
|
||||||
__global__ void paged_attn_prefill_split_q_mma_kernel(PagedAttentionParams<bf16> p) {
|
|
||||||
const int warp = threadIdx.x / 32;
|
|
||||||
const int lane = threadIdx.x % 32;
|
|
||||||
const int gid = lane >> 2;
|
|
||||||
const int tid4 = lane & 3;
|
|
||||||
|
|
||||||
const int q_head = blockIdx.y;
|
|
||||||
const int req_b = blockIdx.z;
|
|
||||||
const int qrow0 = (blockIdx.x * Traits::WARPS + warp) * Traits::BR;
|
|
||||||
|
|
||||||
const int seq_len = p.kv_indptr[req_b + 1] - p.kv_indptr[req_b];
|
|
||||||
const int q_len = p.qo_indptr[req_b + 1] - p.qo_indptr[req_b];
|
|
||||||
const int causal_off = seq_len - q_len;
|
|
||||||
const int64_t req_idx = p.req_pool_indices[req_b];
|
|
||||||
|
|
||||||
// No per-warp early exit — all warps must participate in __syncthreads.
|
|
||||||
// Warps beyond q_len get zero-filled Q frags (va=vb=false) and skip output.
|
|
||||||
const int kv_head = q_head / (p.q_head / p.kv_head);
|
|
||||||
|
|
||||||
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
|
||||||
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
|
||||||
|
|
||||||
// Q base: offset by qo_indptr[req_b] to get absolute token address.
|
|
||||||
const int q_base = p.qo_indptr[req_b] * p.q_stride_l + q_head * p.q_stride_h;
|
|
||||||
const int qra = qrow0 + gid;
|
|
||||||
const int qrb = qrow0 + gid + 8;
|
|
||||||
const bool va = qra < q_len, vb = qrb < q_len;
|
|
||||||
unsigned Qa[Traits::KD][4];
|
|
||||||
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_l, p.q_stride_d,
|
|
||||||
qra, qrb, va, vb, tid4, Qa);
|
|
||||||
|
|
||||||
float Oacc[Traits::DN8][4];
|
|
||||||
#pragma unroll
|
|
||||||
for (int j = 0; j < Traits::DN8; j++)
|
|
||||||
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
|
||||||
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
|
||||||
|
|
||||||
const int64_t pool_stride = (int64_t)p.kv_head * Traits::HEAD_DIM;
|
|
||||||
const int64_t head_off = (int64_t)kv_head * Traits::HEAD_DIM;
|
|
||||||
const int64_t rtt_stride = (int64_t)p.max_context_len;
|
|
||||||
|
|
||||||
const int tiles = (seq_len + Traits::BC - 1) / Traits::BC;
|
|
||||||
const int qr0 = qrow0 + gid;
|
|
||||||
const int qr1 = qrow0 + gid + 8;
|
|
||||||
|
|
||||||
// Causal tile-skip (dead code when IsCausal == false)
|
|
||||||
const int max_kv = qrow0 + Traits::BR - 1 + causal_off;
|
|
||||||
const int block_max_kv =
|
|
||||||
blockIdx.x * Traits::WARPS * Traits::BR + Traits::WARPS * Traits::BR - 1
|
|
||||||
+ causal_off;
|
|
||||||
|
|
||||||
int t_end = tiles - 1;
|
|
||||||
if constexpr (IsCausal) {
|
|
||||||
int bt = block_max_kv / Traits::BC;
|
|
||||||
if (bt < t_end) t_end = bt;
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---- Load tile lambda: SGLang addressing ----
|
|
||||||
auto load_tile = [&](int ti, int buf) {
|
|
||||||
int kv0 = ti * Traits::BC;
|
|
||||||
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
|
||||||
bf16* dV = sV + buf * Traits::BC * Traits::LD;
|
|
||||||
#pragma unroll
|
|
||||||
for (int i = threadIdx.x * Traits::VEC; i < Traits::TOTAL;
|
|
||||||
i += Traits::NUM_THREADS * Traits::VEC) {
|
|
||||||
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
|
||||||
int kc = kv0 + r;
|
|
||||||
bool valid = kc < seq_len;
|
|
||||||
int64_t slot = valid ? p.req_to_token[req_idx * rtt_stride + kc] : 0;
|
|
||||||
valid = valid && (slot >= 0);
|
|
||||||
int64_t gmem_base = slot * pool_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();
|
|
||||||
};
|
|
||||||
|
|
||||||
// ---- Prologue + main loop (FA2-style double-buffer) ----
|
|
||||||
load_tile(0, 0);
|
|
||||||
|
|
||||||
for (int ti = 0; ti <= t_end; ti++) {
|
|
||||||
int buf = ti & 1;
|
|
||||||
|
|
||||||
cp_async_wait_group<0>();
|
|
||||||
__syncthreads();
|
|
||||||
if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1);
|
|
||||||
|
|
||||||
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
|
|
||||||
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
|
|
||||||
int kv0 = ti * Traits::BC;
|
|
||||||
|
|
||||||
if (!IsCausal || kv0 <= max_kv) {
|
|
||||||
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 maxc0 = IsCausal ? min(seq_len, causal_off + qr0 + 1)
|
|
||||||
: seq_len;
|
|
||||||
int maxc1 = IsCausal ? min(seq_len, causal_off + qr1 + 1)
|
|
||||||
: seq_len;
|
|
||||||
// HasMask: mask[batch, q_head, qi, kc] — kc is request-local.
|
|
||||||
mma_softmax_tile<Traits, HasMask>(kv0, maxc0, maxc1,
|
|
||||||
qr0, qr1,
|
|
||||||
p.mask_b_stride, p.mask_h_stride,
|
|
||||||
p.mask_q_stride,
|
|
||||||
req_b, q_head,
|
|
||||||
p.mask,
|
|
||||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
|
||||||
|
|
||||||
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---- write output: packed bf16x2 stores ----
|
|
||||||
float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f;
|
|
||||||
float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f;
|
|
||||||
const int o_base = p.qo_indptr[req_b] * p.q_stride_l + q_head * p.q_stride_h;
|
|
||||||
#pragma unroll
|
|
||||||
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
|
||||||
int d = dn8 * 8 + 2 * tid4;
|
|
||||||
if (qr0 < q_len) {
|
|
||||||
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0,
|
|
||||||
Oacc[dn8][1] * rl0);
|
|
||||||
*reinterpret_cast<__nv_bfloat162*>(
|
|
||||||
&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v;
|
|
||||||
}
|
|
||||||
if (qr1 < q_len) {
|
|
||||||
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1,
|
|
||||||
Oacc[dn8][3] * rl1);
|
|
||||||
*reinterpret_cast<__nv_bfloat162*>(
|
|
||||||
&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -2,13 +2,15 @@
|
|||||||
#include <cfloat>
|
#include <cfloat>
|
||||||
#include <cuda_bf16.h>
|
#include <cuda_bf16.h>
|
||||||
#include "attn_common.h"
|
#include "attn_common.h"
|
||||||
|
#include "attn_kv_source.cuh"
|
||||||
|
|
||||||
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>
|
template <int G>
|
||||||
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
|
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
|
||||||
@@ -30,7 +32,7 @@ __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 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;
|
||||||
|
|
||||||
@@ -41,16 +43,21 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
|||||||
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 = KV::q_len(p, batch);
|
||||||
|
const int causal_off = KV::causal_offset(p, batch);
|
||||||
|
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 = KV::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_stride_l + gpos * DPT * p.q_stride_d;
|
||||||
+ 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[q_off + i * p.q_stride_d]);
|
||||||
@@ -62,10 +69,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 +80,24 @@ __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;
|
KVAddr a = KV::kv_addr(p, kctx, kc, d_dim, true);
|
||||||
sK[i] = p.k[g_off];
|
sK[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
|
||||||
sV[i] = p.v[g_off];
|
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)
|
||||||
@@ -138,9 +144,8 @@ __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_stride_l + gpos * DPT * p.q_stride_d;
|
||||||
+ 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++)
|
||||||
|
|||||||
@@ -2,17 +2,21 @@
|
|||||||
#include <cfloat>
|
#include <cfloat>
|
||||||
#include <cuda_bf16.h>
|
#include <cuda_bf16.h>
|
||||||
#include "attn_common.h"
|
#include "attn_common.h"
|
||||||
|
#include "attn_kv_source.cuh"
|
||||||
#include "attn_mma_utils.cuh"
|
#include "attn_mma_utils.cuh"
|
||||||
|
|
||||||
// Tensor-core prefill flash attention (raw mma.sync PTX).
|
// 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).
|
||||||
//
|
//
|
||||||
|
// 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 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;
|
||||||
@@ -24,16 +28,22 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
const int kv_head = q_head / (p.q_head / p.kv_head);
|
const int kv_head = q_head / (p.q_head / p.kv_head);
|
||||||
const int qrow0 = (blockIdx.x * Traits::WARPS + warp) * Traits::BR;
|
const int qrow0 = (blockIdx.x * Traits::WARPS + warp) * 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 = KV::q_len(p, batch);
|
||||||
|
const int causal_off = KV::causal_offset(p, batch);
|
||||||
|
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).
|
||||||
__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 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 = KV::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 + q_base, p.q_stride_l, p.q_stride_d,
|
||||||
qra, qrb, va, vb, tid4, Qa);
|
qra, qrb, va, vb, tid4, Qa);
|
||||||
@@ -44,17 +54,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;
|
const int max_kv = qrow0 + Traits::BR - 1 + causal_off;
|
||||||
const int block_max_kv =
|
const int block_max_kv =
|
||||||
blockIdx.x * Traits::WARPS * Traits::BR + Traits::WARPS * Traits::BR - 1
|
blockIdx.x * Traits::WARPS * Traits::BR + Traits::WARPS * Traits::BR - 1
|
||||||
+ p.causal_offset;
|
+ causal_off;
|
||||||
|
|
||||||
int t_end = tiles - 1;
|
int t_end = tiles - 1;
|
||||||
if constexpr (IsCausal) {
|
if constexpr (IsCausal) {
|
||||||
@@ -62,7 +70,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,11 +80,11 @@ __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;
|
||||||
|
KVAddr a = KV::kv_addr(p, kctx, kc, d, valid);
|
||||||
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;
|
cp_async_16_pred(&dK[off], a.k, a.valid);
|
||||||
cp_async_16_pred(&dK[off], &p.k[g_off], valid);
|
cp_async_16_pred(&dV[off], a.v, a.valid);
|
||||||
cp_async_16_pred(&dV[off], &p.v[g_off], valid);
|
|
||||||
}
|
}
|
||||||
cp_async_commit();
|
cp_async_commit();
|
||||||
};
|
};
|
||||||
@@ -108,10 +116,10 @@ __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_q_stride,
|
||||||
@@ -126,17 +134,17 @@ __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 = KV::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 (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[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v;
|
||||||
}
|
}
|
||||||
if (qr1 < p.q_len) {
|
if (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*>(
|
||||||
|
|||||||
@@ -212,7 +212,7 @@ static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
|
|||||||
B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref);
|
B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref);
|
||||||
|
|
||||||
// Kernel launch
|
// Kernel launch
|
||||||
PagedAttentionParams<bf16> p;
|
AttentionParams<bf16> p;
|
||||||
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
||||||
p.head_dim = HEAD_DIM; p.total_q = B;
|
p.head_dim = HEAD_DIM; p.total_q = B;
|
||||||
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||||
@@ -347,7 +347,7 @@ static int run_decode_mask_test(int B, int Hq, int Hkv, int max_seq,
|
|||||||
h_mask, max_sl,
|
h_mask, max_sl,
|
||||||
B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref);
|
B, Hq, Hkv, HEAD_DIM, max_ctx, h_o_ref);
|
||||||
|
|
||||||
PagedAttentionParams<bf16> p;
|
AttentionParams<bf16> p;
|
||||||
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
||||||
p.head_dim = HEAD_DIM; p.total_q = B;
|
p.head_dim = HEAD_DIM; p.total_q = B;
|
||||||
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||||
@@ -480,7 +480,7 @@ static int run_prefill_test(int B, int Hq, int Hkv,
|
|||||||
B, Hq, Hkv, HEAD_DIM, max_ctx, causal, h_o_ref);
|
B, Hq, Hkv, HEAD_DIM, max_ctx, causal, h_o_ref);
|
||||||
|
|
||||||
// Kernel launch
|
// Kernel launch
|
||||||
PagedAttentionParams<bf16> p;
|
AttentionParams<bf16> p;
|
||||||
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
||||||
p.head_dim = HEAD_DIM; p.total_q = total_q;
|
p.head_dim = HEAD_DIM; p.total_q = total_q;
|
||||||
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||||
@@ -617,7 +617,7 @@ static int run_prefill_mask_test(int Hq, int Hkv, int q_len, int seed) {
|
|||||||
h_mask, q_len, q_len,
|
h_mask, q_len, q_len,
|
||||||
B, Hq, Hkv, HEAD_DIM, max_ctx, 0, h_o_ref);
|
B, Hq, Hkv, HEAD_DIM, max_ctx, 0, h_o_ref);
|
||||||
|
|
||||||
PagedAttentionParams<bf16> p;
|
AttentionParams<bf16> p;
|
||||||
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
||||||
p.head_dim = HEAD_DIM; p.total_q = total_q;
|
p.head_dim = HEAD_DIM; p.total_q = total_q;
|
||||||
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||||
@@ -707,7 +707,7 @@ static void bench_decode(int B, int Hq, int Hkv, int seq_len) {
|
|||||||
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_len;
|
for (int b = 0; b < B; b++) h_kvi[b + 1] = h_kvi[b] + seq_len;
|
||||||
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
|
cudaMemcpy(d_kvi, h_kvi, sz_kvi, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
PagedAttentionParams<bf16> p;
|
AttentionParams<bf16> p;
|
||||||
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
||||||
p.head_dim = HEAD_DIM; p.total_q = B;
|
p.head_dim = HEAD_DIM; p.total_q = B;
|
||||||
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||||
@@ -784,7 +784,7 @@ static void bench_prefill(int B, int Hq, int Hkv, int q_len, int kv_len, int cau
|
|||||||
for (int b = 0; b < B; b++) h_qoi[b + 1] = h_qoi[b] + q_len;
|
for (int b = 0; b < B; b++) h_qoi[b + 1] = h_qoi[b] + q_len;
|
||||||
cudaMemcpy(d_qoi, h_qoi, sz_qoi, cudaMemcpyHostToDevice);
|
cudaMemcpy(d_qoi, h_qoi, sz_qoi, cudaMemcpyHostToDevice);
|
||||||
|
|
||||||
PagedAttentionParams<bf16> p;
|
AttentionParams<bf16> p;
|
||||||
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
p.batch = B; p.q_head = Hq; p.kv_head = Hkv;
|
||||||
p.head_dim = HEAD_DIM; p.total_q = total_q;
|
p.head_dim = HEAD_DIM; p.total_q = total_q;
|
||||||
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
p.q_stride_l = Hq * HEAD_DIM; p.q_stride_h = HEAD_DIM; p.q_stride_d = 1;
|
||||||
|
|||||||
@@ -83,7 +83,7 @@ static int run_decode_test(int B, int Hq, int Hk, int sl, int D, int causal) {
|
|||||||
for (size_t i=0;i<nQ;i++){
|
for (size_t i=0;i<nQ;i++){
|
||||||
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||||
if(err>max_abs_err) max_abs_err=err;
|
if(err>max_abs_err) max_abs_err=err;
|
||||||
float rel=err/fmaxf(fabsf(ref[i]), 1e-8f);
|
float rel=err/fmaxf(fabsf(ref[i]), 1e-4f);
|
||||||
if(rel>max_rel_err) max_rel_err=rel;
|
if(rel>max_rel_err) max_rel_err=rel;
|
||||||
}
|
}
|
||||||
const float atol=0.01f, rtol=0.01f;
|
const float atol=0.01f, rtol=0.01f;
|
||||||
@@ -206,7 +206,7 @@ static int run_prefill_test(int B, int Hq, int Hk, int ql, int kl, int D, int ca
|
|||||||
for (size_t i=0;i<nQ;i++) {
|
for (size_t i=0;i<nQ;i++) {
|
||||||
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
float err=fabsf(bf2f(hOut[i])-ref[i]);
|
||||||
if(err>max_abs_err) max_abs_err=err;
|
if(err>max_abs_err) max_abs_err=err;
|
||||||
float rel=err/fmaxf(fabsf(ref[i]), 1e-8f);
|
float rel=err/fmaxf(fabsf(ref[i]), 1e-4f);
|
||||||
if(rel>max_rel_err) max_rel_err=rel;
|
if(rel>max_rel_err) max_rel_err=rel;
|
||||||
}
|
}
|
||||||
const float atol=0.01f, rtol=0.01f;
|
const float atol=0.01f, rtol=0.01f;
|
||||||
|
|||||||
@@ -120,7 +120,7 @@ inline void set_default_strides(P& p) {
|
|||||||
p.mask_q_stride = 0;
|
p.mask_q_stride = 0;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Set default Q strides for contiguous b h l d layout on PagedAttentionParams.
|
// Set default Q strides for a paged decode params struct.
|
||||||
template<typename P>
|
template<typename P>
|
||||||
inline void set_default_paged_strides(P& p) {
|
inline void set_default_paged_strides(P& p) {
|
||||||
p.q_stride_b = p.q_head * p.q_len * p.head_dim;
|
p.q_stride_b = p.q_head * p.q_len * p.head_dim;
|
||||||
|
|||||||
+15
-11
@@ -14,7 +14,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">
|
||||||
@@ -33,7 +33,7 @@
|
|||||||
|
|
||||||
## 📖 目录
|
## 📖 目录
|
||||||
|
|
||||||
- [特性](#特性)
|
- [项目概览](#项目概览)
|
||||||
- [快速上手](#快速上手)
|
- [快速上手](#快速上手)
|
||||||
- [演示](#演示)
|
- [演示](#演示)
|
||||||
- [文档](#文档)
|
- [文档](#文档)
|
||||||
@@ -46,15 +46,19 @@
|
|||||||
<a id="chinese"></a>
|
<a id="chinese"></a>
|
||||||
## 中文
|
## 中文
|
||||||
|
|
||||||
### 特性
|
### 项目概览
|
||||||
|
|
||||||
- 🚀 **高性能**: 训练与推理双向优化,高效并行。
|
AstrAI 是一个覆盖模型构建、训练、评测与部署的端到端 Transformer 框架。项目以精简的 PyTorch 代码实现完整模型生命周期,包括声明式数据预处理、分布式训练、连续批处理推理,以及兼容 OpenAI 和 Anthropic 的服务接口。
|
||||||
- 🔧 **灵活**: 支持 seq/sft/dpo/grpo 多种训练方式,可定制模型架构。
|
|
||||||
- 💡 **易用**: 简洁的 API 与丰富的示例、演示。
|
| 领域 | 能力 |
|
||||||
- 📦 **轻量**: 依赖少,部署简单。
|
|---|---|
|
||||||
- 🔬 **研究友好**: 模块化设计,便于实验新想法。
|
| **模型** | 自回归语言模型与嵌入模型,支持 GQA、MLA、MoE、RoPE,以及可扩展的 Attention/FFN 组件 |
|
||||||
- 🤗 **HuggingFace 风格 API**: 类 HuggingFace 的 AutoModel/AutoTokenizer 接口,方便加载模型和分词器。
|
| **训练** | 预训练(`seq`)、监督微调(`sft`)、DPO 和 GRPO,支持梯度累积、检查点、DDP 与 FSDP |
|
||||||
- 🔌 **双 API 兼容**: 同时支持 OpenAI 和 Anthropic 聊天补全 API,开箱即用。
|
| **数据** | 声明式 JSON 预处理、可配置掩码与样本打包、二进制/JSONL 存储和流式数据集 |
|
||||||
|
| **推理** | 连续批处理、分页 KV Cache、Radix 前缀缓存、流式生成,以及 Torch/CUDA/FlashAttention 后端 |
|
||||||
|
| **服务** | 基于 FastAPI 的 OpenAI 与 Anthropic 聊天补全协议,支持 SSE 流式输出和工具调用 |
|
||||||
|
| **评测** | Perplexity、MMLU、HumanEval、IFEval、IFD 和 ROUGE 评测工具 |
|
||||||
|
| **扩展** | 基于工厂与注册表扩展模型、数据集、训练策略、回调、内核和协议组件 |
|
||||||
|
|
||||||
### 快速上手
|
### 快速上手
|
||||||
|
|
||||||
@@ -258,7 +262,7 @@ SSE 流式格式、错误码和统计端点详见[推理文档](guides/inference
|
|||||||
|
|
||||||
### 许可证
|
### 许可证
|
||||||
|
|
||||||
本项目采用 [GPL-3.0 许可证](../LICENSE)。
|
本项目采用 [Apache-2.0 许可证](../LICENSE)。
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
|||||||
@@ -817,12 +817,31 @@ classDiagram
|
|||||||
+AutoModel model
|
+AutoModel model
|
||||||
+AutoTokenizer tokenizer
|
+AutoTokenizer tokenizer
|
||||||
+PagePool kv_cache
|
+PagePool kv_cache
|
||||||
|
+InferenceWorkspace _workspace
|
||||||
+Optional[str] device
|
+Optional[str] device
|
||||||
+Optional[torch.dtype] dtype
|
+Optional[torch.dtype] dtype
|
||||||
+execute_prefill(tasks, prompt_len, start_pos)
|
+execute_prefill(tasks, prompt_len, start_pos=0)
|
||||||
+execute_decode(tasks, return_logprobs=False) Union[List[int], List[Tuple[int, float]]]
|
+execute_decode(tasks, return_logprobs=False) Union[List[int], List[Tuple[int, float]]]
|
||||||
}
|
}
|
||||||
|
|
||||||
|
class InferenceWorkspace {
|
||||||
|
+int max_batch_size
|
||||||
|
+int max_seq_len
|
||||||
|
+torch.device device
|
||||||
|
+torch.dtype dtype
|
||||||
|
+Tensor arange
|
||||||
|
+Tensor input_mask
|
||||||
|
+Tensor input_ids
|
||||||
|
+Tensor req_pool_indices
|
||||||
|
+Tensor seq_lens
|
||||||
|
+Tensor kv_indptr
|
||||||
|
+Tensor qo_indptr
|
||||||
|
+Tensor inc
|
||||||
|
+Tensor out_cache_loc
|
||||||
|
+fill_input_ids(ids) Tensor
|
||||||
|
+decode_mask(position_ids, total_len) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
class InferenceScheduler {
|
class InferenceScheduler {
|
||||||
+PagePool _cache
|
+PagePool _cache
|
||||||
+Executor _executor
|
+Executor _executor
|
||||||
@@ -851,12 +870,21 @@ classDiagram
|
|||||||
+ref_count(idx) int
|
+ref_count(idx) int
|
||||||
}
|
}
|
||||||
|
|
||||||
class PrefixCache {
|
class RadixNode {
|
||||||
|
+RadixNode parent
|
||||||
|
+Dict children
|
||||||
|
+Optional[int] page_idx
|
||||||
|
+Tuple tokens
|
||||||
|
+int lock_ref
|
||||||
|
}
|
||||||
|
|
||||||
|
class RadixCache {
|
||||||
+int _page_size
|
+int _page_size
|
||||||
+evict(page_idx)
|
+evict(page_idx)
|
||||||
+has_page(idx) bool
|
+has_page(idx) bool
|
||||||
+lookup(token_ids) List[int]
|
+lookup(token_ids) List[int]
|
||||||
+record(page_idx, token_ids, logical_page_idx)
|
+record(page_idx, token_ids, logical_page_idx)
|
||||||
|
+release(pages)
|
||||||
}
|
}
|
||||||
|
|
||||||
class KVStorage {
|
class KVStorage {
|
||||||
@@ -886,6 +914,7 @@ classDiagram
|
|||||||
+Tensor out_cache_loc
|
+Tensor out_cache_loc
|
||||||
+int max_len
|
+int max_len
|
||||||
+Optional[Tensor] kv_indptr
|
+Optional[Tensor] kv_indptr
|
||||||
|
+Optional[Tensor] qo_indptr
|
||||||
}
|
}
|
||||||
|
|
||||||
class PagePool {
|
class PagePool {
|
||||||
@@ -894,13 +923,13 @@ classDiagram
|
|||||||
-KVStorage _storage
|
-KVStorage _storage
|
||||||
-ReqToTokenPool _req_pool
|
-ReqToTokenPool _req_pool
|
||||||
-Allocator _alloc
|
-Allocator _alloc
|
||||||
-PrefixCache _prefix
|
-RadixCache _prefix
|
||||||
+task_alloc(task_id, prompt_ids) bool
|
+task_alloc(task_id, prompt_ids) bool
|
||||||
+task_free(task_id)
|
+task_free(task_id)
|
||||||
+task_extend(task_id, pos) bool
|
+task_extend(task_id, pos) bool
|
||||||
+task_cached(task_id) int
|
+task_cached(task_id) int
|
||||||
+task_record_hashes(task_id, prompt_ids, start_logical_page)
|
+task_record_hashes(task_id, prompt_ids, start_logical_page)
|
||||||
+bind_tasks(task_ids, seq_lens, device, start_pos) KVCache
|
+bind_tasks(task_ids, workspace, device, start_pos) KVCache
|
||||||
}
|
}
|
||||||
|
|
||||||
class Task {
|
class Task {
|
||||||
@@ -1301,10 +1330,12 @@ classDiagram
|
|||||||
PagePool *-- KVStorage
|
PagePool *-- KVStorage
|
||||||
PagePool *-- ReqToTokenPool
|
PagePool *-- ReqToTokenPool
|
||||||
PagePool *-- Allocator
|
PagePool *-- Allocator
|
||||||
PagePool *-- PrefixCache
|
PagePool *-- RadixCache
|
||||||
|
RadixCache *-- RadixNode
|
||||||
InferenceEngine *-- InferenceScheduler
|
InferenceEngine *-- InferenceScheduler
|
||||||
InferenceScheduler *-- PagePool
|
InferenceScheduler *-- PagePool
|
||||||
InferenceScheduler *-- Executor
|
InferenceScheduler *-- Executor
|
||||||
|
Executor *-- InferenceWorkspace
|
||||||
InferenceScheduler *-- TaskManager
|
InferenceScheduler *-- TaskManager
|
||||||
AutoRegressiveLM *-- DecoderBlock
|
AutoRegressiveLM *-- DecoderBlock
|
||||||
AutoRegressiveLM *-- RotaryEmbedding
|
AutoRegressiveLM *-- RotaryEmbedding
|
||||||
@@ -1375,6 +1406,7 @@ classDiagram
|
|||||||
Checkpoint ..> Checkpoint : serializes
|
Checkpoint ..> Checkpoint : serializes
|
||||||
CheckpointCallback ..> Checkpoint : creates
|
CheckpointCallback ..> Checkpoint : creates
|
||||||
PagePool ..> KVCache : binds
|
PagePool ..> KVCache : binds
|
||||||
|
PagePool ..> InferenceWorkspace : fills
|
||||||
InferenceEngine ..> GenerationRequest : uses
|
InferenceEngine ..> GenerationRequest : uses
|
||||||
InferenceEngine ..> GenerateResult : creates
|
InferenceEngine ..> GenerateResult : creates
|
||||||
OpenAIResponseBuilder ..> ChatCompletionRequest : receives
|
OpenAIResponseBuilder ..> ChatCompletionRequest : receives
|
||||||
@@ -1411,8 +1443,8 @@ classDiagram
|
|||||||
| **astrai.model** | ModelFactory, AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, LoRAConfig, LoRALinear, RotaryEmbedding, Embedding | Neural network model |
|
| **astrai.model** | ModelFactory, AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, LoRAConfig, LoRALinear, RotaryEmbedding, Embedding | Neural network model |
|
||||||
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
||||||
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–WSDScheduler, SchedulerFactory, TrainCallback(Protocol)–MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow |
|
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–WSDScheduler, SchedulerFactory, TrainCallback(Protocol)–MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow |
|
||||||
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, PagePool, KVStorage, ReqToTokenPool, KVCache, Allocator, PrefixCache, Task, TaskManager, TaskStatus, StreamDecoder, GenerationRequest, GenerateResult, BaseSamplingStrategy–SamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service |
|
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, InferenceWorkspace, PagePool, KVStorage, ReqToTokenPool, KVCache, Allocator, RadixCache, Task, TaskManager, TaskStatus, StreamDecoder, GenerationRequest, GenerateResult, BaseSamplingStrategy–SamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service |
|
||||||
| **astrai.extension** | AttentionBackend, TorchNativeBackend, CudaBackend, attn_backend, ATTN_BACKEND, attn_decode, attn_prefill, attn_paged_decode, rotary_emb, apply_rotary_emb, rotary_backend, is_available | CUDA attention + rotary kernels, backend abstraction, auto-dispatch |
|
| **astrai.extension** | AttentionBackend, TorchNativeBackend, CudaBackend, attn_backend, ATTN_BACKEND, attn_decode, attn_prefill, attn_paged_decode, attn_paged_prefill, rotary_emb, apply_rotary_emb, rotary_backend, is_available | CUDA attention + rotary kernels, backend abstraction, auto-dispatch |
|
||||||
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, LaunchStrategy, TorchrunStrategy, LocalStrategy, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, GradientState, AccumOptimizer, AccumScheduler | Distributed parallel & gradient accumulation |
|
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, LaunchStrategy, TorchrunStrategy, LocalStrategy, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, GradientState, AccumOptimizer, AccumScheduler | Distributed parallel & gradient accumulation |
|
||||||
| **astrai.factory** | BaseFactory | Component registration |
|
| **astrai.factory** | BaseFactory | Component registration |
|
||||||
| **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
|
| **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
|
||||||
@@ -1424,7 +1456,7 @@ classDiagram
|
|||||||
| **Factory** | `ModelFactory`, `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory`, `MaskBuilderFactory`, `StoreWriterFactory`, `PackingStrategyFactory`, `PositionIdStrategyFactory`, `ToolParserFactory` | Decorator-based component creation |
|
| **Factory** | `ModelFactory`, `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory`, `MaskBuilderFactory`, `StoreWriterFactory`, `PackingStrategyFactory`, `PositionIdStrategyFactory`, `ToolParserFactory` | Decorator-based component creation |
|
||||||
| **Registry** | `BaseFactory` | Component registration |
|
| **Registry** | `BaseFactory` | Component registration |
|
||||||
| **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching |
|
| **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching |
|
||||||
| **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations |
|
| **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `FrequencyPenaltyStrategy`, `SamplingPipeline` | Composable logit transformations |
|
||||||
| **Strategy (API)** | `ResponseBuilder`, `OpenAIResponseBuilder`, `AnthropicResponseBuilder` | HTTP API handler with format hooks |
|
| **Strategy (API)** | `ResponseBuilder`, `OpenAIResponseBuilder`, `AnthropicResponseBuilder` | HTTP API handler with format hooks |
|
||||||
| **Builder** | `TrainContextBuilder` | Chain-building training context |
|
| **Builder** | `TrainContextBuilder` | Chain-building training context |
|
||||||
| **Observer** | `TrainCallback`, callback implementations | Training process monitoring |
|
| **Observer** | `TrainCallback`, callback implementations | Training process monitoring |
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ AstrAI includes optional custom CUDA kernels for attention and rotary embedding.
|
|||||||
| `attn_decode` | `attn_decode.cu` | GQA decode attention (split-KV) |
|
| `attn_decode` | `attn_decode.cu` | GQA decode attention (split-KV) |
|
||||||
| `attn_prefill` | `attn_prefill.cu` | GQA prefill attention (split-Q) |
|
| `attn_prefill` | `attn_prefill.cu` | GQA prefill attention (split-Q) |
|
||||||
| `attn_paged_decode` | `attn_paged_decode.cu` | Paged KV cache decode attention |
|
| `attn_paged_decode` | `attn_paged_decode.cu` | Paged KV cache decode attention |
|
||||||
|
| `attn_paged_prefill` | `attn_paged_prefill.cu` | Paged KV cache prefill attention (ragged batch) |
|
||||||
| `rotary_emb` | `rotary_emb.cu` | Fused rotary embedding (cos/sin lookup + rotation) |
|
| `rotary_emb` | `rotary_emb.cu` | Fused rotary embedding (cos/sin lookup + rotation) |
|
||||||
|
|
||||||
Additionally, optimized `.cuh` variants with tensor-core MMA (Matrix Multiply-Accumulate) exist:
|
Additionally, optimized `.cuh` variants with tensor-core MMA (Matrix Multiply-Accumulate) exist:
|
||||||
@@ -17,7 +18,10 @@ Additionally, optimized `.cuh` variants with tensor-core MMA (Matrix Multiply-Ac
|
|||||||
|---------|------|--------------|
|
|---------|------|--------------|
|
||||||
| Split-KV MMA decode | `attn_decode_split_kv_mma.cuh` | Split KV across warps + MMA (sm_80+) |
|
| Split-KV MMA decode | `attn_decode_split_kv_mma.cuh` | Split KV across warps + MMA (sm_80+) |
|
||||||
| Split-Q MMA prefill | `attn_prefill_split_q_mma.cuh` | Split Q across warps + MMA (sm_80+) |
|
| Split-Q MMA prefill | `attn_prefill_split_q_mma.cuh` | Split Q across warps + MMA (sm_80+) |
|
||||||
| Paged split-KV MMA decode | `attn_paged_decode_split_kv_mma.cuh` | Paged cache + split-KV + MMA |
|
|
||||||
|
> The paged and non-paged paths are ONE kernel templated on a `KVSource`
|
||||||
|
> policy (`ContigKV` / `PagedKV` in `attn_kv_source.cuh`); there are no
|
||||||
|
> separate `attn_paged_*.cuh` files anymore.
|
||||||
|
|
||||||
### Rotary Embedding Kernel
|
### Rotary Embedding Kernel
|
||||||
|
|
||||||
@@ -50,23 +54,32 @@ CSRC_KERNELS=true pip install -e . --no-build-isolation
|
|||||||
# Rebuild after editing .cu/.cuh files
|
# Rebuild after editing .cu/.cuh files
|
||||||
CSRC_KERNELS=true python setup.py build_ext --inplace
|
CSRC_KERNELS=true python setup.py build_ext --inplace
|
||||||
# Output: astrai/extension/lib/*.so
|
# Output: astrai/extension/lib/*.so
|
||||||
|
|
||||||
|
# Or invoke CMake directly
|
||||||
|
cmake -S csrc -B build/cmake \
|
||||||
|
-DTORCH_HOME=<site-packages>/torch \
|
||||||
|
-DPYTHON_INCLUDE_DIR=<python include> \
|
||||||
|
-DPY_SOABI=cpython-312-x86_64-linux-gnu
|
||||||
|
cmake --build build/cmake -j 16
|
||||||
```
|
```
|
||||||
|
|
||||||
### Architecture flags
|
### Architecture flags
|
||||||
|
|
||||||
`csrc/build.py` auto-detects the GPU compute capability and generates the appropriate `nvcc` gencode flag:
|
`setup.py` passes the GPU compute capability to CMake via `ASTRAI_CUDA_ARCH` (default `89`, i.e. sm_89 / L20):
|
||||||
|
|
||||||
- **sm_80+** (Ampere and later): enables tensor-core MMA path (`mma.sync.m16n8k16.bf16`)
|
- **sm_80+** (Ampere and later): enables tensor-core MMA path (`mma.sync.m16n8k16.bf16`)
|
||||||
- **Below sm_80**: adds `-DASTRAI_NO_MMA` to disable the MMA path at compile time
|
- **Below sm_80**: adds `-DASTRAI_NO_MMA` to disable the MMA path at compile time
|
||||||
|
|
||||||
### Build configuration
|
### Build configuration
|
||||||
|
|
||||||
|
`csrc/CMakeLists.txt` defines the CUDA extension build:
|
||||||
|
|
||||||
```
|
```
|
||||||
NVCC_FLAGS = -O3 --expt-relaxed-constexpr --use_fast_math
|
NVCC_FLAGS = -O3 --expt-relaxed-constexpr --use_fast_math
|
||||||
--ptxas-options=-O3,-v --extra-device-vectorization --threads=8
|
--ptxas-options=-O3,-v --extra-device-vectorization --threads=16
|
||||||
```
|
```
|
||||||
|
|
||||||
The `REGISTRY` in `csrc/build.py` lists all registered kernels (currently 4). Each entry maps a kernel name to its source files and build flags.
|
Each kernel in `astrai/extension/lib` is compiled as an independent pybind11 module (one `.so` per kernel, named `<kernel>.cpython-*-x86_64-linux-gnu.so`). CMake builds all five kernel targets in parallel via `cmake --build -j N`.
|
||||||
|
|
||||||
## Attention Backend
|
## Attention Backend
|
||||||
|
|
||||||
@@ -74,7 +87,7 @@ The `REGISTRY` in `csrc/build.py` lists all registered kernels (currently 4). Ea
|
|||||||
|
|
||||||
- **`AttentionBackend`** (ABC): `fwd_decode` / `fwd_prefill` abstract methods, `forward` dispatches by q_len
|
- **`AttentionBackend`** (ABC): `fwd_decode` / `fwd_prefill` abstract methods, `forward` dispatches by q_len
|
||||||
- **`TorchNativeBackend`**: SDPA with indirect KV cache gather (default)
|
- **`TorchNativeBackend`**: SDPA with indirect KV cache gather (default)
|
||||||
- **`CudaBackend`**: CUDA kernel dispatch — decode via `attn_paged_decode` (page_size=1), prefill via `attn_prefill`
|
- **`CudaBackend`**: CUDA kernel dispatch — decode via `attn_paged_decode` (page_size=1), prefill via `attn_paged_prefill` (ragged batch, `qo_indptr` + `kv_indptr`)
|
||||||
|
|
||||||
Select a backend via context manager (mirrors `torch.nn.attention.sdpa_kernel`):
|
Select a backend via context manager (mirrors `torch.nn.attention.sdpa_kernel`):
|
||||||
|
|
||||||
@@ -145,20 +158,20 @@ nvcc -I csrc -arch=sm_89 -O3 --use_fast_math \
|
|||||||
|
|
||||||
```
|
```
|
||||||
csrc/
|
csrc/
|
||||||
├── build.py # Build system: REGISTRY, _arch_flags, nvcc flags
|
├── CMakeLists.txt # CMake build: 5 kernel targets, torch/pybind11 linking
|
||||||
├── kernels/
|
├── kernels/
|
||||||
│ ├── attn_common.h # Shared attention params (AttentionParams, PagedAttentionParams)
|
│ ├── attn_common.h # Unified attention params (contig + paged modes)
|
||||||
│ ├── attn_decode.cu # Basic decode kernel (registered)
|
│ ├── attn_decode.cu # Basic decode kernel (registered)
|
||||||
│ ├── attn_prefill.cu # Basic prefill kernel (registered)
|
│ ├── attn_prefill.cu # Basic prefill kernel (registered)
|
||||||
│ ├── attn_paged_decode.cu # Paged decode kernel (registered)
|
│ ├── attn_paged_decode.cu # Paged decode kernel (registered)
|
||||||
|
│ ├── attn_paged_prefill.cu # Paged prefill kernel (registered)
|
||||||
│ ├── rotary_emb.cu # Fused rotary embedding kernel (registered)
|
│ ├── rotary_emb.cu # Fused rotary embedding kernel (registered)
|
||||||
│ ├── attn_decode_split_kv.cuh # Split-KV variant
|
│ ├── attn_decode_split_kv.cuh # Split-KV variant (contig + paged via KVSource)
|
||||||
│ ├── attn_decode_split_kv_mma.cuh # Split-KV + MMA variant
|
│ ├── attn_decode_split_kv_mma.cuh # Split-KV + MMA variant (contig + paged)
|
||||||
│ ├── attn_prefill_split_q.cuh # Split-Q variant
|
│ ├── attn_prefill_split_q.cuh # Split-Q variant (contig + paged via KVSource)
|
||||||
│ ├── attn_prefill_split_q_mma.cuh # Split-Q + MMA variant
|
│ ├── attn_prefill_split_q_mma.cuh # Split-Q + MMA variant (contig + paged)
|
||||||
│ ├── attn_paged_decode_split_kv.cuh # Paged + split-KV variant
|
│ ├── attn_kv_source.cuh # KVSource policies (ContigKV / PagedKV)
|
||||||
│ ├── attn_paged_decode_split_kv_mma.cuh # Paged + split-KV + MMA variant
|
│ ├── attn_dispatchers.cuh # Kernel dispatch macros + KV-templated launchers
|
||||||
│ ├── attn_dispatchers.cuh # Kernel dispatch macros
|
|
||||||
│ ├── attn_entry_utils.cuh # Entry point helpers
|
│ ├── attn_entry_utils.cuh # Entry point helpers
|
||||||
│ ├── attn_mma_utils.cuh # MMA utilities
|
│ ├── attn_mma_utils.cuh # MMA utilities
|
||||||
│ └── attn_warp_utils.cuh # Warp-level utilities
|
│ └── attn_warp_utils.cuh # Warp-level utilities
|
||||||
|
|||||||
@@ -41,12 +41,14 @@ RoPE embeds position into Q/K vectors via complex rotation:
|
|||||||
|
|
||||||
$$ q_i = R_i W_q x_i, \quad k_j = R_j W_k x_j, \quad q_i^T k_j = x_i^T W_q^T R_{i-j} W_k x_j $$
|
$$ q_i = R_i W_q x_i, \quad k_j = R_j W_k x_j, \quad q_i^T k_j = x_i^T W_q^T R_{i-j} W_k x_j $$
|
||||||
|
|
||||||
`RotaryEmbedding` pre-computes a complex `freqs_cis` buffer. `forward()` returns
|
`RotaryEmbedding` pre-computes a cos/sin table `freqs_cis` of shape
|
||||||
a tensor indexed by `position_ids`. `apply_rotary_emb` applies the rotation:
|
`[max_len, dim/2, 2]` (f32 — `[cos, sin]` pairs). `forward()` returns
|
||||||
during training it uses torch complex multiply (autograd-compatible); during
|
a `[batch, seq_len, dim/2, 2]` slice indexed by `position_ids`.
|
||||||
inference it auto-dispatches to a fused CUDA kernel when available. The key
|
`apply_rotary_emb` applies the rotation: during training it uses torch
|
||||||
property is that the dot product $q_i^T k_j$ depends only on the relative
|
complex multiply (autograd-compatible); during inference it auto-dispatches
|
||||||
position $i - j$, not the absolute positions.
|
to a fused CUDA kernel when available. The key property is that the dot
|
||||||
|
product $q_i^T k_j$ depends only on the relative position $i - j$, not the
|
||||||
|
absolute positions.
|
||||||
|
|
||||||
**Critical for inference**: RoPE is applied **before** KV cache write, not after. If applied after caching, position encoding drift occurs because cached K/V would have stale rotation factors.
|
**Critical for inference**: RoPE is applied **before** KV cache write, not after. If applied after caching, position encoding drift occurs because cached K/V would have stale rotation factors.
|
||||||
|
|
||||||
@@ -166,16 +168,18 @@ Three-layer separation (SGLang-inspired):
|
|||||||
|
|
||||||
- **KVStorage**: Flat token-level buffers `[n_layers, size, n_kv_heads, head_dim]`.
|
- **KVStorage**: Flat token-level buffers `[n_layers, size, n_kv_heads, head_dim]`.
|
||||||
- **ReqToTokenPool**: Index table `[req_idx, pos] → physical token slot`, shared across all layers.
|
- **ReqToTokenPool**: Index table `[req_idx, pos] → physical token slot`, shared across all layers.
|
||||||
- **Allocator + PrefixCache**: Paged-mode slot allocation with ref-counting, LRU eviction, and hash-based prefix sharing.
|
- **Allocator + RadixCache**: Paged-mode allocation with ref-counting, LRU eviction, and exact page-aligned prefix sharing when `page_size > 1`.
|
||||||
|
|
||||||
`PagePool` orchestrates all three. In contiguous mode (default), `req_to_token` is a trivial linear mapping. In paged mode, slots are allocated on demand with prefix caching support. `bind_tasks()` returns a `KVCache` dataclass with `kv_indptr`, a prefix-sum index over sequence lengths computed once per step and shared across layers. Attention layers access buffers directly — no methods, no abstraction.
|
`PagePool` orchestrates all three. In contiguous mode (default), `req_to_token` is a trivial linear mapping. In paged mode, slots are allocated on demand. `RadixCache` walks exact token-page edges from the root, preserving parent-prefix context instead of treating a page hash as a globally unique key. Only complete pages whose KV entries have been materialized are shared; partial pages remain request-private and are released at completion. The final sampled token is excluded because it has not yet been decoded into KV.
|
||||||
|
|
||||||
|
`bind_tasks()` returns a `KVCache` dataclass with `kv_indptr`, a prefix-sum index over sequence lengths computed once per step and shared across layers. Attention layers access buffers directly — no methods, no abstraction.
|
||||||
|
|
||||||
### Attention Backend
|
### Attention Backend
|
||||||
|
|
||||||
Attention computation is decoupled from the model via `AttentionBackend` ABC (`astrai/extension/attention_backend.py`):
|
Attention computation is decoupled from the model via `AttentionBackend` ABC (`astrai/extension/attention_backend.py`):
|
||||||
|
|
||||||
- **`TorchNativeBackend`** (default): writes K/V to cache, gathers via `req_to_token` indirect indexing, calls `F.scaled_dot_product_attention`.
|
- **`TorchNativeBackend`** (default): writes K/V to cache, gathers via `req_to_token` indirect indexing, calls `F.scaled_dot_product_attention`.
|
||||||
- **`CudaBackend`**: decode path uses `attn_paged_decode` with `page_size=1` (the `req_to_token` table serves as the page table, each token slot is a single-token "page"); prefill path gathers K/V then calls `attn_prefill`. Falls back to `TorchNativeBackend` when kernel unavailable.
|
- **`CudaBackend`**: decode path uses `attn_paged_decode` with `page_size=1` (the `req_to_token` table serves as the page table, each token slot is a single-token "page"); prefill path uses the ragged-batch `attn_paged_prefill` (addresses each request via `qo_indptr` + `kv_indptr` directly against the flat pool). Falls back to `TorchNativeBackend` when kernel unavailable.
|
||||||
|
|
||||||
Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_backend.py`, which auto-dispatches to the fused CUDA kernel (`rotary_emb.cu`) during inference or torch complex multiply during training (for autograd compatibility). Both attention backends share the same rotary dispatch.
|
Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_backend.py`, which auto-dispatches to the fused CUDA kernel (`rotary_emb.cu`) during inference or torch complex multiply during training (for autograd compatibility). Both attention backends share the same rotary dispatch.
|
||||||
|
|
||||||
|
|||||||
+16
-10
@@ -31,13 +31,17 @@ PagePool (top-level manager, orchestrates all layers)
|
|||||||
├── KVStorage k_buffer / v_buffer [n_layers, size, n_kv_heads, head_dim]
|
├── KVStorage k_buffer / v_buffer [n_layers, size, n_kv_heads, head_dim]
|
||||||
├── ReqToTokenPool req_to_token [num_reqs, max_ctx_len] → physical token slot
|
├── ReqToTokenPool req_to_token [num_reqs, max_ctx_len] → physical token slot
|
||||||
├── Allocator bitmask-based page allocator + ref-count + LRU (paged mode only)
|
├── Allocator bitmask-based page allocator + ref-count + LRU (paged mode only)
|
||||||
└── PrefixCache hash-based prefix matching (paged mode only)
|
└── RadixCache exact, page-aligned prefix matching (paged mode, page_size > 1)
|
||||||
```
|
```
|
||||||
|
|
||||||
`PagePool` supports two modes:
|
`PagePool` supports two modes:
|
||||||
|
|
||||||
- **Contiguous (default)**: pre-allocates `max_batch_size * max_seq_len` token slots. `req_to_token` is a trivial linear mapping (`slot = req_idx * max_seq_len + pos`). No dynamic allocation.
|
- **Contiguous (default)**: pre-allocates `max_batch_size * max_seq_len` token slots. `req_to_token` is a trivial linear mapping (`slot = req_idx * max_seq_len + pos`). No dynamic allocation.
|
||||||
- **Paged** (`page_size=1` or `>1` with `n_tokens` set): shared token pool with on-demand allocation. Allocator + PrefixCache enable prefix sharing and LRU eviction.
|
- **Paged** (`page_size=1` or `>1` with `n_tokens` set): shared token pool with on-demand allocation. `Allocator` provides ref-counted allocation and LRU eviction. When `page_size > 1`, `RadixCache` also enables prefix sharing.
|
||||||
|
|
||||||
|
`RadixCache` indexes complete token pages as parent-linked radix edges. Lookup walks from the root and compares each page's exact token tuple, so an identical page can only be reused under the same parent prefix. Hash values are retained for introspection, but never determine a match.
|
||||||
|
|
||||||
|
Only fully materialized KV pages enter the radix. A partial final page remains private to its request and is released when the request ends. On completion, the scheduler records the prompt plus generated tokens already decoded into KV; it excludes the final sampled token because that token has not yet passed through the model. A later request resumes prefill immediately after the longest complete-page hit.
|
||||||
|
|
||||||
`bind_tasks()` returns a `KVCache` dataclass — pure data, no methods:
|
`bind_tasks()` returns a `KVCache` dataclass — pure data, no methods:
|
||||||
|
|
||||||
@@ -49,7 +53,8 @@ KVCache
|
|||||||
├── seq_lens [batch_size]
|
├── seq_lens [batch_size]
|
||||||
├── out_cache_loc [batch, seq_len] — write indices for this forward
|
├── out_cache_loc [batch, seq_len] — write indices for this forward
|
||||||
├── max_len int — max(seq_lens), avoids GPU sync in decode
|
├── max_len int — max(seq_lens), avoids GPU sync in decode
|
||||||
└── kv_indptr [batch + 1] int32 — prefix sum of seq_lens, precomputed once per step
|
├── kv_indptr [batch + 1] int32 — prefix sum of seq_lens, precomputed once per step
|
||||||
|
└── qo_indptr [batch + 1] int32 — prefix sum of per-request q_lens (prefill), precomputed once per step
|
||||||
```
|
```
|
||||||
|
|
||||||
Attention layers do raw buffer indexing: `k_buffer[layer_id, out_cache_loc] = k` to write, `k_buffer[layer_id, indices]` to gather.
|
Attention layers do raw buffer indexing: `k_buffer[layer_id, out_cache_loc] = k` to write, `k_buffer[layer_id, indices]` to gather.
|
||||||
@@ -61,7 +66,7 @@ Attention computation (cache I/O + SDPA/kernel dispatch) is decoupled from the m
|
|||||||
```
|
```
|
||||||
AttentionBackend (ABC)
|
AttentionBackend (ABC)
|
||||||
├── TorchNativeBackend SDPA + indirect KV cache gather (default)
|
├── TorchNativeBackend SDPA + indirect KV cache gather (default)
|
||||||
└── CudaBackend CUDA kernel dispatch (attn_paged_decode, attn_prefill)
|
└── CudaBackend CUDA kernel dispatch (attn_paged_decode, attn_paged_prefill)
|
||||||
```
|
```
|
||||||
|
|
||||||
Select via context manager (mirrors `torch.nn.attention.sdpa_kernel`):
|
Select via context manager (mirrors `torch.nn.attention.sdpa_kernel`):
|
||||||
@@ -75,7 +80,7 @@ with attn_backend(ATTN_BACKEND.CUDA):
|
|||||||
|
|
||||||
`CudaBackend` decode path: writes K/V to cache, then calls `attn_paged_decode` with `page_size=1` — the `req_to_token` table serves directly as the page table, each token slot is a single-token "page". No explicit K/V gather needed.
|
`CudaBackend` decode path: writes K/V to cache, then calls `attn_paged_decode` with `page_size=1` — the `req_to_token` table serves directly as the page table, each token slot is a single-token "page". No explicit K/V gather needed.
|
||||||
|
|
||||||
`CudaBackend` prefill path: writes K/V, gathers full-sequence K/V via indirect indexing (same as `TorchNativeBackend`), then calls `attn_prefill`.
|
`CudaBackend` prefill path: writes K/V, then calls `attn_paged_prefill` — a ragged-batch (paged) prefill kernel that reads K/V directly from the flat pool via `req_to_token`, addressing each request's `q_len`/`kv_len` through `qo_indptr` and `kv_indptr`. No explicit K/V gather needed.
|
||||||
|
|
||||||
Fallback: `CudaBackend` delegates to `TorchNativeBackend` when a CUDA kernel is not available.
|
Fallback: `CudaBackend` delegates to `TorchNativeBackend` when a CUDA kernel is not available.
|
||||||
|
|
||||||
@@ -83,19 +88,20 @@ Fallback: `CudaBackend` delegates to `TorchNativeBackend` when a CUDA kernel is
|
|||||||
|
|
||||||
Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_backend.py`, which auto-dispatches:
|
Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_backend.py`, which auto-dispatches:
|
||||||
|
|
||||||
- **CUDA kernel** (`rotary_emb.cu`): fused cos/sin lookup + rotation in a single kernel, used when the kernel is available, input is on CUDA, and `torch.is_grad_enabled()` is `False` (inference mode)
|
- **CUDA kernel** (`rotary_emb.cu`): fused cos/sin lookup + rotation in a single kernel, used when the kernel is available, the input is bf16 on CUDA, and `torch.is_grad_enabled()` is `False` (inference mode)
|
||||||
- **Torch fallback**: complex multiply path (`torch.view_as_complex` → `torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd backward) or when the CUDA kernel is not available
|
- **Torch fallback**: complex multiply path (`torch.view_as_complex` → `torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd backward) or when the CUDA kernel is not available
|
||||||
|
|
||||||
`RotaryEmbedding` stores a complex `freqs_cis` buffer and returns a tensor
|
`RotaryEmbedding` stores a cos/sin table `freqs_cis` of shape
|
||||||
from `forward()`. Both attention backends share the same rotary dispatch — it
|
`[max_len, dim/2, 2]` (f32 — `[cos, sin]` pairs) and `forward()` returns
|
||||||
is backend-agnostic.
|
a `[batch, seq_len, dim/2, 2]` slice indexed by `position_ids`. Both
|
||||||
|
attention backends share the same rotary dispatch — it is backend-agnostic.
|
||||||
|
|
||||||
## Continuous Batching
|
## Continuous Batching
|
||||||
|
|
||||||
`InferenceScheduler` runs a daemon thread with a 4-phase loop:
|
`InferenceScheduler` runs a daemon thread with a 4-phase loop:
|
||||||
|
|
||||||
```
|
```
|
||||||
1. Cleanup → Remove finished tasks, free KV cache slots/pages
|
1. Cleanup → Record complete materialized pages, then release task-owned KV resources
|
||||||
2. Refill → Pop from waiting_queue, task_alloc resources, activate
|
2. Refill → Pop from waiting_queue, task_alloc resources, activate
|
||||||
3. Prefill → Group by (prompt_len, start_pos), run full forward
|
3. Prefill → Group by (prompt_len, start_pos), run full forward
|
||||||
4. Decode → Run single-token forward for each same-position group
|
4. Decode → Run single-token forward for each same-position group
|
||||||
|
|||||||
@@ -41,10 +41,12 @@ RoPE embeds position into Q/K vectors via complex rotation:
|
|||||||
|
|
||||||
$$ q_i = R_i W_q x_i, \quad k_j = R_j W_k x_j, \quad q_i^T k_j = x_i^T W_q^T R_{i-j} W_k x_j $$
|
$$ q_i = R_i W_q x_i, \quad k_j = R_j W_k x_j, \quad q_i^T k_j = x_i^T W_q^T R_{i-j} W_k x_j $$
|
||||||
|
|
||||||
`RotaryEmbedding` pre-computes a complex `freqs_cis` buffer. `forward()` returns
|
`RotaryEmbedding` pre-computes a cos/sin table `freqs_cis` of shape
|
||||||
a tensor indexed by `position_ids`. `apply_rotary_emb` applies the rotation:
|
`[max_len, dim/2, 2]` (f32 — `[cos, sin]` pairs). `forward()` returns
|
||||||
during training it uses torch complex multiply (autograd-compatible); during
|
a `[batch, seq_len, dim/2, 2]` slice indexed by `position_ids`.
|
||||||
inference it auto-dispatches to a fused CUDA kernel when available.
|
`apply_rotary_emb` applies the rotation: during training it uses torch
|
||||||
|
complex multiply (autograd-compatible); during inference it auto-dispatches
|
||||||
|
to a fused CUDA kernel when available.
|
||||||
|
|
||||||
## Training Loop
|
## Training Loop
|
||||||
|
|
||||||
|
|||||||
+3
-2
@@ -22,16 +22,17 @@ dependencies = [
|
|||||||
"pyyaml>=6.0",
|
"pyyaml>=6.0",
|
||||||
]
|
]
|
||||||
keywords = ["nlp", "datasets", "language-models", "machine-learning"]
|
keywords = ["nlp", "datasets", "language-models", "machine-learning"]
|
||||||
license = { text = "GPL-3.0" }
|
license = { text = "Apache-2.0" }
|
||||||
classifiers = [
|
classifiers = [
|
||||||
"Programming Language :: Python :: 3",
|
"Programming Language :: Python :: 3",
|
||||||
"License :: OSI Approved :: GPL-3.0",
|
"License :: OSI Approved :: Apache Software License",
|
||||||
"Operating System :: OS Independent",
|
"Operating System :: OS Independent",
|
||||||
]
|
]
|
||||||
urls = { Homepage = "https://github.com/ViperEkura/AstrAI" }
|
urls = { Homepage = "https://github.com/ViperEkura/AstrAI" }
|
||||||
|
|
||||||
[project.optional-dependencies]
|
[project.optional-dependencies]
|
||||||
dev = ["pytest==9.0.2", "ruff", "httpx2"]
|
dev = ["pytest==9.0.2", "ruff", "httpx2"]
|
||||||
|
flash = ["flash-attn>=2.6"]
|
||||||
|
|
||||||
[tool.setuptools.packages.find]
|
[tool.setuptools.packages.find]
|
||||||
where = ["."]
|
where = ["."]
|
||||||
|
|||||||
@@ -1,23 +1,18 @@
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Optional
|
from typing import Optional, Union
|
||||||
|
|
||||||
import click
|
import click
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai import setup_logging
|
from astrai import setup_logging
|
||||||
from astrai.config import AutoRegressiveLMConfig
|
from astrai.config import AutoRegressiveLMConfig
|
||||||
from astrai.extension import ATTN_BACKEND, attn_backend
|
from astrai.extension import ATTN_BACKEND, AttentionBackendFactory, attn_backend
|
||||||
from astrai.inference.core.cache import PagePool
|
from astrai.inference.core.cache import PagePool
|
||||||
from astrai.model import AutoModel
|
from astrai.model import AutoModel
|
||||||
|
|
||||||
_DTYPES = ["bfloat16", "float16", "float32"]
|
_DTYPES = ["bfloat16", "float16", "float32"]
|
||||||
_CACHES = ["contiguous", "paged"]
|
_CACHES = ["contiguous", "paged"]
|
||||||
_BACKENDS = ["cuda", "torch_native"]
|
_BACKENDS = AttentionBackendFactory.list_registered()
|
||||||
|
|
||||||
_BACKEND_MAP = {
|
|
||||||
"cuda": ATTN_BACKEND.CUDA,
|
|
||||||
"torch_native": ATTN_BACKEND.TORCH_NATIVE,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
class BenchmarkResult:
|
class BenchmarkResult:
|
||||||
@@ -46,7 +41,7 @@ class GenerationBenchmark:
|
|||||||
device: str = "cuda",
|
device: str = "cuda",
|
||||||
dtype: torch.dtype = torch.bfloat16,
|
dtype: torch.dtype = torch.bfloat16,
|
||||||
cache_type: str = "contiguous",
|
cache_type: str = "contiguous",
|
||||||
backend: ATTN_BACKEND = ATTN_BACKEND.CUDA,
|
backend: Union[str, ATTN_BACKEND] = ATTN_BACKEND.CUDA,
|
||||||
):
|
):
|
||||||
self.device = device
|
self.device = device
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
@@ -297,7 +292,7 @@ def benchmark_command(
|
|||||||
device=device,
|
device=device,
|
||||||
dtype=dtype_map[dtype],
|
dtype=dtype_map[dtype],
|
||||||
cache_type=cache,
|
cache_type=cache,
|
||||||
backend=_BACKEND_MAP[name],
|
backend=name,
|
||||||
)
|
)
|
||||||
|
|
||||||
click.secho(
|
click.secho(
|
||||||
|
|||||||
@@ -289,6 +289,13 @@ _START_METHODS = ["spawn", "fork", "forkserver"]
|
|||||||
group="Data Loading",
|
group="Data Loading",
|
||||||
help="Label smoothing.",
|
help="Label smoothing.",
|
||||||
)
|
)
|
||||||
|
@opt(
|
||||||
|
"--moe_aux_loss_coef",
|
||||||
|
type=float,
|
||||||
|
default=0.01,
|
||||||
|
group="Algorithm",
|
||||||
|
help="MoE load balancing auxiliary loss coefficient (0=disable).",
|
||||||
|
)
|
||||||
@opt(
|
@opt(
|
||||||
"--rollout_interval",
|
"--rollout_interval",
|
||||||
type=int,
|
type=int,
|
||||||
@@ -813,6 +820,7 @@ def train(
|
|||||||
rollout_top_p=rollout_top_p,
|
rollout_top_p=rollout_top_p,
|
||||||
rollout_max_tokens=rollout_max_tokens,
|
rollout_max_tokens=rollout_max_tokens,
|
||||||
reward_model_fn=reward_model_fn,
|
reward_model_fn=reward_model_fn,
|
||||||
|
moe_aux_loss_coef=kwargs.pop("moe_aux_loss_coef", 0.01),
|
||||||
)
|
)
|
||||||
|
|
||||||
trainer = Trainer(train_config)
|
trainer = Trainer(train_config)
|
||||||
|
|||||||
@@ -19,8 +19,6 @@ def _should_build():
|
|||||||
if force == "false":
|
if force == "false":
|
||||||
return False
|
return False
|
||||||
try:
|
try:
|
||||||
import shutil
|
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
return shutil.which("nvcc") is not None and torch.cuda.is_available()
|
return shutil.which("nvcc") is not None and torch.cuda.is_available()
|
||||||
@@ -28,125 +26,133 @@ def _should_build():
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
ext_modules = []
|
def _torch_prefix():
|
||||||
cmdclass = {}
|
"""Return the torch install dir (site-packages/torch) used for headers/libs."""
|
||||||
|
try:
|
||||||
|
import torch
|
||||||
|
|
||||||
if _should_build():
|
return str(Path(torch.__file__).parent.resolve())
|
||||||
import torch
|
except Exception:
|
||||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
return os.environ.get("TORCH_HOME", "")
|
||||||
|
|
||||||
from csrc.build import REGISTRY, cuda_toolkit_version
|
|
||||||
|
|
||||||
# Preflight: warn if nvcc major version != torch's bundled CUDA major version.
|
def _python_include():
|
||||||
# A mismatch (e.g. nvcc 13.0 + cu128 torch) causes cryptic ABI/header errors.
|
import sysconfig
|
||||||
nvcc_ver = cuda_toolkit_version()
|
|
||||||
torch_cuda = torch.version.cuda
|
return sysconfig.get_path("include")
|
||||||
if nvcc_ver is not None and torch_cuda is not None:
|
|
||||||
torch_major = int(torch_cuda.split(".")[0])
|
|
||||||
if nvcc_ver[0] != torch_major:
|
def _python_soabi():
|
||||||
|
import sysconfig
|
||||||
|
|
||||||
|
ext = sysconfig.get_config_var("EXT_SUFFIX").lstrip(".")
|
||||||
|
return ext[: -len(".so")]
|
||||||
|
|
||||||
|
|
||||||
|
class _CMakeBuildExt(_build_ext):
|
||||||
|
def run(self):
|
||||||
|
src = Path(__file__).parent
|
||||||
|
build_dir = src / "build" / "cmake"
|
||||||
|
torch_home = _torch_prefix()
|
||||||
|
if not torch_home:
|
||||||
|
raise RuntimeError(
|
||||||
|
"torch not found; cannot build kernels. "
|
||||||
|
"Activate the environment or set TORCH_HOME."
|
||||||
|
)
|
||||||
|
|
||||||
|
nvcc_ver = _cuda_toolkit_version()
|
||||||
|
torch_cuda = _torch_cuda_version()
|
||||||
|
if (
|
||||||
|
nvcc_ver is not None
|
||||||
|
and torch_cuda is not None
|
||||||
|
and nvcc_ver[0] != int(torch_cuda.split(".")[0])
|
||||||
|
):
|
||||||
warnings.warn(
|
warnings.warn(
|
||||||
f"CUDA version mismatch: nvcc is {nvcc_ver[0]}.{nvcc_ver[1]} "
|
f"CUDA version mismatch: nvcc is {nvcc_ver[0]}.{nvcc_ver[1]} "
|
||||||
f"but torch was built with CUDA {torch_cuda}. "
|
f"but torch was built with CUDA {torch_cuda}. "
|
||||||
f"This may cause compilation errors. "
|
f"Install a matching torch wheel.",
|
||||||
f"Install a matching torch wheel: "
|
|
||||||
f"pip install torch --index-url "
|
|
||||||
f"https://download.pytorch.org/whl/cu{nvcc_ver[0]}{nvcc_ver[1]}",
|
|
||||||
stacklevel=2,
|
stacklevel=2,
|
||||||
)
|
)
|
||||||
|
|
||||||
_torch_lib = torch.utils.cpp_extension.library_paths()[0]
|
cmake = shutil.which("cmake")
|
||||||
|
if cmake is None:
|
||||||
|
raise RuntimeError("cmake not found on PATH; install it to build kernels")
|
||||||
|
|
||||||
for name, info in REGISTRY.items():
|
parallel = os.environ.get("BUILD_PARALLEL", "16")
|
||||||
ext_modules.append(
|
cfg = [
|
||||||
CUDAExtension(
|
cmake,
|
||||||
f"astrai.extension.lib.{name}",
|
"-S",
|
||||||
info["sources"],
|
str(src / "csrc"),
|
||||||
extra_compile_args={
|
"-B",
|
||||||
"cxx": info["cxx_flags"],
|
str(build_dir),
|
||||||
"nvcc": info["nvcc_flags"],
|
f"-DTORCH_HOME={torch_home}",
|
||||||
},
|
f"-DPYTHON_INCLUDE_DIR={_python_include()}",
|
||||||
extra_link_args=[f"-Wl,-rpath,{_torch_lib}"],
|
f"-DPY_SOABI={_python_soabi()}",
|
||||||
)
|
]
|
||||||
|
arch = os.environ.get("ASTRAI_CUDA_ARCH")
|
||||||
|
if not arch:
|
||||||
|
arch = _detect_cuda_arch()
|
||||||
|
if arch:
|
||||||
|
cfg.append(f"-DASTRAI_CUDA_ARCH={arch}")
|
||||||
|
subprocess.run(cfg, check=True)
|
||||||
|
subprocess.run([cmake, "--build", str(build_dir), "-j", parallel], check=True)
|
||||||
|
|
||||||
|
|
||||||
|
def _cuda_toolkit_version():
|
||||||
|
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()
|
||||||
|
return tuple(int(x) for x in ver.split("."))
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return None
|
||||||
|
|
||||||
# Parallel build — each extension is an independent ninja project, so we
|
|
||||||
# can compile them concurrently. BuildExtension compiles them serially by
|
|
||||||
# default; this subclass dispatches each extension to a subprocess.
|
|
||||||
# Set BUILD_PARALLEL=N to override (default: min(n_exts, 4)).
|
|
||||||
_single_ext = os.environ.get("ASTRAI_BUILD_SINGLE_EXT", "")
|
|
||||||
|
|
||||||
class ParallelBuildExtension(BuildExtension):
|
def _detect_cuda_arch():
|
||||||
def build_extensions(self):
|
"""Detect real GPU compute capability via torch (nvidia-smi may be spoofed).
|
||||||
if _single_ext:
|
|
||||||
self.extensions = [e for e in self.extensions if e.name == _single_ext]
|
|
||||||
if not self.extensions:
|
|
||||||
return
|
|
||||||
super().build_extensions()
|
|
||||||
return
|
|
||||||
|
|
||||||
n = len(self.extensions)
|
Returns something like ``"89"`` or ``"103"``, or ``None`` if unavailable.
|
||||||
max_workers = int(os.environ.get("BUILD_PARALLEL", 8))
|
"""
|
||||||
if max_workers <= 1 or n <= 1:
|
try:
|
||||||
super().build_extensions()
|
import torch
|
||||||
return
|
|
||||||
|
|
||||||
# Each subprocess gets its own build-temp / build-lib so the
|
if torch.cuda.is_available():
|
||||||
# ninja files (build.ninja, .ninja_log) never race. The built
|
major, minor = torch.cuda.get_device_capability()
|
||||||
# .so files are then collected into the parent's build_lib so the
|
return f"{major}{minor}"
|
||||||
# normal setuptools copy steps (inplace / editable wheel) work.
|
except Exception:
|
||||||
names = [e.name for e in self.extensions]
|
pass
|
||||||
env = {**os.environ, "BUILD_PARALLEL": "1"}
|
return None
|
||||||
base = os.path.join("build", "parallel")
|
|
||||||
os.makedirs(base, exist_ok=True)
|
|
||||||
procs = {}
|
|
||||||
for i in range(0, len(names), max_workers):
|
|
||||||
batch = names[i : i + max_workers]
|
|
||||||
for name in batch:
|
|
||||||
e = {**env, "ASTRAI_BUILD_SINGLE_EXT": name}
|
|
||||||
tag = name.replace(".", "_")
|
|
||||||
subdir = os.path.join(base, tag)
|
|
||||||
cmd = [
|
|
||||||
sys.executable,
|
|
||||||
__file__,
|
|
||||||
"build_ext",
|
|
||||||
"--build-temp",
|
|
||||||
os.path.join(subdir, "temp"),
|
|
||||||
"--build-lib",
|
|
||||||
os.path.join(subdir, "lib"),
|
|
||||||
]
|
|
||||||
procs[name] = subprocess.Popen(
|
|
||||||
cmd, env=e, stdout=subprocess.PIPE, stderr=subprocess.STDOUT
|
|
||||||
)
|
|
||||||
for name in batch:
|
|
||||||
out, _ = procs[name].communicate()
|
|
||||||
if procs[name].returncode != 0:
|
|
||||||
sys.stdout.write(out.decode())
|
|
||||||
raise RuntimeError(
|
|
||||||
f"parallel build failed for {name} "
|
|
||||||
f"(exit {procs[name].returncode})"
|
|
||||||
)
|
|
||||||
self._collect_extensions(
|
|
||||||
os.path.join(base, name.replace(".", "_"), "lib")
|
|
||||||
)
|
|
||||||
|
|
||||||
def _collect_extensions(self, sub_lib):
|
|
||||||
src = os.path.join(sub_lib, "astrai", "extension", "lib")
|
|
||||||
if not os.path.isdir(src):
|
|
||||||
return
|
|
||||||
dst = os.path.join(self.build_lib, "astrai", "extension", "lib")
|
|
||||||
os.makedirs(dst, exist_ok=True)
|
|
||||||
for f in os.listdir(src):
|
|
||||||
if f.endswith(".so"):
|
|
||||||
shutil.copy2(os.path.join(src, f), os.path.join(dst, f))
|
|
||||||
|
|
||||||
cmdclass["build_ext"] = ParallelBuildExtension
|
def _torch_cuda_version():
|
||||||
|
try:
|
||||||
|
import torch
|
||||||
|
|
||||||
if not cmdclass:
|
return torch.version.cuda
|
||||||
|
except Exception:
|
||||||
|
return None
|
||||||
|
|
||||||
class _NullBuildExt(_build_ext):
|
|
||||||
def build_extensions(self):
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
class _NullBuildExt(_build_ext):
|
||||||
|
def build_extensions(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
cmdclass = {}
|
||||||
|
|
||||||
|
if _should_build():
|
||||||
|
cmdclass["build_ext"] = _CMakeBuildExt
|
||||||
|
else:
|
||||||
cmdclass["build_ext"] = _NullBuildExt
|
cmdclass["build_ext"] = _NullBuildExt
|
||||||
|
|
||||||
setup(ext_modules=ext_modules, cmdclass=cmdclass)
|
setup(ext_modules=[], cmdclass=cmdclass)
|
||||||
|
|||||||
@@ -369,6 +369,19 @@ def test_dpo_missing_field_is_none(chat_tokenizer, builder):
|
|||||||
assert builder.build({"chosen": [], "rejected": []}, config, chat_tokenizer) is None
|
assert builder.build({"chosen": [], "rejected": []}, config, chat_tokenizer) is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("missing", ["chosen", "rejected"])
|
||||||
|
def test_dpo_partial_record_is_none(chat_tokenizer, builder, missing):
|
||||||
|
config = make_dpo_chat_config()
|
||||||
|
item = {
|
||||||
|
"chosen": [{"role": "assistant", "content": "Good"}],
|
||||||
|
"rejected": [{"role": "assistant", "content": "Bad"}],
|
||||||
|
}
|
||||||
|
item.pop(missing)
|
||||||
|
|
||||||
|
assert builder.build(item, config, chat_tokenizer) is None
|
||||||
|
assert builder.build_batch([item], config, chat_tokenizer) == [None]
|
||||||
|
|
||||||
|
|
||||||
def test_grpo_basic(chat_tokenizer, builder):
|
def test_grpo_basic(chat_tokenizer, builder):
|
||||||
config = make_grpo_config()
|
config = make_grpo_config()
|
||||||
item = {
|
item = {
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import pytest
|
|||||||
|
|
||||||
from astrai.extension import (
|
from astrai.extension import (
|
||||||
ATTN_BACKEND,
|
ATTN_BACKEND,
|
||||||
|
AttentionBackendFactory,
|
||||||
CudaBackend,
|
CudaBackend,
|
||||||
TorchNativeBackend,
|
TorchNativeBackend,
|
||||||
attn_backend,
|
attn_backend,
|
||||||
@@ -26,6 +27,26 @@ def test_attn_backend_context_with_enum():
|
|||||||
assert isinstance(get_backend(), TorchNativeBackend)
|
assert isinstance(get_backend(), TorchNativeBackend)
|
||||||
|
|
||||||
|
|
||||||
|
def test_attn_backend_context_with_registered_name():
|
||||||
|
with attn_backend("cuda"):
|
||||||
|
assert isinstance(get_backend(), CudaBackend)
|
||||||
|
assert isinstance(get_backend(), TorchNativeBackend)
|
||||||
|
|
||||||
|
|
||||||
|
def test_attention_backend_factory_lists_builtin_backends():
|
||||||
|
assert AttentionBackendFactory.list_registered() == [
|
||||||
|
"cuda",
|
||||||
|
"flash",
|
||||||
|
"torch_native",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_attn_backend_rejects_unknown_registered_name():
|
||||||
|
with pytest.raises(ValueError, match="Unknown component: 'unknown'"):
|
||||||
|
with attn_backend("unknown"):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
def test_attn_backend_context_with_class():
|
def test_attn_backend_context_with_class():
|
||||||
with attn_backend(CudaBackend):
|
with attn_backend(CudaBackend):
|
||||||
assert isinstance(get_backend(), CudaBackend)
|
assert isinstance(get_backend(), CudaBackend)
|
||||||
|
|||||||
@@ -8,9 +8,16 @@ import torch
|
|||||||
|
|
||||||
from astrai.extension import ATTN_BACKEND, attn_backend
|
from astrai.extension import ATTN_BACKEND, attn_backend
|
||||||
from astrai.inference.core.cache import PagePool
|
from astrai.inference.core.cache import PagePool
|
||||||
|
from astrai.inference.core.workspace import InferenceWorkspace
|
||||||
from tests.extension.conftest import D, skip_no_kernel
|
from tests.extension.conftest import D, skip_no_kernel
|
||||||
|
|
||||||
|
|
||||||
|
def _ws(pool: PagePool) -> InferenceWorkspace:
|
||||||
|
return InferenceWorkspace(
|
||||||
|
pool.max_batch_size, pool.max_seq_len, pool.device, pool.dtype
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@skip_no_kernel
|
@skip_no_kernel
|
||||||
def test_training_forward_matches_torch(cuda_model):
|
def test_training_forward_matches_torch(cuda_model):
|
||||||
"""Training forward (kv_cache=None) should produce identical logits.
|
"""Training forward (kv_cache=None) should produce identical logits.
|
||||||
@@ -62,11 +69,10 @@ def test_prefill_with_kv_cache_matches_torch(cuda_model):
|
|||||||
dtype=torch.bfloat16,
|
dtype=torch.bfloat16,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
ws = _ws(cache)
|
||||||
cache.task_alloc("t1", prompt_ids[0])
|
cache.task_alloc("t1", prompt_ids[0])
|
||||||
cache.task_alloc("t2", prompt_ids[1])
|
cache.task_alloc("t2", prompt_ids[1])
|
||||||
kv1 = cache.bind_tasks(
|
kv1 = cache.bind_tasks(["t1", "t2"], ws, start_pos=0)
|
||||||
["t1", "t2"], [len(prompt_ids[0]), len(prompt_ids[1])], device, start_pos=0
|
|
||||||
)
|
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
out_torch = model(
|
out_torch = model(
|
||||||
input_ids, input_mask=input_mask, kv_cache=kv1, position_ids=position_ids
|
input_ids, input_mask=input_mask, kv_cache=kv1, position_ids=position_ids
|
||||||
@@ -76,9 +82,7 @@ def test_prefill_with_kv_cache_matches_torch(cuda_model):
|
|||||||
cache.task_free("t2")
|
cache.task_free("t2")
|
||||||
cache.task_alloc("t1", prompt_ids[0])
|
cache.task_alloc("t1", prompt_ids[0])
|
||||||
cache.task_alloc("t2", prompt_ids[1])
|
cache.task_alloc("t2", prompt_ids[1])
|
||||||
kv2 = cache.bind_tasks(
|
kv2 = cache.bind_tasks(["t1", "t2"], ws, start_pos=0)
|
||||||
["t1", "t2"], [len(prompt_ids[0]), len(prompt_ids[1])], device, start_pos=0
|
|
||||||
)
|
|
||||||
with attn_backend(ATTN_BACKEND.CUDA):
|
with attn_backend(ATTN_BACKEND.CUDA):
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
out_cuda = model(
|
out_cuda = model(
|
||||||
@@ -129,11 +133,10 @@ def test_decode_mixed_seq_lens_matches_torch(cuda_model):
|
|||||||
input_mask[i, : len(p)] = True
|
input_mask[i, : len(p)] = True
|
||||||
position_ids[i, : len(p)] = torch.arange(len(p), device=device)
|
position_ids[i, : len(p)] = torch.arange(len(p), device=device)
|
||||||
|
|
||||||
|
ws = _ws(cache)
|
||||||
cache.task_alloc("t1", prompt_ids[0])
|
cache.task_alloc("t1", prompt_ids[0])
|
||||||
cache.task_alloc("t2", prompt_ids[1])
|
cache.task_alloc("t2", prompt_ids[1])
|
||||||
kv = cache.bind_tasks(
|
kv = cache.bind_tasks(["t1", "t2"], ws, start_pos=0)
|
||||||
["t1", "t2"], [len(prompt_ids[0]), len(prompt_ids[1])], device, start_pos=0
|
|
||||||
)
|
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
model(input_ids, input_mask=input_mask, kv_cache=kv, position_ids=position_ids)
|
model(input_ids, input_mask=input_mask, kv_cache=kv, position_ids=position_ids)
|
||||||
|
|
||||||
@@ -143,13 +146,15 @@ def test_decode_mixed_seq_lens_matches_torch(cuda_model):
|
|||||||
total_len = 9
|
total_len = 9
|
||||||
dec_mask = dec_pos[:, None, None] >= torch.arange(total_len, device=device)
|
dec_mask = dec_pos[:, None, None] >= torch.arange(total_len, device=device)
|
||||||
|
|
||||||
kv_t = cache.bind_tasks(["t1", "t2"], [9, 7], device)
|
cache.task_extend("t1", 8)
|
||||||
|
cache.task_extend("t2", 6)
|
||||||
|
kv_t = cache.bind_tasks(["t1", "t2"], ws)
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
out_torch = model(
|
out_torch = model(
|
||||||
dec_ids, input_mask=dec_mask, kv_cache=kv_t, position_ids=dec_pos
|
dec_ids, input_mask=dec_mask, kv_cache=kv_t, position_ids=dec_pos
|
||||||
)
|
)
|
||||||
|
|
||||||
kv_c = cache.bind_tasks(["t1", "t2"], [9, 7], device)
|
kv_c = cache.bind_tasks(["t1", "t2"], ws)
|
||||||
with attn_backend(ATTN_BACKEND.CUDA):
|
with attn_backend(ATTN_BACKEND.CUDA):
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
out_cuda = model(
|
out_cuda = model(
|
||||||
|
|||||||
@@ -6,10 +6,19 @@ from astrai.inference import (
|
|||||||
Allocator,
|
Allocator,
|
||||||
KVStorage,
|
KVStorage,
|
||||||
PagePool,
|
PagePool,
|
||||||
PrefixCache,
|
RadixCache,
|
||||||
ReqToTokenPool,
|
ReqToTokenPool,
|
||||||
page_hash,
|
page_hash,
|
||||||
)
|
)
|
||||||
|
from astrai.inference.core.workspace import InferenceWorkspace
|
||||||
|
|
||||||
|
|
||||||
|
def _ws(pool: PagePool) -> InferenceWorkspace:
|
||||||
|
"""Workspace sized to the pool (bind_tasks requires it)."""
|
||||||
|
return InferenceWorkspace(
|
||||||
|
pool.max_batch_size, pool.max_seq_len, pool.device, pool.dtype
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
# ---- page_hash ----
|
# ---- page_hash ----
|
||||||
|
|
||||||
@@ -68,12 +77,12 @@ def test_allocator_inc_ref_and_free():
|
|||||||
assert alloc._refs[p] == 0
|
assert alloc._refs[p] == 0
|
||||||
|
|
||||||
|
|
||||||
# ---- PrefixCache ----
|
# ---- RadixCache ----
|
||||||
|
|
||||||
|
|
||||||
def test_prefix_cache_lookup_returns_hits():
|
def test_prefix_cache_lookup_returns_hits():
|
||||||
token_ids = list(range(256))
|
token_ids = list(range(256))
|
||||||
prefix = PrefixCache(64)
|
prefix = RadixCache(64)
|
||||||
pages = [0, 1, 2, 3]
|
pages = [0, 1, 2, 3]
|
||||||
for i, p in enumerate(pages):
|
for i, p in enumerate(pages):
|
||||||
prefix.record(p, token_ids, i)
|
prefix.record(p, token_ids, i)
|
||||||
@@ -83,7 +92,7 @@ def test_prefix_cache_lookup_returns_hits():
|
|||||||
|
|
||||||
def test_prefix_cache_lookup_stops_at_first_miss():
|
def test_prefix_cache_lookup_stops_at_first_miss():
|
||||||
token_ids = list(range(256))
|
token_ids = list(range(256))
|
||||||
prefix = PrefixCache(64)
|
prefix = RadixCache(64)
|
||||||
prefix.record(0, token_ids, 0)
|
prefix.record(0, token_ids, 0)
|
||||||
prefix.record(1, [99] * 64, 1)
|
prefix.record(1, [99] * 64, 1)
|
||||||
hits = prefix.lookup(token_ids)
|
hits = prefix.lookup(token_ids)
|
||||||
@@ -93,14 +102,14 @@ def test_prefix_cache_lookup_stops_at_first_miss():
|
|||||||
|
|
||||||
def test_prefix_cache_ignores_partial_last_page():
|
def test_prefix_cache_ignores_partial_last_page():
|
||||||
token_ids = list(range(100))
|
token_ids = list(range(100))
|
||||||
prefix = PrefixCache(64)
|
prefix = RadixCache(64)
|
||||||
prefix.record(0, token_ids, 0)
|
prefix.record(0, token_ids, 0)
|
||||||
hits = prefix.lookup(token_ids)
|
hits = prefix.lookup(token_ids)
|
||||||
assert len(hits) == 1
|
assert len(hits) == 1
|
||||||
|
|
||||||
|
|
||||||
def test_prefix_cache_on_evict_clears_mappings():
|
def test_prefix_cache_on_evict_clears_mappings():
|
||||||
prefix = PrefixCache(64)
|
prefix = RadixCache(64)
|
||||||
prefix.record(0, list(range(64)), 0)
|
prefix.record(0, list(range(64)), 0)
|
||||||
assert 0 in prefix._page_to_hash
|
assert 0 in prefix._page_to_hash
|
||||||
prefix.evict(0)
|
prefix.evict(0)
|
||||||
@@ -108,12 +117,49 @@ def test_prefix_cache_on_evict_clears_mappings():
|
|||||||
|
|
||||||
|
|
||||||
def test_prefix_cache_has_page():
|
def test_prefix_cache_has_page():
|
||||||
prefix = PrefixCache(64)
|
prefix = RadixCache(64)
|
||||||
assert not prefix.has_page(0)
|
assert not prefix.has_page(0)
|
||||||
prefix.record(0, list(range(64)), 0)
|
prefix.record(0, list(range(64)), 0)
|
||||||
assert prefix.has_page(0)
|
assert prefix.has_page(0)
|
||||||
|
|
||||||
|
|
||||||
|
def test_prefix_cache_does_not_reuse_page_without_parent_prefix():
|
||||||
|
prefix = RadixCache(2)
|
||||||
|
prefix.record(0, [1, 2, 3, 4], 0)
|
||||||
|
prefix.record(1, [1, 2, 3, 4, 5, 6], 1)
|
||||||
|
prefix.record(2, [9, 10, 5, 6], 0)
|
||||||
|
prefix.record(3, [9, 10, 5, 6, 7, 8], 1)
|
||||||
|
assert prefix.lookup([1, 2, 3, 4, 5, 6]) == [0, 1]
|
||||||
|
assert prefix.lookup([9, 10, 5, 6, 7, 8]) == [2, 3]
|
||||||
|
|
||||||
|
|
||||||
|
def test_prefix_cache_shares_branch_prefix():
|
||||||
|
prefix = RadixCache(2)
|
||||||
|
prefix.record(0, [1, 2, 3, 4], 0)
|
||||||
|
prefix.record(1, [1, 2, 3, 4], 1)
|
||||||
|
prefix.record(2, [1, 2, 7, 8], 1)
|
||||||
|
assert prefix.lookup([1, 2, 3, 4]) == [0, 1]
|
||||||
|
assert prefix.lookup([1, 2, 7, 8]) == [0, 2]
|
||||||
|
prefix.evict(1)
|
||||||
|
assert prefix.lookup([1, 2, 3, 4]) == [0]
|
||||||
|
assert prefix.lookup([1, 2, 7, 8]) == [0, 2]
|
||||||
|
|
||||||
|
|
||||||
|
def test_prefix_cache_does_not_record_partial_page():
|
||||||
|
prefix = RadixCache(4)
|
||||||
|
prefix.record(0, [1, 2, 3, 4, 5, 6], 0)
|
||||||
|
prefix.record(1, [1, 2, 3, 4, 5, 6], 1)
|
||||||
|
assert prefix.lookup([1, 2, 3, 4, 5, 6]) == [0]
|
||||||
|
|
||||||
|
prefix.record(1, [1, 2, 3, 4, 5, 6, 7, 8], 1)
|
||||||
|
assert prefix.lookup([1, 2, 3, 4, 5, 6, 7, 8]) == [0, 1]
|
||||||
|
|
||||||
|
|
||||||
|
def test_page_pool_task_cacheable_ids_excludes_unmaterialized_tail():
|
||||||
|
pool = _make_paged_pool_ps64()
|
||||||
|
assert pool.task_cacheable_ids("missing", [1, 2], [3, 4]) == [1, 2, 3]
|
||||||
|
|
||||||
|
|
||||||
# ---- ReqToTokenPool ----
|
# ---- ReqToTokenPool ----
|
||||||
|
|
||||||
|
|
||||||
@@ -216,7 +262,7 @@ def test_page_pool_contiguous_bind_tasks_prefill():
|
|||||||
pool = _make_contiguous_pool()
|
pool = _make_contiguous_pool()
|
||||||
pool.task_alloc("t1", list(range(10)))
|
pool.task_alloc("t1", list(range(10)))
|
||||||
pool.task_alloc("t2", list(range(10)))
|
pool.task_alloc("t2", list(range(10)))
|
||||||
kv = pool.bind_tasks(["t1", "t2"], [10, 10], torch.device("cpu"), start_pos=0)
|
kv = pool.bind_tasks(["t1", "t2"], _ws(pool), start_pos=0)
|
||||||
assert kv.out_cache_loc.shape == (2, 10)
|
assert kv.out_cache_loc.shape == (2, 10)
|
||||||
assert kv.seq_lens.tolist() == [10, 10]
|
assert kv.seq_lens.tolist() == [10, 10]
|
||||||
assert kv.req_pool_indices.shape == (2,)
|
assert kv.req_pool_indices.shape == (2,)
|
||||||
@@ -226,7 +272,10 @@ def test_page_pool_contiguous_bind_tasks_decode():
|
|||||||
pool = _make_contiguous_pool()
|
pool = _make_contiguous_pool()
|
||||||
pool.task_alloc("t1", list(range(10)))
|
pool.task_alloc("t1", list(range(10)))
|
||||||
pool.task_alloc("t2", list(range(8)))
|
pool.task_alloc("t2", list(range(8)))
|
||||||
kv = pool.bind_tasks(["t1", "t2"], [11, 9], torch.device("cpu"))
|
# Simulate one decode extension so seq_lens advance to 11 and 9.
|
||||||
|
assert pool.task_extend("t1", 10)
|
||||||
|
assert pool.task_extend("t2", 8)
|
||||||
|
kv = pool.bind_tasks(["t1", "t2"], _ws(pool))
|
||||||
assert kv.out_cache_loc.shape == (2, 1)
|
assert kv.out_cache_loc.shape == (2, 1)
|
||||||
assert kv.seq_lens.tolist() == [11, 9]
|
assert kv.seq_lens.tolist() == [11, 9]
|
||||||
|
|
||||||
@@ -236,7 +285,7 @@ def test_page_pool_contiguous_bind_roundtrip():
|
|||||||
pool = _make_contiguous_pool(n_layers=1, n_kv_heads=2, head_dim=4)
|
pool = _make_contiguous_pool(n_layers=1, n_kv_heads=2, head_dim=4)
|
||||||
pool.task_alloc("t1", list(range(4)))
|
pool.task_alloc("t1", list(range(4)))
|
||||||
|
|
||||||
kv = pool.bind_tasks(["t1"], [4], torch.device("cpu"), start_pos=0)
|
kv = pool.bind_tasks(["t1"], _ws(pool), start_pos=0)
|
||||||
k = torch.randn(1, 4, 2, 4)
|
k = torch.randn(1, 4, 2, 4)
|
||||||
v = torch.randn(1, 4, 2, 4)
|
v = torch.randn(1, 4, 2, 4)
|
||||||
kv.k_buffer[0, kv.out_cache_loc] = k
|
kv.k_buffer[0, kv.out_cache_loc] = k
|
||||||
@@ -298,7 +347,7 @@ def test_page_pool_paged_bind_roundtrip():
|
|||||||
pool = _make_paged_pool(n_layers=1, n_kv_heads=2, head_dim=4)
|
pool = _make_paged_pool(n_layers=1, n_kv_heads=2, head_dim=4)
|
||||||
pool.task_alloc("t1", list(range(4)))
|
pool.task_alloc("t1", list(range(4)))
|
||||||
|
|
||||||
kv = pool.bind_tasks(["t1"], [4], torch.device("cpu"), start_pos=0)
|
kv = pool.bind_tasks(["t1"], _ws(pool), start_pos=0)
|
||||||
k = torch.randn(1, 4, 2, 4)
|
k = torch.randn(1, 4, 2, 4)
|
||||||
v = torch.randn(1, 4, 2, 4)
|
v = torch.randn(1, 4, 2, 4)
|
||||||
kv.k_buffer[0, kv.out_cache_loc] = k
|
kv.k_buffer[0, kv.out_cache_loc] = k
|
||||||
@@ -349,7 +398,7 @@ def test_page_pool_paged_ps64_bind_roundtrip():
|
|||||||
prompt = list(range(128))
|
prompt = list(range(128))
|
||||||
pool.task_alloc("t1", prompt)
|
pool.task_alloc("t1", prompt)
|
||||||
|
|
||||||
kv = pool.bind_tasks(["t1"], [128], torch.device("cpu"), start_pos=0)
|
kv = pool.bind_tasks(["t1"], _ws(pool), start_pos=0)
|
||||||
k = torch.randn(1, 128, 2, 4)
|
k = torch.randn(1, 128, 2, 4)
|
||||||
v = torch.randn(1, 128, 2, 4)
|
v = torch.randn(1, 128, 2, 4)
|
||||||
kv.k_buffer[0, kv.out_cache_loc] = k
|
kv.k_buffer[0, kv.out_cache_loc] = k
|
||||||
|
|||||||
@@ -205,6 +205,36 @@ def test_run_batch_returns_token_sequences(device):
|
|||||||
scheduler.stop()
|
scheduler.stop()
|
||||||
|
|
||||||
|
|
||||||
|
def test_run_batch_tokens_match_full_sequence_forward(device):
|
||||||
|
scheduler, _tok, model = _make_real_scheduler(device)
|
||||||
|
prompt = [10, 20, 30, 40]
|
||||||
|
try:
|
||||||
|
expected = []
|
||||||
|
sequence = list(prompt)
|
||||||
|
for _ in range(2):
|
||||||
|
input_ids = torch.tensor([sequence], dtype=torch.long, device=device)
|
||||||
|
position_ids = torch.arange(len(sequence), device=device).unsqueeze(0)
|
||||||
|
input_mask = torch.ones(
|
||||||
|
1, len(sequence), len(sequence), dtype=torch.bool, device=device
|
||||||
|
).tril()
|
||||||
|
with torch.inference_mode():
|
||||||
|
logits = model(
|
||||||
|
input_ids,
|
||||||
|
input_mask=input_mask,
|
||||||
|
position_ids=position_ids,
|
||||||
|
)["logits"][:, -1, :]
|
||||||
|
token = logits.argmax(dim=-1).item()
|
||||||
|
expected.append(token)
|
||||||
|
sequence.append(token)
|
||||||
|
|
||||||
|
result = scheduler.run_batch(
|
||||||
|
prompt_ids_list=[prompt], max_tokens=2, temperature=0
|
||||||
|
)
|
||||||
|
assert result == [expected]
|
||||||
|
finally:
|
||||||
|
scheduler.stop()
|
||||||
|
|
||||||
|
|
||||||
def test_run_batch_return_logprobs_aligned(device):
|
def test_run_batch_return_logprobs_aligned(device):
|
||||||
"""return_logprobs=True gives (token_ids, logprobs) tuples with equal len."""
|
"""return_logprobs=True gives (token_ids, logprobs) tuples with equal len."""
|
||||||
scheduler, _tok, _model = _make_real_scheduler(device)
|
scheduler, _tok, _model = _make_real_scheduler(device)
|
||||||
|
|||||||
@@ -22,6 +22,8 @@ def test_task_next_pos():
|
|||||||
task.input_tokens = 5
|
task.input_tokens = 5
|
||||||
assert task.next_pos == 5
|
assert task.next_pos == 5
|
||||||
task.output_ids.append(4)
|
task.output_ids.append(4)
|
||||||
|
assert task.next_pos == 5
|
||||||
|
task.output_ids.append(5)
|
||||||
assert task.next_pos == 6
|
assert task.next_pos == 6
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -265,6 +265,68 @@ def test_moe_defaults_preserve_normalized_routing():
|
|||||||
assert model.layers[0].mlp.norm_topk_prob is True
|
assert model.layers[0].mlp.norm_topk_prob is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_moe_router_stats_in_output_during_training():
|
||||||
|
"""Verify forward output carries per-layer router_stats in training mode."""
|
||||||
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
|
|
||||||
|
config = AutoRegressiveLMConfig(
|
||||||
|
**TINY_CONFIG,
|
||||||
|
ffn_type="moe",
|
||||||
|
n_routed_experts=4,
|
||||||
|
n_shared_experts=1,
|
||||||
|
n_activated_experts=2,
|
||||||
|
topk_method="greedy",
|
||||||
|
)
|
||||||
|
model = AutoRegressiveLM(config)
|
||||||
|
model.train()
|
||||||
|
input_ids = torch.randint(0, config.vocab_size, (2, 8))
|
||||||
|
|
||||||
|
with torch.enable_grad():
|
||||||
|
outputs = model(input_ids)
|
||||||
|
|
||||||
|
stats = outputs["router_stats"]
|
||||||
|
assert isinstance(stats, list)
|
||||||
|
assert len(stats) == config.num_hidden_layers
|
||||||
|
for s in stats:
|
||||||
|
assert s["probs"].shape == (2 * 8, 4) # (N, n_routed_experts)
|
||||||
|
assert s["topk_indices"].shape == (2 * 8, 2) # (N, n_activated_experts)
|
||||||
|
|
||||||
|
|
||||||
|
def test_moe_router_stats_absent_in_eval():
|
||||||
|
"""Verify no router_stats are emitted outside training."""
|
||||||
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
|
|
||||||
|
config = AutoRegressiveLMConfig(
|
||||||
|
**TINY_CONFIG,
|
||||||
|
ffn_type="moe",
|
||||||
|
n_routed_experts=4,
|
||||||
|
n_shared_experts=1,
|
||||||
|
n_activated_experts=2,
|
||||||
|
)
|
||||||
|
model = AutoRegressiveLM(config)
|
||||||
|
model.eval()
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
outputs = model(torch.randint(0, config.vocab_size, (2, 8)))
|
||||||
|
|
||||||
|
assert "router_stats" not in outputs
|
||||||
|
|
||||||
|
|
||||||
|
def test_no_router_stats_for_mlp_model():
|
||||||
|
"""Verify pure MLP models emit no router_stats and no aux_loss."""
|
||||||
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
|
|
||||||
|
config = AutoRegressiveLMConfig(**TINY_CONFIG, ffn_type="mlp")
|
||||||
|
model = AutoRegressiveLM(config)
|
||||||
|
model.train()
|
||||||
|
|
||||||
|
with torch.enable_grad():
|
||||||
|
outputs = model(torch.randint(0, config.vocab_size, (2, 8)))
|
||||||
|
|
||||||
|
assert "router_stats" not in outputs
|
||||||
|
assert "aux_loss" not in outputs
|
||||||
|
|
||||||
|
|
||||||
def test_moe_aux_loss_only_emitted_during_training():
|
def test_moe_aux_loss_only_emitted_during_training():
|
||||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,315 @@
|
|||||||
|
"""Smoke tests for MoE aux loss and diagnostic metrics integration.
|
||||||
|
|
||||||
|
Does NOT load real data or weights. Uses a tiny randomly-initialized
|
||||||
|
MoE model and verifies that aux loss computation and MoE routing
|
||||||
|
diagnostics flow end‑to‑end through the strategy layer.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
|
from astrai.model.transformer import AutoRegressiveLM
|
||||||
|
from astrai.trainer.strategy import (
|
||||||
|
SEQStrategy,
|
||||||
|
SFTStrategy,
|
||||||
|
StrategyFactory,
|
||||||
|
_collect_moe_diagnostics,
|
||||||
|
)
|
||||||
|
from tests.helpers import TINY_CONFIG
|
||||||
|
|
||||||
|
|
||||||
|
def _make_tiny_moe_config(**overrides) -> AutoRegressiveLMConfig:
|
||||||
|
return AutoRegressiveLMConfig(
|
||||||
|
**{
|
||||||
|
**TINY_CONFIG,
|
||||||
|
"ffn_type": "moe",
|
||||||
|
"n_routed_experts": 4,
|
||||||
|
"n_shared_experts": 1,
|
||||||
|
"n_activated_experts": 2,
|
||||||
|
"topk_method": "greedy",
|
||||||
|
**overrides,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_model(config=None) -> AutoRegressiveLM:
|
||||||
|
if config is None:
|
||||||
|
config = _make_tiny_moe_config()
|
||||||
|
return AutoRegressiveLM(config)
|
||||||
|
|
||||||
|
|
||||||
|
def _router_stats(probs, topk_indices):
|
||||||
|
return {"probs": probs, "topk_indices": topk_indices}
|
||||||
|
|
||||||
|
|
||||||
|
def test_collect_moe_diagnostics_returns_all_keys():
|
||||||
|
"""_collect_moe_diagnostics should return the four expected keys."""
|
||||||
|
# Simulate two MoE layers with uniform routing probabilities
|
||||||
|
probs = torch.ones(128, 4) / 4.0
|
||||||
|
topk = torch.zeros(128, 2, dtype=torch.long)
|
||||||
|
diag = _collect_moe_diagnostics([_router_stats(probs, topk)] * 2)
|
||||||
|
|
||||||
|
assert set(diag.keys()) == {
|
||||||
|
"router_entropy",
|
||||||
|
"dead_expert_fraction",
|
||||||
|
"load_imbalance_mean",
|
||||||
|
"load_imbalance_max",
|
||||||
|
}
|
||||||
|
for v in diag.values():
|
||||||
|
assert isinstance(v, float)
|
||||||
|
|
||||||
|
|
||||||
|
def test_collect_moe_diagnostics_empty_list():
|
||||||
|
"""Empty list returns empty dict."""
|
||||||
|
assert _collect_moe_diagnostics([]) == {}
|
||||||
|
|
||||||
|
|
||||||
|
def test_collect_moe_diagnostics_uniform_routing():
|
||||||
|
"""Uniform routing with top_k=2 → tie-breaking by index.
|
||||||
|
|
||||||
|
torch.topk breaks ties by index, so with equal probabilities
|
||||||
|
experts 0 and 1 always win over experts 2 and 3:
|
||||||
|
- dead_expert_fraction = 2/4 = 0.5
|
||||||
|
- load_ratios = [2, 2, 0, 0] → |ratio-1| = [1, 1, 1, 1] → mean = 1.0
|
||||||
|
- load_imbalance_max = 2.0
|
||||||
|
"""
|
||||||
|
probs = torch.ones(128, 4) / 4.0
|
||||||
|
topk = torch.tensor([[0, 1]] * 128)
|
||||||
|
diag = _collect_moe_diagnostics([_router_stats(probs, topk)])
|
||||||
|
|
||||||
|
assert diag["dead_expert_fraction"] == pytest.approx(0.5, abs=1e-6)
|
||||||
|
assert diag["load_imbalance_mean"] == pytest.approx(1.0, abs=1e-6)
|
||||||
|
assert diag["load_imbalance_max"] == pytest.approx(2.0, abs=1e-6)
|
||||||
|
|
||||||
|
|
||||||
|
def test_collect_moe_diagnostics_max_entropy():
|
||||||
|
"""Uniform probabilities should give log(num_experts) entropy."""
|
||||||
|
num_experts = 4
|
||||||
|
probs = torch.ones(128, num_experts) / num_experts
|
||||||
|
topk = torch.zeros(128, 2, dtype=torch.long)
|
||||||
|
diag = _collect_moe_diagnostics([_router_stats(probs, topk)])
|
||||||
|
expected_entropy = float(torch.log(torch.tensor(num_experts, dtype=torch.float32)))
|
||||||
|
assert diag["router_entropy"] == pytest.approx(expected_entropy, abs=1e-5)
|
||||||
|
|
||||||
|
|
||||||
|
def test_moe_metrics_flow_through_wrapped_model(device):
|
||||||
|
"""DDP-like wrappers (no .config / get_moe_router_probs) still collect MoE metrics."""
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
from astrai.trainer.strategy import SEQStrategy
|
||||||
|
|
||||||
|
class ForwardOnlyWrapper(nn.Module):
|
||||||
|
def __init__(self, model):
|
||||||
|
super().__init__()
|
||||||
|
self.module = model
|
||||||
|
|
||||||
|
def forward(self, *args, **kwargs):
|
||||||
|
return self.module(*args, **kwargs)
|
||||||
|
|
||||||
|
config = _make_tiny_moe_config()
|
||||||
|
model = AutoRegressiveLM(config).to(device)
|
||||||
|
wrapped = ForwardOnlyWrapper(model)
|
||||||
|
wrapped.train()
|
||||||
|
|
||||||
|
strategy = SEQStrategy(wrapped, device, moe_aux_loss_coef=0.01)
|
||||||
|
output = strategy.compute_loss_output(
|
||||||
|
{
|
||||||
|
"input_ids": torch.randint(0, config.vocab_size, (2, 8)),
|
||||||
|
"target_ids": torch.randint(0, config.vocab_size, (2, 8)),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert "moe_aux_loss" in output["metrics"]
|
||||||
|
assert "router_entropy" in strategy._moe_metrics
|
||||||
|
|
||||||
|
|
||||||
|
class TestSEQStrategyMoE:
|
||||||
|
"""End‑to‑end tests for SEQStrategy with MoE aux loss."""
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def setup(self, device):
|
||||||
|
self.device = device
|
||||||
|
self.config = _make_tiny_moe_config()
|
||||||
|
self.model = _make_model(self.config).to(device)
|
||||||
|
self.model.train()
|
||||||
|
|
||||||
|
def _make_batch(self, batch_size=2, seq_len=8):
|
||||||
|
vocab = self.config.vocab_size
|
||||||
|
input_ids = torch.randint(0, vocab, (batch_size, seq_len))
|
||||||
|
# target = input shifted right
|
||||||
|
target_ids = torch.randint(0, vocab, (batch_size, seq_len))
|
||||||
|
return {"input_ids": input_ids, "target_ids": target_ids}
|
||||||
|
|
||||||
|
def test_compute_loss_returns_scalar(self):
|
||||||
|
"""compute_loss should return a scalar tensor."""
|
||||||
|
strategy = SEQStrategy(
|
||||||
|
self.model,
|
||||||
|
self.device,
|
||||||
|
moe_aux_loss_coef=0.01,
|
||||||
|
)
|
||||||
|
loss = strategy.compute_loss(self._make_batch())
|
||||||
|
assert loss.ndim == 0
|
||||||
|
assert loss.requires_grad
|
||||||
|
|
||||||
|
def test_compute_loss_output_has_metrics(self):
|
||||||
|
"""compute_loss_output dict with moe_aux_loss_coef > 0 includes MoE metrics."""
|
||||||
|
strategy = SEQStrategy(
|
||||||
|
self.model,
|
||||||
|
self.device,
|
||||||
|
moe_aux_loss_coef=0.01,
|
||||||
|
)
|
||||||
|
output = strategy.compute_loss_output(self._make_batch())
|
||||||
|
|
||||||
|
assert "loss" in output
|
||||||
|
assert "metrics" in output
|
||||||
|
assert output["loss"].ndim == 0
|
||||||
|
assert output["loss"].requires_grad
|
||||||
|
|
||||||
|
metrics = output["metrics"]
|
||||||
|
# MoE metrics should appear when coef > 0 and model has MoE layers
|
||||||
|
for key in ("moe_aux_loss", "moe_aux_loss_weighted", "task_loss", "loss"):
|
||||||
|
assert key in metrics, f"Missing metric: {key}"
|
||||||
|
assert isinstance(metrics[key], float)
|
||||||
|
|
||||||
|
def test_moe_metrics_populated_after_forward(self):
|
||||||
|
"""strategy._moe_metrics populated after compute_loss_output."""
|
||||||
|
strategy = SEQStrategy(
|
||||||
|
self.model,
|
||||||
|
self.device,
|
||||||
|
moe_aux_loss_coef=0.01,
|
||||||
|
)
|
||||||
|
strategy.compute_loss_output(self._make_batch())
|
||||||
|
|
||||||
|
moe_metrics = strategy._moe_metrics
|
||||||
|
assert moe_metrics, "_moe_metrics should not be empty for MoE model"
|
||||||
|
for key in (
|
||||||
|
"aux_loss",
|
||||||
|
"router_entropy",
|
||||||
|
"dead_expert_fraction",
|
||||||
|
"load_imbalance_mean",
|
||||||
|
"load_imbalance_max",
|
||||||
|
):
|
||||||
|
assert key in moe_metrics, f"Missing _moe_metrics key: {key}"
|
||||||
|
assert isinstance(moe_metrics[key], float)
|
||||||
|
|
||||||
|
def test_zero_coef_zeroes_weighted_aux(self):
|
||||||
|
"""moe_aux_loss_coef=0 → weighted_aux_loss is zero, task_loss == loss."""
|
||||||
|
strategy = SEQStrategy(
|
||||||
|
self.model,
|
||||||
|
self.device,
|
||||||
|
moe_aux_loss_coef=0.0,
|
||||||
|
)
|
||||||
|
output = strategy.compute_loss_output(self._make_batch())
|
||||||
|
metrics = output["metrics"]
|
||||||
|
|
||||||
|
# task_loss and loss should be equal (aux weighted by zero)
|
||||||
|
assert "task_loss" in metrics
|
||||||
|
assert "loss" in metrics
|
||||||
|
assert metrics["loss"] == pytest.approx(metrics["task_loss"], abs=1e-6)
|
||||||
|
|
||||||
|
# weighted aux loss is zero
|
||||||
|
assert metrics.get("moe_aux_loss_weighted") == pytest.approx(0.0, abs=1e-6)
|
||||||
|
|
||||||
|
# MoE diagnostics are still collected (monitoring purposes)
|
||||||
|
assert strategy._moe_metrics
|
||||||
|
assert "router_entropy" in strategy._moe_metrics
|
||||||
|
|
||||||
|
def test_aux_loss_added_to_total_loss(self):
|
||||||
|
"""Total loss > task_loss when moe_aux_loss_coef > 0."""
|
||||||
|
strategy = SEQStrategy(
|
||||||
|
self.model,
|
||||||
|
self.device,
|
||||||
|
moe_aux_loss_coef=0.01,
|
||||||
|
)
|
||||||
|
output = strategy.compute_loss_output(self._make_batch())
|
||||||
|
assert output["metrics"]["loss"] > output["metrics"]["task_loss"] + 1e-12
|
||||||
|
|
||||||
|
def test_factory_creates_strategy_with_coef(self):
|
||||||
|
"""StrategyFactory.create passes moe_aux_loss_coef to strategy."""
|
||||||
|
strategy = StrategyFactory.create(
|
||||||
|
"seq",
|
||||||
|
model=self.model,
|
||||||
|
device=self.device,
|
||||||
|
moe_aux_loss_coef=0.02,
|
||||||
|
)
|
||||||
|
assert strategy.moe_aux_loss_coef == 0.02
|
||||||
|
|
||||||
|
def test_no_aux_loss_for_mlp_model(self):
|
||||||
|
"""Pure MLP model: model outputs no aux_loss → no MoE metrics."""
|
||||||
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
|
|
||||||
|
mlp_config = AutoRegressiveLMConfig(**{**TINY_CONFIG, "ffn_type": "mlp"})
|
||||||
|
mlp_model = AutoRegressiveLM(mlp_config).to(self.device)
|
||||||
|
mlp_model.train()
|
||||||
|
|
||||||
|
strategy = SEQStrategy(
|
||||||
|
mlp_model,
|
||||||
|
self.device,
|
||||||
|
moe_aux_loss_coef=0.01,
|
||||||
|
)
|
||||||
|
output = strategy.compute_loss_output(self._make_batch())
|
||||||
|
metrics = output["metrics"]
|
||||||
|
|
||||||
|
assert "moe_aux_loss" not in metrics
|
||||||
|
assert "moe_aux_loss_weighted" not in metrics
|
||||||
|
assert metrics["loss"] == pytest.approx(metrics["task_loss"], abs=1e-6)
|
||||||
|
assert strategy._moe_metrics == {}
|
||||||
|
|
||||||
|
|
||||||
|
class TestSFTStrategyMoE:
|
||||||
|
"""End‑to‑end tests for SFTStrategy with MoE aux loss."""
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def setup(self, device):
|
||||||
|
self.device = device
|
||||||
|
self.config = _make_tiny_moe_config()
|
||||||
|
self.model = _make_model(self.config).to(device)
|
||||||
|
self.model.train()
|
||||||
|
|
||||||
|
def _make_batch(self, batch_size=2, seq_len=8):
|
||||||
|
vocab = self.config.vocab_size
|
||||||
|
input_ids = torch.randint(0, vocab, (batch_size, seq_len))
|
||||||
|
target_ids = torch.randint(0, vocab, (batch_size, seq_len))
|
||||||
|
position_ids = torch.arange(seq_len).unsqueeze(0).expand(batch_size, -1)
|
||||||
|
loss_mask = torch.ones(batch_size, seq_len, dtype=torch.bool)
|
||||||
|
return {
|
||||||
|
"input_ids": input_ids,
|
||||||
|
"target_ids": target_ids,
|
||||||
|
"position_ids": position_ids,
|
||||||
|
"loss_mask": loss_mask,
|
||||||
|
}
|
||||||
|
|
||||||
|
def test_compute_loss_output_with_aux_loss(self):
|
||||||
|
"""SFTStrategy produces MoE metrics when coef > 0."""
|
||||||
|
strategy = SFTStrategy(
|
||||||
|
self.model,
|
||||||
|
self.device,
|
||||||
|
moe_aux_loss_coef=0.01,
|
||||||
|
)
|
||||||
|
output = strategy.compute_loss_output(self._make_batch())
|
||||||
|
|
||||||
|
metrics = output["metrics"]
|
||||||
|
assert "moe_aux_loss" in metrics
|
||||||
|
assert "moe_aux_loss_weighted" in metrics
|
||||||
|
assert metrics["loss"] > metrics["task_loss"] + 1e-12
|
||||||
|
|
||||||
|
moe_metrics = strategy._moe_metrics
|
||||||
|
assert "router_entropy" in moe_metrics
|
||||||
|
assert "dead_expert_fraction" in moe_metrics
|
||||||
|
|
||||||
|
def test_sft_zero_coef_zeroes_weighted_aux(self):
|
||||||
|
"""SFTStrategy with zero coef: weighted aux is zero, loss == task_loss."""
|
||||||
|
strategy = SFTStrategy(
|
||||||
|
self.model,
|
||||||
|
self.device,
|
||||||
|
moe_aux_loss_coef=0.0,
|
||||||
|
)
|
||||||
|
output = strategy.compute_loss_output(self._make_batch())
|
||||||
|
metrics = output["metrics"]
|
||||||
|
|
||||||
|
assert metrics["loss"] == pytest.approx(metrics["task_loss"], abs=1e-6)
|
||||||
|
assert metrics.get("moe_aux_loss_weighted") == pytest.approx(0.0, abs=1e-6)
|
||||||
|
# Diagnostics still collected
|
||||||
|
assert strategy._moe_metrics
|
||||||
|
assert "router_entropy" in strategy._moe_metrics
|
||||||
Reference in New Issue
Block a user