18 Commits
Author SHA1 Message Date
ViperEkura 6f09b1d2ee docs : clarify radix cache architecture
- document exact page-aligned radix prefix matching
- explain partial-page ownership and materialized KV boundaries
- remove bilingual wording from project overview
2026-08-06 11:50:45 +08:00
ViperEkura b2230fefd8 feat : add radix prefix cache
- replace hash-only lookup with page-granular radix matching
- keep partial pages private and cache only materialized KV prefixes
- integrate completed-request caching and add radix behavior tests
2026-08-06 11:45:52 +08:00
ViperEkura 654e6eb0d1 fix : correct prefill sampling and record alignment
- sample the first token from prefill logits without duplicating the prompt tail
- reject incomplete multi-output records before preprocessing alignment
- cover cached generation and partial DPO records with regression tests
2026-08-05 22:20:29 +08:00
ViperEkura a317a4756b refactor: stateless MoE routing with grouped dispatch
- replace per-expert mask scan with sort+bincount grouped dispatch
- carry router stats in forward output instead of module state
- keep MoE diagnostics working under DDP/FSDP wrappers
- remove unused _load_balancing_loss helper
2026-08-05 18:42:12 +08:00
ViperEkura 9b7e6c205f feat: add moe auxloss and metrics 2026-08-05 18:12:28 +08:00
ViperEkura 602b5ce216 docs : add project capability overview
- summarize the end-to-end model lifecycle
- add matching capability tables in both READMEs
2026-08-05 15:47:42 +08:00
ViperEkura 8152760b5f refactor : use factory for attention backends
- register built-in backends through BaseFactory
- derive benchmark choices from registered backends
- cover string selection and invalid backend names
2026-08-05 15:37:22 +08:00
ViperEkura 8c052c99ee feat: add optional FlashAttention (FA2/FA3) backend
- add FlashAttnBackend (ATTN_BACKEND.FLASH) using flash_attn_func with KV-cache gather + GQA, mirroring TorchNativeBackend
- add flash_attn_available() probe gated on compute capability plus a real-kernel smoke test, cached at first use
- lazy-import flash-attn via importlib so it stays an optional dependency, raising clear errors when unusable
- add 'flash' optional extra (flash-attn>=2.6) and export the new backend
2026-08-05 15:27:26 +08:00
ViperEkura 2667b8116d refactor: unify paged and contiguous attention kernels via KVSource policy
- merge AttentionParams and PagedAttentionParams into one struct
- add attn_kv_source.cuh with ContigKV/PagedKV addressing policies
- template prefill/decode kernels (MMA + scalar) on the KV policy, deleting the four duplicated attn_paged_*.cuh variants
- template dispatcher launchers on KV; single combine kernel
- verify: all correctness tests pass and SASS matches baseline (no perf regression)
2026-08-05 14:06:13 +08:00
ViperEkura 6dffb0305a fix: satisfy ruff format and import lint in setup.py
- Merge nested if for CUDA version mismatch check
- Convert try-except-pass to return None (S110)
- Apply ruff format
2026-08-04 21:32:33 +08:00
ViperEkura 49a9c6b3d2 build: migrate CUDA kernel build to CMake
Replace torch CUDAExtension/ParallelBuildExtension with a CMake-based build. Each kernel compiles as an independent pybind11 module in parallel via cmake --build -j, outputting to astrai/extension/lib.

- Add csrc/CMakeLists.txt (5 kernel targets, torch/pybind11 linking)
- setup.py: _CMakeBuildExt invokes cmake; auto-detect CUDA arch via torch
- Remove csrc/build.py (REGISTRY/build flags now in CMakeLists)
- Fix rel-err eps in attn_test.cu (1e-8 -> 1e-4, bf16 scale)
- Update docs/developer/cuda_kernels.md build section
- .gitignore: allow csrc/CMakeLists.txt
2026-08-04 21:27:22 +08:00
ViperEkura cdf9145ecf docs: align CUDA kernel and RoPE docs with code
- Fix rotary docs to describe cos/sin freqs_cis table, not complex buffer
- Replace attn_prefill with attn_paged_prefill for the CudaBackend path
- Register attn_paged_prefill in kernel overview, layout, and module list
- Add qo_indptr and InferenceWorkspace to architecture class diagram
- Add FrequencyPenaltyStrategy to sampling design patterns
2026-08-03 20:54:40 +08:00
ViperEkura 85f0461b3b docs: update license refs from GPL-3.0 to Apache-2.0 2026-08-03 20:21:36 +08:00
ViperEkura 9f0e9195f7 Update LICENSE 2026-08-03 20:18:27 +08:00
ViperEkura 88751d0b08 refactor: share prefill+decode step between scheduler paths
- Extract _step() as the single prefill-group + task_extend + decode primitive
- _run_generation_loop and run_batch now both call it, so the two cannot drift
- run_batch now records prefix hashes (paged mode) and uses input order for
  decode, matching the loop thread
2026-08-03 13:45:27 +08:00
ViperEkura d0e5d910de perf: reduce remaining per-step allocations
- hoist prefill qo_indptr into the workspace so CudaBackend.fwd_prefill does not rebuild it per layer
- cache has_freq in SamplingBatchInfo to drop the per-step GPU any() sync
- drop pin_memory host staging for input_ids; sync copy suffices for a small batch
2026-08-03 01:10:06 +08:00
ViperEkura a03504a280 perf: preallocate inference decode buffers
- add InferenceWorkspace with fixed-shape per-step buffers (input_ids, decode mask, KV bind metadata) for CUDA-graph capture
- bind_tasks derives seq_lens from the pool's own _task_len tracking, dropping the seq_lens parameter
- update decode metadata in-place (position_ids, seq_lens, kv_indptr) instead of re-allocating per step
- task_extend advances _task_len in contiguous mode so the pool tracks current length
- skip log_softmax when logprobs are not requested
2026-08-03 00:55:26 +08:00
ViperEkura d033b2ef0f perf: cache per-step decode tensor construction
- SamplingBatchInfo: sample params built once per task set (top_k int32, pinned async H2D)
- position_ids advances by +1 on steady-state decode instead of re-building
- DecodeBindCache: bind_tasks increments seq_lens/kv_indptr, reuses req_pool_indices
- saves ~240us of python/launch overhead per decode step
2026-08-02 20:32:53 +08:00
59 changed files with 2404 additions and 2097 deletions
+1
View File
@@ -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
View File
@@ -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).
--- ---
+201 -674
View File
@@ -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>.
+15 -11
View File
@@ -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. |---|---|
- 🔬 **ResearchFriendly**: 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).
--- ---
+1
View File
@@ -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:
+4
View File
@@ -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",
+183 -14
View File
@@ -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)
+2 -2
View File
@@ -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",
+2 -2
View File
@@ -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
View File
@@ -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 ----
+146 -79
View File
@@ -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()
+94 -79
View File
@@ -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:
+2 -1
View File
@@ -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:
+111
View File
@@ -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
+1 -1
View File
@@ -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
+7 -2
View File
@@ -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"),
}
+54 -18
View File
@@ -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,
}
+5 -1
View File
@@ -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
+12 -7
View File
@@ -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)
] ]
+20
View File
@@ -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
View File
@@ -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:
+10
View File
@@ -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):
+74
View File
@@ -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()
-76
View File
@@ -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
View File
@@ -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;
}; };
+33 -17
View File
@@ -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);
} }
+24 -17
View File
@@ -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
View File
@@ -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
} }
+2 -2
View File
@@ -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));
+159
View File
@@ -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};
}
};
+1 -1
View File
@@ -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);
-151
View File
@@ -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;
}
}
}
+1 -1
View File
@@ -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,
-126
View File
@@ -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;
}
}
}
+25 -20
View File
@@ -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++)
+29 -21
View File
@@ -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*>(
+6 -6
View File
@@ -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;
+2 -2
View File
@@ -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;
+1 -1
View File
@@ -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
View File
@@ -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)。
--- ---
+40 -8
View File
@@ -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, BaseStrategyGRPOStrategy, StrategyFactory, BaseSchedulerWSDScheduler, SchedulerFactory, TrainCallback(Protocol)MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow | | **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategyGRPOStrategy, StrategyFactory, BaseSchedulerWSDScheduler, 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, BaseSamplingStrategySamplingPipeline, 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, BaseSamplingStrategySamplingPipeline, 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 |
+27 -14
View File
@@ -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
+13 -9
View File
@@ -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
View File
@@ -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
+6 -4
View File
@@ -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
View File
@@ -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 = ["."]
+5 -10
View File
@@ -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(
+8
View File
@@ -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)
+108 -102
View File
@@ -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)
+13
View File
@@ -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 = {
+21
View File
@@ -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)
+16 -11
View File
@@ -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(
+61 -12
View File
@@ -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
+30
View File
@@ -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)
+2
View File
@@ -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
+62
View File
@@ -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
+315
View File
@@ -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 endtoend 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:
"""Endtoend 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:
"""Endtoend 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