Compare commits
466
Commits
v1.3.4
...
3d3ea47d37
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3d3ea47d37 | ||
|
|
f7f14d0e5f | ||
|
|
7580d80d45 | ||
|
|
cb51a3587b | ||
|
|
1bcd8f53ab | ||
|
|
0d0dc64884 | ||
|
|
6ac3b51496 | ||
|
|
3406157431 | ||
|
|
0dd9a417b7 | ||
|
|
a01c1fd427 | ||
|
|
f8d9ab344d | ||
|
|
3fb4b8ab13 | ||
|
|
b5afe3d7a4 | ||
|
|
69f35c46e0 | ||
|
|
0378e62e17 | ||
|
|
5244f1a8fc | ||
|
|
5104638447 | ||
|
|
a711d9f478 | ||
|
|
15862d4b56 | ||
|
|
f9efb705b8 | ||
|
|
c6a82a5029 | ||
|
|
a5b238dd86 | ||
|
|
da6d94492d | ||
|
|
71b6e3aaaf | ||
|
|
f95722a277 | ||
|
|
9f48cb8928 | ||
|
|
9b58fef222 | ||
|
|
c5fba9c238 | ||
|
|
cd31f1f62f | ||
|
|
a5a3cc1fc2 | ||
|
|
d565d44c43 | ||
|
|
596c35fd71 | ||
|
|
47b3ed4e44 | ||
|
|
c1d05ae11d | ||
|
|
cf4f5ab9f6 | ||
|
|
3416f98c58 | ||
|
|
d28552f878 | ||
|
|
be90dfe2bd | ||
|
|
a33ca04f60 | ||
|
|
7f0e8bb8c2 | ||
|
|
0c1b7664c1 | ||
|
|
3fa7e66676 | ||
|
|
ca50fe4721 | ||
|
|
d9240ab149 | ||
|
|
d7cd69fef5 | ||
|
|
9bff61fb91 | ||
|
|
0b661bae85 | ||
|
|
ae9fd546ef | ||
|
|
e3ea850dc9 | ||
|
|
6e5088cc7d | ||
|
|
cbc584470d | ||
|
|
cb60713a72 | ||
|
|
c52a2487ae | ||
|
|
49aaa9a714 | ||
|
|
056c1382ff | ||
|
|
f163520fff | ||
|
|
1b1f1a0707 | ||
|
|
184fbbce5c | ||
|
|
02469887f5 | ||
|
|
05739629fc | ||
|
|
e0f7fa8e13 | ||
|
|
af25833fab | ||
|
|
6572be4f98 | ||
|
|
81788faef4 | ||
|
|
0e7fe57d96 | ||
|
|
55ee258e95 | ||
|
|
ef1bb6f401 | ||
|
|
6f49738991 | ||
|
|
a59ae8f32e | ||
|
|
6054b8dbd4 | ||
|
|
6f67ba8942 | ||
|
|
d0c5debbab | ||
|
|
4f2e03880b | ||
|
|
5c180cfa90 | ||
|
|
6f09b1d2ee | ||
|
|
b2230fefd8 | ||
|
|
654e6eb0d1 | ||
|
|
a317a4756b | ||
|
|
9b7e6c205f | ||
|
|
602b5ce216 | ||
|
|
8152760b5f | ||
|
|
8c052c99ee | ||
|
|
2667b8116d | ||
|
|
6dffb0305a | ||
|
|
49a9c6b3d2 | ||
|
|
cdf9145ecf | ||
|
|
85f0461b3b | ||
|
|
9f0e9195f7 | ||
|
|
88751d0b08 | ||
|
|
d0e5d910de | ||
|
|
a03504a280 | ||
|
|
d033b2ef0f | ||
|
|
8447f88f61 | ||
|
|
b1b65a657e | ||
|
|
3439e3104e | ||
|
|
288ba20db1 | ||
|
|
020e2eff4e | ||
|
|
1c7369f293 | ||
|
|
0fc1b1bd46 | ||
|
|
d7db37a70f | ||
|
|
6d98bb4f9f | ||
|
|
925cbedc93 | ||
|
|
fda82ee232 | ||
|
|
4b25664c79 | ||
|
|
a27c8a819d | ||
|
|
91acaf4b0b | ||
|
|
41dcf0feb9 | ||
|
|
9960f79920 | ||
|
|
7feeb0b93e | ||
|
|
3639b50b4a | ||
|
|
d855c09cf3 | ||
|
|
d6bfb09863 | ||
|
|
6db276f37a | ||
|
|
6c76c16480 | ||
|
|
11073bd1d2 | ||
|
|
25c9e81b2b | ||
|
|
ffbd9b57c9 | ||
|
|
04899a2b15 | ||
|
|
530d280e33 | ||
|
|
21ddead238 | ||
|
|
7aa5ed09d9 | ||
|
|
75411ce0cc | ||
|
|
9f83d982ec | ||
|
|
3e67b4f88d | ||
|
|
50cfd0d555 | ||
|
|
5756054d38 | ||
|
|
738cb8f128 | ||
|
|
28d1bd07cf | ||
|
|
02625739fe | ||
|
|
f688cd9c5a | ||
|
|
8055027df7 | ||
|
|
3067a8e1a6 | ||
|
|
97114b95a4 | ||
|
|
32fd03a025 | ||
|
|
21bf37dd83 | ||
|
|
5b67d5865a | ||
|
|
df979b4469 | ||
|
|
deb2d7e127 | ||
|
|
fc47319240 | ||
|
|
22cf798d81 | ||
|
|
164be9708b | ||
|
|
6a97524db4 | ||
|
|
c8b1e40f71 | ||
|
|
bcaa2d1ae0 | ||
|
|
8206afefd9 | ||
|
|
646b1b0f46 | ||
|
|
8150ab6c32 | ||
|
|
0b0693a0a2 | ||
|
|
115192c67c | ||
|
|
c2b04d8458 | ||
|
|
db487ab48b | ||
|
|
a95794d3db | ||
|
|
39f84f3b4c | ||
|
|
9f7cf50c56 | ||
|
|
d9a0c72149 | ||
|
|
5ab18bec48 | ||
|
|
2e29ed45d3 | ||
|
|
5ba21f4eb3 | ||
|
|
c26a47b0df | ||
|
|
b1a87b22bb | ||
|
|
07625057f2 | ||
|
|
53c804e233 | ||
|
|
05c7432964 | ||
|
|
4de42d83c2 | ||
|
|
b99485f462 | ||
|
|
20041d7aa9 | ||
|
|
59248032dc | ||
|
|
ceadc34ea9 | ||
|
|
8ab5631446 | ||
|
|
99b5d2b2da | ||
|
|
021e6f3788 | ||
|
|
4e38183e86 | ||
|
|
4eeb23e2b3 | ||
|
|
ef8783b7e3 | ||
|
|
60d7ee614a | ||
|
|
f7a16efc9d | ||
|
|
a01e8bbe98 | ||
|
|
ccf728a1b7 | ||
|
|
f1b4b05d08 | ||
|
|
0c86c89af4 | ||
|
|
d7ac66fb73 | ||
|
|
a6e920fdb0 | ||
|
|
958df58f9d | ||
|
|
e0f102c4d9 | ||
|
|
5a942527b2 | ||
|
|
37a3036934 | ||
|
|
121a7bf8b4 | ||
|
|
a5678c9185 | ||
|
|
2c50b3cf37 | ||
|
|
eee7f54789 | ||
|
|
06eeeead79 | ||
|
|
e8ff7f5321 | ||
|
|
a6e1f26cd4 | ||
|
|
95c43368ae | ||
|
|
754624acf0 | ||
|
|
0b6a17330f | ||
|
|
74b9308883 | ||
|
|
e5f9b1a3a9 | ||
|
|
31d33ccdf0 | ||
|
|
88ec786e39 | ||
|
|
663ef900fc | ||
|
|
7d478a54db | ||
|
|
f3eaaef842 | ||
|
|
d655b65027 | ||
|
|
31c22dc043 | ||
|
|
17127f8b3c | ||
|
|
d7695b40e3 | ||
|
|
fc62890e70 | ||
|
|
f433672140 | ||
|
|
7e1e5b6e6a | ||
|
|
553a42702d | ||
|
|
b133fc9c07 | ||
|
|
b33250dc28 | ||
|
|
a74e5b91a3 | ||
|
|
28886e4241 | ||
|
|
9d3ccfdffc | ||
|
|
a24a7b4da5 | ||
|
|
f7df02f9a3 | ||
|
|
ee450686f3 | ||
|
|
2565755e45 | ||
|
|
d08a92c7bd | ||
|
|
a1ea26d367 | ||
|
|
c17aa0dc54 | ||
|
|
b12b24eadc | ||
|
|
cd14d53707 | ||
|
|
e220413035 | ||
|
|
84ed2327f5 | ||
|
|
b14f301730 | ||
|
|
0654b4b916 | ||
|
|
1f0be382ad | ||
|
|
bb175fda91 | ||
|
|
13998da15a | ||
|
|
57729fd92d | ||
|
|
2c7a71a9c0 | ||
|
|
3e0007fc91 | ||
|
|
b092316385 | ||
|
|
9bcd696580 | ||
|
|
8f89c82d55 | ||
|
|
21871197d7 | ||
|
|
4c35d36146 | ||
|
|
9aca62c26c | ||
|
|
b5cdea98ad | ||
|
|
69fecaf387 | ||
|
|
fd6d25ad86 | ||
|
|
2c3cef1c87 | ||
|
|
89ece26c25 | ||
|
|
2c0b5d0b5e | ||
|
|
a4ae7d17fb | ||
|
|
8a8550184f | ||
|
|
b8b439b713 | ||
|
|
41cd40363a | ||
|
|
d923ebe38d | ||
|
|
29b0423c4e | ||
|
|
88f8dca2c2 | ||
|
|
9027fdc546 | ||
|
|
cbd140340d | ||
|
|
988e01314d | ||
|
|
7ba43a7c6f | ||
|
|
dea59f7e1d | ||
|
|
85dc771460 | ||
|
|
2c5629b81d | ||
|
|
841a582b28 | ||
|
|
c8567a6f65 | ||
|
|
8035be9b1f | ||
|
|
e9b03f4fca | ||
|
|
fd65b9bc23 | ||
|
|
9ebaea840f | ||
|
|
6adc221c10 | ||
|
|
9e63cb9ed0 | ||
|
|
4225518cf3 | ||
|
|
c50adbaac0 | ||
|
|
536dbc0c9a | ||
|
|
4af7acd449 | ||
|
|
53ed52b4b8 | ||
|
|
f1cc7cedce | ||
|
|
ddc4bd1cf6 | ||
|
|
cc36530c73 | ||
|
|
11fa807cfc | ||
|
|
bcdd93e0eb | ||
|
|
579b8c3129 | ||
|
|
d7da51569f | ||
|
|
e8e228d035 | ||
|
|
2579658e15 | ||
|
|
f0cd0134c6 | ||
|
|
abb96996f8 | ||
|
|
bbe6ff2d8f | ||
|
|
db9b39b084 | ||
|
|
849e1e00a3 | ||
|
|
5416c2e8fb | ||
|
|
599a51f4f7 | ||
|
|
17d6eaa2f2 | ||
|
|
2d908639e9 | ||
|
|
c7158418dd | ||
|
|
4d3c9341c1 | ||
|
|
4e508afa2d | ||
|
|
8999ca89b8 | ||
|
|
1adca39cd8 | ||
|
|
204873fa2f | ||
|
|
a5c1de6b1b | ||
|
|
27524ad085 | ||
|
|
27d1921d9c | ||
|
|
70c0e5de90 | ||
|
|
dfb151537b | ||
|
|
500c605fad | ||
|
|
dc9faca3b1 | ||
|
|
aabb0d83e9 | ||
|
|
44579ea6dc | ||
|
|
0f1fcb079f | ||
|
|
84d4769163 | ||
|
|
bf09a35c95 | ||
|
|
6715461a36 | ||
|
|
b4587c5d08 | ||
|
|
88ec63121d | ||
|
|
01d2da2893 | ||
|
|
25d4ea3f91 | ||
|
|
39985840c7 | ||
|
|
b1adc40cfb | ||
|
|
7348bac6ab | ||
|
|
8ab7564d02 | ||
|
|
d096b6e29e | ||
|
|
d88a41f8f1 | ||
|
|
376e9eba80 | ||
|
|
a62c2e11a2 | ||
|
|
a4e5a8c81c | ||
|
|
3e234c46f6 | ||
|
|
7a04b1f8ce | ||
|
|
a30e3d5114 | ||
|
|
1818d06576 | ||
|
|
4e8d1ee24e | ||
|
|
fec376b0dd | ||
|
|
a2512f8a5a | ||
|
|
457e16ea3c | ||
|
|
daf627a6de | ||
|
|
445378667f | ||
|
|
6ae1828449 | ||
|
|
e7b18b7c03 | ||
|
|
9e31d4ef2b | ||
|
|
52aa4d01d5 | ||
|
|
986be957ec | ||
|
|
cf9c60841b | ||
|
|
31bc7f5c2a | ||
|
|
3057741de9 | ||
|
|
acd1103bd0 | ||
|
|
dc7d2cfbca | ||
|
|
b36a78c612 | ||
|
|
985d940db6 | ||
|
|
5e73ca20aa | ||
|
|
438dc10391 | ||
|
|
615ba5d8ef | ||
|
|
02a7cb9fa0 | ||
|
|
9fe2121743 | ||
|
|
0422d6d38e | ||
|
|
9b416c1bbb | ||
|
|
d6899100ac | ||
|
|
0deee48602 | ||
|
|
746a1475b2 | ||
|
|
01ce1fb9e3 | ||
|
|
14f83cbdac | ||
|
|
dbe5891201 | ||
|
|
2a65c3314c | ||
|
|
1c2ff05a6d | ||
|
|
31ae2deeba | ||
|
|
69207e2c57 | ||
|
|
138c5bcc08 | ||
|
|
a923e0a23a | ||
|
|
f521a30b22 | ||
|
|
d4451f6afb | ||
|
|
a3275423a4 | ||
|
|
b37c3d000c | ||
|
|
6031020e37 | ||
|
|
c424dfc293 | ||
|
|
3a28e52e98 | ||
|
|
e371908b54 | ||
|
|
7c99da155c | ||
|
|
629e72385b | ||
|
|
0a708fff24 | ||
|
|
6e150ea6d0 | ||
|
|
cb8dcb97ea | ||
|
|
2d5dc93b3d | ||
|
|
4145d35e3c | ||
|
|
34c6c45bd6 | ||
|
|
e9def84ce7 | ||
|
|
836e02a166 | ||
|
|
b558e61f63 | ||
|
|
65ab69543b | ||
|
|
1d26aa2e93 | ||
|
|
a548d4553e | ||
|
|
dd1b39f435 | ||
|
|
94d6e713e9 | ||
|
|
47c37e4876 | ||
|
|
737585a32a | ||
|
|
a4688021bf | ||
|
|
7df6eb9211 | ||
|
|
82a3f2626f | ||
|
|
7fa69572c0 | ||
|
|
3ab4f237e5 | ||
|
|
8cbf3f36e2 | ||
|
|
0594ce1017 | ||
|
|
ff509ff39f | ||
|
|
785d65436c | ||
|
|
64be81b7b3 | ||
|
|
45479b5731 | ||
|
|
e0a3337c22 | ||
|
|
812238060b | ||
|
|
14b0d56197 | ||
|
|
6c8533f1d2 | ||
|
|
2c2697390d | ||
|
|
7621f05d3f | ||
|
|
10ebd7211f | ||
|
|
42a391f0fb | ||
|
|
97c7ac0f4f | ||
|
|
8f1b32f2b6 | ||
|
|
c241a5dcef | ||
|
|
44dab27fdc | ||
|
|
a44fd22a99 | ||
|
|
8a11a7d444 | ||
|
|
1d54491809 | ||
|
|
ad9f4d9cf6 | ||
|
|
e1638a7ade | ||
|
|
f91bfee33e | ||
|
|
d7a7f570ed | ||
|
|
7dea929788 | ||
|
|
026d1fc33d | ||
|
|
7242eedbf4 | ||
|
|
04c0dc7a47 | ||
|
|
48a53121ba | ||
|
|
0ba8c70ce1 | ||
|
|
3d12a03909 | ||
|
|
c169659611 | ||
|
|
e12f1a7ee5 | ||
|
|
ef25efffa2 | ||
|
|
19532440b4 | ||
|
|
9096e413c3 | ||
|
|
9d5e9fa6c4 | ||
|
|
08dde46778 | ||
|
|
513f1f7826 | ||
|
|
e3382f6bb5 | ||
|
|
f0339022c1 | ||
|
|
d8da2cf17c | ||
|
|
205b40bd28 | ||
|
|
18fe6e9339 | ||
|
|
2196c34c52 | ||
|
|
466c2e1efd | ||
|
|
7e26d848ab | ||
|
|
ed95ef245c | ||
|
|
6d6ef99e66 | ||
|
|
a8e2a1ba45 | ||
|
|
6269bacfc3 | ||
|
|
c0effc9f5b | ||
|
|
df0845e916 | ||
|
|
7440e9c809 | ||
|
|
7d4029c2a4 | ||
|
|
0ca6c9e6eb | ||
|
|
6e49d27057 | ||
|
|
5203b7f53e | ||
|
|
5889179c54 | ||
|
|
38e18fdfd3 | ||
|
|
4753958f92 | ||
|
|
73d6cc0f26 | ||
|
|
317ed90bac | ||
|
|
951df8155c | ||
|
|
a58fab8d6e | ||
|
|
a3c8296135 | ||
|
|
c95ace41aa | ||
|
|
3da428e0e4 | ||
|
|
133a9de98f |
+3
-1
@@ -4,6 +4,8 @@
|
||||
# Allow necessary files
|
||||
!astrai/
|
||||
!scripts/
|
||||
!assets/
|
||||
!docs/
|
||||
!csrc/
|
||||
!setup.py
|
||||
!pyproject.toml
|
||||
!README.md
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
name: Bug report
|
||||
about: Create a report to help us improve
|
||||
title: "[BUG]"
|
||||
labels: enhancement
|
||||
labels: bug
|
||||
assignees: ''
|
||||
|
||||
---
|
||||
|
||||
@@ -16,9 +16,9 @@ Please delete options that are not relevant.
|
||||
Please describe the tests that you ran to verify your changes. Provide instructions so we can reproduce.
|
||||
|
||||
## Checklist:
|
||||
- [ ] My code follows the style guidelines of this project (run `ruff format .` and `ruff check --fix .`)
|
||||
- [ ] My code follows the style guidelines of this project (run `ruff format .` and `ruff check . --select I`)
|
||||
- [ ] I have performed a self-review of my own code
|
||||
- [ ] I have commented my code, particularly in hard-to-understand areas
|
||||
- [ ] Code is self-documenting (no unnecessary comments)
|
||||
- [ ] I have made corresponding changes to the documentation
|
||||
- [ ] My changes generate no new warnings
|
||||
- [ ] I have added tests that prove my fix is effective or that my feature works
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
name: Release
|
||||
|
||||
on:
|
||||
push:
|
||||
tags:
|
||||
- "v*"
|
||||
|
||||
jobs:
|
||||
build-pure:
|
||||
name: Build pure-Python wheel
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Build wheel (no CUDA)
|
||||
run: |
|
||||
pip wheel . --no-deps -w dist/
|
||||
|
||||
- uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: pure-wheel
|
||||
path: dist/*.whl
|
||||
if-no-files-found: error
|
||||
|
||||
build-cuda-linux:
|
||||
name: Build CUDA wheel (Linux, ${{ matrix.cuda_tag }})
|
||||
runs-on: ubuntu-latest
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
include:
|
||||
- cuda_tag: "cu128"
|
||||
cuda_ver: "12.8.0"
|
||||
- cuda_tag: "cu130"
|
||||
cuda_ver: "13.0.0"
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Install torch (${{ matrix.cuda_tag }})
|
||||
run: |
|
||||
pip install torch --index-url https://download.pytorch.org/whl/${{ matrix.cuda_tag }}
|
||||
|
||||
- name: Setup CUDA (${{ matrix.cuda_ver }})
|
||||
uses: Jimver/cuda-toolkit@v0.2.35
|
||||
with:
|
||||
cuda: "${{ matrix.cuda_ver }}"
|
||||
|
||||
- name: Build wheel (with CUDA kernels)
|
||||
run: |
|
||||
CSRC_KERNELS=true pip wheel . --no-deps --no-build-isolation -w dist/
|
||||
for f in dist/*.whl; do
|
||||
mv "$f" "dist/$(basename "$f" .whl)+${{ matrix.cuda_tag }}.whl"
|
||||
done
|
||||
|
||||
- uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: cuda-wheel-linux-${{ matrix.cuda_tag }}
|
||||
path: dist/*.whl
|
||||
if-no-files-found: error
|
||||
|
||||
release:
|
||||
name: Attach wheels to release
|
||||
needs: [build-pure, build-cuda-linux]
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: write
|
||||
steps:
|
||||
- name: Download pure-Python wheel
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: pure-wheel
|
||||
path: release-assets/pure
|
||||
|
||||
- name: Download CUDA wheels (all variants)
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
pattern: cuda-wheel-linux-*
|
||||
merge-multiple: true
|
||||
path: release-assets/cuda
|
||||
|
||||
- name: Verify release assets
|
||||
shell: bash
|
||||
run: |
|
||||
set -euo pipefail
|
||||
pure_wheels=(release-assets/pure/*.whl)
|
||||
cuda_wheels=(release-assets/cuda/*.whl)
|
||||
test "${#pure_wheels[@]}" -eq 1
|
||||
test "${#cuda_wheels[@]}" -ge 1
|
||||
|
||||
- name: Create release & upload assets
|
||||
uses: softprops/action-gh-release@v2
|
||||
with:
|
||||
files: |
|
||||
release-assets/pure/*.whl
|
||||
release-assets/cuda/*.whl
|
||||
tag_name: ${{ github.ref_name }}
|
||||
generate_release_notes: true
|
||||
+18
-4
@@ -5,8 +5,17 @@
|
||||
!*/
|
||||
|
||||
# Allow specific file types and root files
|
||||
!*.py
|
||||
!*.sh
|
||||
!astrai/**/*.py
|
||||
!scripts/**/*.py
|
||||
!tests/**/*.py
|
||||
!csrc/**/*.py
|
||||
!csrc/CMakeLists.txt
|
||||
|
||||
!csrc/**/*.cu
|
||||
!csrc/**/*.h
|
||||
!csrc/**/*.cuh
|
||||
|
||||
!scripts/**/*.sh
|
||||
|
||||
# Allow GitHub files
|
||||
!/.github/**
|
||||
@@ -16,8 +25,13 @@
|
||||
!/.dockerignore
|
||||
!/Dockerfile
|
||||
!/docker-compose.yml
|
||||
!/assets/**
|
||||
!/docs/**
|
||||
!/CONTRIBUTING.md
|
||||
!/LICENSE
|
||||
!/pyproject.toml
|
||||
!/README.md
|
||||
!/README.md
|
||||
# Allow extension modules (only source .py)
|
||||
!/astrai/extension/**/*.py
|
||||
|
||||
# Allow build files
|
||||
!/setup.py
|
||||
|
||||
+82
-48
@@ -1,68 +1,102 @@
|
||||
# Contributing to AstrAI
|
||||
|
||||
Thank you for your interest in contributing to AstrAI! This document provides guidelines and steps for contributing.
|
||||
Thank you for your interest in contributing! This document provides step-by-step guidelines.
|
||||
|
||||
## How to Contribute
|
||||
## Quick Start
|
||||
|
||||
### Reporting Issues
|
||||
If you encounter a bug or have a feature request, please open an issue on GitHub. Include as much detail as possible:
|
||||
- A clear description of the problem or request.
|
||||
- Steps to reproduce (for bugs).
|
||||
- Your environment (Python version, OS, etc.).
|
||||
```bash
|
||||
git clone https://github.com/ViperEkura/AstrAI.git
|
||||
cd AstrAI
|
||||
pip install -e ".[dev]" # install with dev dependencies (pytest, ruff)
|
||||
```
|
||||
|
||||
### Submitting Changes
|
||||
1. **Fork** the repository.
|
||||
2. **Clone** your fork:
|
||||
```bash
|
||||
git clone https://github.com/your-username/AstrAI.git
|
||||
cd AstrAI
|
||||
```
|
||||
3. **Create a feature branch**:
|
||||
```bash
|
||||
git checkout -b feature/your-feature-name
|
||||
```
|
||||
4. **Make your changes**. Follow the code style guidelines below.
|
||||
5. **Commit your changes** with a descriptive commit message:
|
||||
```bash
|
||||
git commit -m "Add: brief description of the change"
|
||||
```
|
||||
6. **Push** to your fork:
|
||||
```bash
|
||||
git push origin feature/your-feature-name
|
||||
```
|
||||
7. **Open a Pull Request** (PR) against the `main` branch of the upstream repository.
|
||||
## Before You Commit
|
||||
|
||||
## Code Style
|
||||
Run the following checks **in order** — CI will reject if any fail.
|
||||
|
||||
AstrAI uses [Ruff](https://docs.astral.sh/ruff/) for code formatting and linting. Please ensure your code is formatted before submitting.
|
||||
### 1. Format
|
||||
|
||||
- Run Ruff to format and lint:
|
||||
```bash
|
||||
ruff format .
|
||||
ruff check --fix .
|
||||
```
|
||||
- The project uses **double quotes** for strings and **4‑space indentation** (as configured in `pyproject.toml`).
|
||||
```bash
|
||||
ruff format .
|
||||
```
|
||||
|
||||
## Testing
|
||||
### 2. Import sorting
|
||||
|
||||
If you add or modify functionality, please include appropriate tests.
|
||||
```bash
|
||||
ruff check . --select I
|
||||
```
|
||||
|
||||
- Run the test suite with:
|
||||
```bash
|
||||
pytest
|
||||
```
|
||||
- Ensure all tests pass before submitting your PR.
|
||||
If this fails, **manually fix** import ordering (ruff does not auto-fix in this project's CI):
|
||||
|
||||
```bash
|
||||
ruff check . --select I --fix .
|
||||
ruff format . # re-format after fix
|
||||
```
|
||||
|
||||
### 3. Run tests
|
||||
|
||||
```bash
|
||||
python -u -m pytest tests/ -v
|
||||
```
|
||||
|
||||
> Failed tests may leave orphan tempdirs under `%TEMP%`. Clean them manually if needed.
|
||||
|
||||
### 4. (Optional) Full pre-commit check script
|
||||
|
||||
If you have Git Bash available:
|
||||
|
||||
```bash
|
||||
bash scripts/pre_commit.sh
|
||||
```
|
||||
|
||||
The script installs development dependencies by default, then runs the format
|
||||
check, import sort check, and tests. If dependencies are already installed, use:
|
||||
|
||||
```bash
|
||||
bash scripts/pre_commit.sh --skip-deps
|
||||
```
|
||||
|
||||
## Commit Style
|
||||
|
||||
```
|
||||
type: short description (~50 chars)
|
||||
|
||||
- bullet point body (each ~60 chars)
|
||||
```
|
||||
|
||||
- **Type** must be one of: `fix`, `feat`, `chore`, `docs`, `refactor`, `perf`, `test`, `style`, `ci`, `build`, `revert`.
|
||||
- **Subject line** ends with no period.
|
||||
- **Body** uses bullet points starting with `-`.
|
||||
- No `(scope)` parentheses.
|
||||
|
||||
## Common Issues
|
||||
|
||||
| Problem | Cause | Fix |
|
||||
|---------|-------|-----|
|
||||
| `ruff check --select I` fails | Wrong import order | `ruff check . --select I --fix .` then `ruff format .` |
|
||||
| `ruff format` changed many files | Not formatted before commit | Review diff carefully before staging |
|
||||
| Pre-commit check script fails | Dependency install, tests, or lint failed | Fix the failing step; use `--skip-deps` only when dependencies are already installed |
|
||||
| Tests fail with tempdir left | Test crash | Clean `%TEMP%` manually |
|
||||
|
||||
## Submitting Changes
|
||||
|
||||
1. Fork the repo.
|
||||
2. Create a feature branch: `git checkout -b feat/my-feature`
|
||||
3. Make changes following the steps above.
|
||||
4. Commit with the commit style above.
|
||||
5. Push: `git push origin feat/my-feature`
|
||||
6. Open a Pull Request against `main`.
|
||||
|
||||
## Code Review
|
||||
|
||||
All submissions will be reviewed. We may request changes or discuss alternatives. Please be responsive to feedback.
|
||||
- All PRs are reviewed. We may request changes.
|
||||
- CI runs `ruff format --check .` then `ruff check . --select I` (no `--fix` in CI).
|
||||
- Ensure all tests pass.
|
||||
|
||||
## License
|
||||
|
||||
By contributing, you agree that your contributions will be licensed under the same [GPL-3.0 License](LICENSE) that covers the project.
|
||||
By contributing, you agree that your contributions will be licensed under the [Apache-2.0 License](LICENSE).
|
||||
|
||||
---
|
||||
|
||||
If you have any questions, feel free to ask in the [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions) or open an issue.
|
||||
|
||||
Happy contributing!
|
||||
Questions? Ask in [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions) or open an issue.
|
||||
|
||||
+24
-8
@@ -1,7 +1,15 @@
|
||||
# AstrAI Dockerfile - Multi-stage Build (Optimized)
|
||||
#
|
||||
# CUDA version selection:
|
||||
# docker build -t astrai .
|
||||
# docker build -t astrai --build-arg CUDA_TAG=cu128 .
|
||||
# docker build -t astrai --build-arg CUDA_TAG=cu130 .
|
||||
# Default: cu128
|
||||
|
||||
# Build stage - use base image with minimal build tools
|
||||
FROM nvidia/cuda:12.6.0-base-ubuntu24.04 AS builder
|
||||
FROM ubuntu:24.04 AS builder
|
||||
|
||||
ARG CUDA_TAG=cu128
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
@@ -18,21 +26,24 @@ RUN apt-get update && DEBIAN_FRONTEND=noninteractive apt-get install -y --no-ins
|
||||
RUN python3.12 -m venv --copies /opt/venv
|
||||
ENV PATH="/opt/venv/bin:$PATH"
|
||||
|
||||
# Copy source code and install dependencies
|
||||
# Copy source code and install (deps read from pyproject.toml)
|
||||
COPY astrai/ ./astrai/
|
||||
COPY csrc/ ./csrc/
|
||||
COPY setup.py .
|
||||
COPY pyproject.toml .
|
||||
RUN pip install --no-cache-dir --upgrade pip \
|
||||
&& pip install --no-cache-dir . \
|
||||
--extra-index-url https://download.pytorch.org/whl/cu126
|
||||
--extra-index-url "https://download.pytorch.org/whl/${CUDA_TAG}"
|
||||
|
||||
# Production stage
|
||||
FROM nvidia/cuda:12.6.0-base-ubuntu24.04 AS production
|
||||
FROM ubuntu:24.04 AS production
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# Install Python 3.12 runtime
|
||||
# Install Python 3.12 runtime and healthcheck dependency
|
||||
RUN apt-get update && DEBIAN_FRONTEND=noninteractive apt-get install -y --no-install-recommends \
|
||||
python3.12 \
|
||||
curl \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Copy virtual environment from builder
|
||||
@@ -42,12 +53,17 @@ ENV PATH="/opt/venv/bin:$PATH"
|
||||
# Copy application code
|
||||
COPY astrai/ ./astrai/
|
||||
COPY scripts/ ./scripts/
|
||||
COPY assets/ ./assets/
|
||||
COPY docs/ ./docs/
|
||||
COPY pyproject.toml .
|
||||
COPY README.md .
|
||||
|
||||
# Create non-root user
|
||||
RUN useradd -m astrai && chown -R astrai:astrai /app
|
||||
# Create non-root user matching the host uid/gid (passed via build args)
|
||||
ARG USER_UID=1000
|
||||
ARG USER_GID=1000
|
||||
RUN groupadd -g "${USER_GID}" astrai \
|
||||
&& useradd -m -u "${USER_UID}" -g astrai astrai \
|
||||
&& chown -R astrai:astrai /app
|
||||
ENV HOME=/home/astrai
|
||||
USER astrai
|
||||
|
||||
ENV PYTHONUNBUFFERED=1 \
|
||||
|
||||
@@ -1,674 +1,201 @@
|
||||
GNU GENERAL PUBLIC LICENSE
|
||||
Version 3, 29 June 2007
|
||||
|
||||
Copyright (C) 2007 Free Software Foundation, Inc. <https://fsf.org/>
|
||||
Everyone is permitted to copy and distribute verbatim copies
|
||||
of this license document, but changing it is not allowed.
|
||||
|
||||
Preamble
|
||||
|
||||
The GNU General Public License is a free, copyleft license for
|
||||
software and other kinds of works.
|
||||
|
||||
The licenses for most software and other practical works are designed
|
||||
to take away your freedom to share and change the works. By contrast,
|
||||
the GNU General Public License is intended to guarantee your freedom to
|
||||
share and change all versions of a program--to make sure it remains free
|
||||
software for all its users. We, the Free Software Foundation, use the
|
||||
GNU General Public License for most of our software; it applies also to
|
||||
any other work released this way by its authors. You can apply it to
|
||||
your programs, too.
|
||||
|
||||
When we speak of free software, we are referring to freedom, not
|
||||
price. Our General Public Licenses are designed to make sure that you
|
||||
have the freedom to distribute copies of free software (and charge for
|
||||
them if you wish), that you receive source code or can get it if you
|
||||
want it, that you can change the software or use pieces of it in new
|
||||
free programs, and that you know you can do these things.
|
||||
|
||||
To protect your rights, we need to prevent others from denying you
|
||||
these rights or asking you to surrender the rights. Therefore, you have
|
||||
certain responsibilities if you distribute copies of the software, or if
|
||||
you modify it: responsibilities to respect the freedom of others.
|
||||
|
||||
For example, if you distribute copies of such a program, whether
|
||||
gratis or for a fee, you must pass on to the recipients the same
|
||||
freedoms that you received. You must make sure that they, too, receive
|
||||
or can get the source code. And you must show them these terms so they
|
||||
know their rights.
|
||||
|
||||
Developers that use the GNU GPL protect your rights with two steps:
|
||||
(1) assert copyright on the software, and (2) offer you this License
|
||||
giving you legal permission to copy, distribute and/or modify it.
|
||||
|
||||
For the developers' and authors' protection, the GPL clearly explains
|
||||
that there is no warranty for this free software. For both users' and
|
||||
authors' sake, the GPL requires that modified versions be marked as
|
||||
changed, so that their problems will not be attributed erroneously to
|
||||
authors of previous versions.
|
||||
|
||||
Some devices are designed to deny users access to install or run
|
||||
modified versions of the software inside them, although the manufacturer
|
||||
can do so. This is fundamentally incompatible with the aim of
|
||||
protecting users' freedom to change the software. The systematic
|
||||
pattern of such abuse occurs in the area of products for individuals to
|
||||
use, which is precisely where it is most unacceptable. Therefore, we
|
||||
have designed this version of the GPL to prohibit the practice for those
|
||||
products. If such problems arise substantially in other domains, we
|
||||
stand ready to extend this provision to those domains in future versions
|
||||
of the GPL, as needed to protect the freedom of users.
|
||||
|
||||
Finally, every program is threatened constantly by software patents.
|
||||
States should not allow patents to restrict development and use of
|
||||
software on general-purpose computers, but in those that do, we wish to
|
||||
avoid the special danger that patents applied to a free program could
|
||||
make it effectively proprietary. To prevent this, the GPL assures that
|
||||
patents cannot be used to render the program non-free.
|
||||
|
||||
The precise terms and conditions for copying, distribution and
|
||||
modification follow.
|
||||
|
||||
TERMS AND CONDITIONS
|
||||
|
||||
0. Definitions.
|
||||
|
||||
"This License" refers to version 3 of the GNU General Public License.
|
||||
|
||||
"Copyright" also means copyright-like laws that apply to other kinds of
|
||||
works, such as semiconductor masks.
|
||||
|
||||
"The Program" refers to any copyrightable work licensed under this
|
||||
License. Each licensee is addressed as "you". "Licensees" and
|
||||
"recipients" may be individuals or organizations.
|
||||
|
||||
To "modify" a work means to copy from or adapt all or part of the work
|
||||
in a fashion requiring copyright permission, other than the making of an
|
||||
exact copy. The resulting work is called a "modified version" of the
|
||||
earlier work or a work "based on" the earlier work.
|
||||
|
||||
A "covered work" means either the unmodified Program or a work based
|
||||
on the Program.
|
||||
|
||||
To "propagate" a work means to do anything with it that, without
|
||||
permission, would make you directly or secondarily liable for
|
||||
infringement under applicable copyright law, except executing it on a
|
||||
computer or modifying a private copy. Propagation includes copying,
|
||||
distribution (with or without modification), making available to the
|
||||
public, and in some countries other activities as well.
|
||||
|
||||
To "convey" a work means any kind of propagation that enables other
|
||||
parties to make or receive copies. Mere interaction with a user through
|
||||
a computer network, with no transfer of a copy, is not conveying.
|
||||
|
||||
An interactive user interface displays "Appropriate Legal Notices"
|
||||
to the extent that it includes a convenient and prominently visible
|
||||
feature that (1) displays an appropriate copyright notice, and (2)
|
||||
tells the user that there is no warranty for the work (except to the
|
||||
extent that warranties are provided), that licensees may convey the
|
||||
work under this License, and how to view a copy of this License. If
|
||||
the interface presents a list of user commands or options, such as a
|
||||
menu, a prominent item in the list meets this criterion.
|
||||
|
||||
1. Source Code.
|
||||
|
||||
The "source code" for a work means the preferred form of the work
|
||||
for making modifications to it. "Object code" means any non-source
|
||||
form of a work.
|
||||
|
||||
A "Standard Interface" means an interface that either is an official
|
||||
standard defined by a recognized standards body, or, in the case of
|
||||
interfaces specified for a particular programming language, one that
|
||||
is widely used among developers working in that language.
|
||||
|
||||
The "System Libraries" of an executable work include anything, other
|
||||
than the work as a whole, that (a) is included in the normal form of
|
||||
packaging a Major Component, but which is not part of that Major
|
||||
Component, and (b) serves only to enable use of the work with that
|
||||
Major Component, or to implement a Standard Interface for which an
|
||||
implementation is available to the public in source code form. A
|
||||
"Major Component", in this context, means a major essential component
|
||||
(kernel, window system, and so on) of the specific operating system
|
||||
(if any) on which the executable work runs, or a compiler used to
|
||||
produce the work, or an object code interpreter used to run it.
|
||||
|
||||
The "Corresponding Source" for a work in object code form means all
|
||||
the source code needed to generate, install, and (for an executable
|
||||
work) run the object code and to modify the work, including scripts to
|
||||
control those activities. However, it does not include the work's
|
||||
System Libraries, or general-purpose tools or generally available free
|
||||
programs which are used unmodified in performing those activities but
|
||||
which are not part of the work. For example, Corresponding Source
|
||||
includes interface definition files associated with source files for
|
||||
the work, and the source code for shared libraries and dynamically
|
||||
linked subprograms that the work is specifically designed to require,
|
||||
such as by intimate data communication or control flow between those
|
||||
subprograms and other parts of the work.
|
||||
|
||||
The Corresponding Source need not include anything that users
|
||||
can regenerate automatically from other parts of the Corresponding
|
||||
Source.
|
||||
|
||||
The Corresponding Source for a work in source code form is that
|
||||
same work.
|
||||
|
||||
2. Basic Permissions.
|
||||
|
||||
All rights granted under this License are granted for the term of
|
||||
copyright on the Program, and are irrevocable provided the stated
|
||||
conditions are met. This License explicitly affirms your unlimited
|
||||
permission to run the unmodified Program. The output from running a
|
||||
covered work is covered by this License only if the output, given its
|
||||
content, constitutes a covered work. This License acknowledges your
|
||||
rights of fair use or other equivalent, as provided by copyright law.
|
||||
|
||||
You may make, run and propagate covered works that you do not
|
||||
convey, without conditions so long as your license otherwise remains
|
||||
in force. You may convey covered works to others for the sole purpose
|
||||
of having them make modifications exclusively for you, or provide you
|
||||
with facilities for running those works, provided that you comply with
|
||||
the terms of this License in conveying all material for which you do
|
||||
not control copyright. Those thus making or running the covered works
|
||||
for you must do so exclusively on your behalf, under your direction
|
||||
and control, on terms that prohibit them from making any copies of
|
||||
your copyrighted material outside their relationship with you.
|
||||
|
||||
Conveying under any other circumstances is permitted solely under
|
||||
the conditions stated below. Sublicensing is not allowed; section 10
|
||||
makes it unnecessary.
|
||||
|
||||
3. Protecting Users' Legal Rights From Anti-Circumvention Law.
|
||||
|
||||
No covered work shall be deemed part of an effective technological
|
||||
measure under any applicable law fulfilling obligations under article
|
||||
11 of the WIPO copyright treaty adopted on 20 December 1996, or
|
||||
similar laws prohibiting or restricting circumvention of such
|
||||
measures.
|
||||
|
||||
When you convey a covered work, you waive any legal power to forbid
|
||||
circumvention of technological measures to the extent such circumvention
|
||||
is effected by exercising rights under this License with respect to
|
||||
the covered work, and you disclaim any intention to limit operation or
|
||||
modification of the work as a means of enforcing, against the work's
|
||||
users, your or third parties' legal rights to forbid circumvention of
|
||||
technological measures.
|
||||
|
||||
4. Conveying Verbatim Copies.
|
||||
|
||||
You may convey verbatim copies of the Program's source code as you
|
||||
receive it, in any medium, provided that you conspicuously and
|
||||
appropriately publish on each copy an appropriate copyright notice;
|
||||
keep intact all notices stating that this License and any
|
||||
non-permissive terms added in accord with section 7 apply to the code;
|
||||
keep intact all notices of the absence of any warranty; and give all
|
||||
recipients a copy of this License along with the Program.
|
||||
|
||||
You may charge any price or no price for each copy that you convey,
|
||||
and you may offer support or warranty protection for a fee.
|
||||
|
||||
5. Conveying Modified Source Versions.
|
||||
|
||||
You may convey a work based on the Program, or the modifications to
|
||||
produce it from the Program, in the form of source code under the
|
||||
terms of section 4, provided that you also meet all of these conditions:
|
||||
|
||||
a) The work must carry prominent notices stating that you modified
|
||||
it, and giving a relevant date.
|
||||
|
||||
b) The work must carry prominent notices stating that it is
|
||||
released under this License and any conditions added under section
|
||||
7. This requirement modifies the requirement in section 4 to
|
||||
"keep intact all notices".
|
||||
|
||||
c) You must license the entire work, as a whole, under this
|
||||
License to anyone who comes into possession of a copy. This
|
||||
License will therefore apply, along with any applicable section 7
|
||||
additional terms, to the whole of the work, and all its parts,
|
||||
regardless of how they are packaged. This License gives no
|
||||
permission to license the work in any other way, but it does not
|
||||
invalidate such permission if you have separately received it.
|
||||
|
||||
d) If the work has interactive user interfaces, each must display
|
||||
Appropriate Legal Notices; however, if the Program has interactive
|
||||
interfaces that do not display Appropriate Legal Notices, your
|
||||
work need not make them do so.
|
||||
|
||||
A compilation of a covered work with other separate and independent
|
||||
works, which are not by their nature extensions of the covered work,
|
||||
and which are not combined with it such as to form a larger program,
|
||||
in or on a volume of a storage or distribution medium, is called an
|
||||
"aggregate" if the compilation and its resulting copyright are not
|
||||
used to limit the access or legal rights of the compilation's users
|
||||
beyond what the individual works permit. Inclusion of a covered work
|
||||
in an aggregate does not cause this License to apply to the other
|
||||
parts of the aggregate.
|
||||
|
||||
6. Conveying Non-Source Forms.
|
||||
|
||||
You may convey a covered work in object code form under the terms
|
||||
of sections 4 and 5, provided that you also convey the
|
||||
machine-readable Corresponding Source under the terms of this License,
|
||||
in one of these ways:
|
||||
|
||||
a) Convey the object code in, or embodied in, a physical product
|
||||
(including a physical distribution medium), accompanied by the
|
||||
Corresponding Source fixed on a durable physical medium
|
||||
customarily used for software interchange.
|
||||
|
||||
b) Convey the object code in, or embodied in, a physical product
|
||||
(including a physical distribution medium), accompanied by a
|
||||
written offer, valid for at least three years and valid for as
|
||||
long as you offer spare parts or customer support for that product
|
||||
model, to give anyone who possesses the object code either (1) a
|
||||
copy of the Corresponding Source for all the software in the
|
||||
product that is covered by this License, on a durable physical
|
||||
medium customarily used for software interchange, for a price no
|
||||
more than your reasonable cost of physically performing this
|
||||
conveying of source, or (2) access to copy the
|
||||
Corresponding Source from a network server at no charge.
|
||||
|
||||
c) Convey individual copies of the object code with a copy of the
|
||||
written offer to provide the Corresponding Source. This
|
||||
alternative is allowed only occasionally and noncommercially, and
|
||||
only if you received the object code with such an offer, in accord
|
||||
with subsection 6b.
|
||||
|
||||
d) Convey the object code by offering access from a designated
|
||||
place (gratis or for a charge), and offer equivalent access to the
|
||||
Corresponding Source in the same way through the same place at no
|
||||
further charge. You need not require recipients to copy the
|
||||
Corresponding Source along with the object code. If the place to
|
||||
copy the object code is a network server, the Corresponding Source
|
||||
may be on a different server (operated by you or a third party)
|
||||
that supports equivalent copying facilities, provided you maintain
|
||||
clear directions next to the object code saying where to find the
|
||||
Corresponding Source. Regardless of what server hosts the
|
||||
Corresponding Source, you remain obligated to ensure that it is
|
||||
available for as long as needed to satisfy these requirements.
|
||||
|
||||
e) Convey the object code using peer-to-peer transmission, provided
|
||||
you inform other peers where the object code and Corresponding
|
||||
Source of the work are being offered to the general public at no
|
||||
charge under subsection 6d.
|
||||
|
||||
A separable portion of the object code, whose source code is excluded
|
||||
from the Corresponding Source as a System Library, need not be
|
||||
included in conveying the object code work.
|
||||
|
||||
A "User Product" is either (1) a "consumer product", which means any
|
||||
tangible personal property which is normally used for personal, family,
|
||||
or household purposes, or (2) anything designed or sold for incorporation
|
||||
into a dwelling. In determining whether a product is a consumer product,
|
||||
doubtful cases shall be resolved in favor of coverage. For a particular
|
||||
product received by a particular user, "normally used" refers to a
|
||||
typical or common use of that class of product, regardless of the status
|
||||
of the particular user or of the way in which the particular user
|
||||
actually uses, or expects or is expected to use, the product. A product
|
||||
is a consumer product regardless of whether the product has substantial
|
||||
commercial, industrial or non-consumer uses, unless such uses represent
|
||||
the only significant mode of use of the product.
|
||||
|
||||
"Installation Information" for a User Product means any methods,
|
||||
procedures, authorization keys, or other information required to install
|
||||
and execute modified versions of a covered work in that User Product from
|
||||
a modified version of its Corresponding Source. The information must
|
||||
suffice to ensure that the continued functioning of the modified object
|
||||
code is in no case prevented or interfered with solely because
|
||||
modification has been made.
|
||||
|
||||
If you convey an object code work under this section in, or with, or
|
||||
specifically for use in, a User Product, and the conveying occurs as
|
||||
part of a transaction in which the right of possession and use of the
|
||||
User Product is transferred to the recipient in perpetuity or for a
|
||||
fixed term (regardless of how the transaction is characterized), the
|
||||
Corresponding Source conveyed under this section must be accompanied
|
||||
by the Installation Information. But this requirement does not apply
|
||||
if neither you nor any third party retains the ability to install
|
||||
modified object code on the User Product (for example, the work has
|
||||
been installed in ROM).
|
||||
|
||||
The requirement to provide Installation Information does not include a
|
||||
requirement to continue to provide support service, warranty, or updates
|
||||
for a work that has been modified or installed by the recipient, or for
|
||||
the User Product in which it has been modified or installed. Access to a
|
||||
network may be denied when the modification itself materially and
|
||||
adversely affects the operation of the network or violates the rules and
|
||||
protocols for communication across the network.
|
||||
|
||||
Corresponding Source conveyed, and Installation Information provided,
|
||||
in accord with this section must be in a format that is publicly
|
||||
documented (and with an implementation available to the public in
|
||||
source code form), and must require no special password or key for
|
||||
unpacking, reading or copying.
|
||||
|
||||
7. Additional Terms.
|
||||
|
||||
"Additional permissions" are terms that supplement the terms of this
|
||||
License by making exceptions from one or more of its conditions.
|
||||
Additional permissions that are applicable to the entire Program shall
|
||||
be treated as though they were included in this License, to the extent
|
||||
that they are valid under applicable law. If additional permissions
|
||||
apply only to part of the Program, that part may be used separately
|
||||
under those permissions, but the entire Program remains governed by
|
||||
this License without regard to the additional permissions.
|
||||
|
||||
When you convey a copy of a covered work, you may at your option
|
||||
remove any additional permissions from that copy, or from any part of
|
||||
it. (Additional permissions may be written to require their own
|
||||
removal in certain cases when you modify the work.) You may place
|
||||
additional permissions on material, added by you to a covered work,
|
||||
for which you have or can give appropriate copyright permission.
|
||||
|
||||
Notwithstanding any other provision of this License, for material you
|
||||
add to a covered work, you may (if authorized by the copyright holders of
|
||||
that material) supplement the terms of this License with terms:
|
||||
|
||||
a) Disclaiming warranty or limiting liability differently from the
|
||||
terms of sections 15 and 16 of this License; or
|
||||
|
||||
b) Requiring preservation of specified reasonable legal notices or
|
||||
author attributions in that material or in the Appropriate Legal
|
||||
Notices displayed by works containing it; or
|
||||
|
||||
c) Prohibiting misrepresentation of the origin of that material, or
|
||||
requiring that modified versions of such material be marked in
|
||||
reasonable ways as different from the original version; or
|
||||
|
||||
d) Limiting the use for publicity purposes of names of licensors or
|
||||
authors of the material; or
|
||||
|
||||
e) Declining to grant rights under trademark law for use of some
|
||||
trade names, trademarks, or service marks; or
|
||||
|
||||
f) Requiring indemnification of licensors and authors of that
|
||||
material by anyone who conveys the material (or modified versions of
|
||||
it) with contractual assumptions of liability to the recipient, for
|
||||
any liability that these contractual assumptions directly impose on
|
||||
those licensors and authors.
|
||||
|
||||
All other non-permissive additional terms are considered "further
|
||||
restrictions" within the meaning of section 10. If the Program as you
|
||||
received it, or any part of it, contains a notice stating that it is
|
||||
governed by this License along with a term that is a further
|
||||
restriction, you may remove that term. If a license document contains
|
||||
a further restriction but permits relicensing or conveying under this
|
||||
License, you may add to a covered work material governed by the terms
|
||||
of that license document, provided that the further restriction does
|
||||
not survive such relicensing or conveying.
|
||||
|
||||
If you add terms to a covered work in accord with this section, you
|
||||
must place, in the relevant source files, a statement of the
|
||||
additional terms that apply to those files, or a notice indicating
|
||||
where to find the applicable terms.
|
||||
|
||||
Additional terms, permissive or non-permissive, may be stated in the
|
||||
form of a separately written license, or stated as exceptions;
|
||||
the above requirements apply either way.
|
||||
|
||||
8. Termination.
|
||||
|
||||
You may not propagate or modify a covered work except as expressly
|
||||
provided under this License. Any attempt otherwise to propagate or
|
||||
modify it is void, and will automatically terminate your rights under
|
||||
this License (including any patent licenses granted under the third
|
||||
paragraph of section 11).
|
||||
|
||||
However, if you cease all violation of this License, then your
|
||||
license from a particular copyright holder is reinstated (a)
|
||||
provisionally, unless and until the copyright holder explicitly and
|
||||
finally terminates your license, and (b) permanently, if the copyright
|
||||
holder fails to notify you of the violation by some reasonable means
|
||||
prior to 60 days after the cessation.
|
||||
|
||||
Moreover, your license from a particular copyright holder is
|
||||
reinstated permanently if the copyright holder notifies you of the
|
||||
violation by some reasonable means, this is the first time you have
|
||||
received notice of violation of this License (for any work) from that
|
||||
copyright holder, and you cure the violation prior to 30 days after
|
||||
your receipt of the notice.
|
||||
|
||||
Termination of your rights under this section does not terminate the
|
||||
licenses of parties who have received copies or rights from you under
|
||||
this License. If your rights have been terminated and not permanently
|
||||
reinstated, you do not qualify to receive new licenses for the same
|
||||
material under section 10.
|
||||
|
||||
9. Acceptance Not Required for Having Copies.
|
||||
|
||||
You are not required to accept this License in order to receive or
|
||||
run a copy of the Program. Ancillary propagation of a covered work
|
||||
occurring solely as a consequence of using peer-to-peer transmission
|
||||
to receive a copy likewise does not require acceptance. However,
|
||||
nothing other than this License grants you permission to propagate or
|
||||
modify any covered work. These actions infringe copyright if you do
|
||||
not accept this License. Therefore, by modifying or propagating a
|
||||
covered work, you indicate your acceptance of this License to do so.
|
||||
|
||||
10. Automatic Licensing of Downstream Recipients.
|
||||
|
||||
Each time you convey a covered work, the recipient automatically
|
||||
receives a license from the original licensors, to run, modify and
|
||||
propagate that work, subject to this License. You are not responsible
|
||||
for enforcing compliance by third parties with this License.
|
||||
|
||||
An "entity transaction" is a transaction transferring control of an
|
||||
organization, or substantially all assets of one, or subdividing an
|
||||
organization, or merging organizations. If propagation of a covered
|
||||
work results from an entity transaction, each party to that
|
||||
transaction who receives a copy of the work also receives whatever
|
||||
licenses to the work the party's predecessor in interest had or could
|
||||
give under the previous paragraph, plus a right to possession of the
|
||||
Corresponding Source of the work from the predecessor in interest, if
|
||||
the predecessor has it or can get it with reasonable efforts.
|
||||
|
||||
You may not impose any further restrictions on the exercise of the
|
||||
rights granted or affirmed under this License. For example, you may
|
||||
not impose a license fee, royalty, or other charge for exercise of
|
||||
rights granted under this License, and you may not initiate litigation
|
||||
(including a cross-claim or counterclaim in a lawsuit) alleging that
|
||||
any patent claim is infringed by making, using, selling, offering for
|
||||
sale, or importing the Program or any portion of it.
|
||||
|
||||
11. Patents.
|
||||
|
||||
A "contributor" is a copyright holder who authorizes use under this
|
||||
License of the Program or a work on which the Program is based. The
|
||||
work thus licensed is called the contributor's "contributor version".
|
||||
|
||||
A contributor's "essential patent claims" are all patent claims
|
||||
owned or controlled by the contributor, whether already acquired or
|
||||
hereafter acquired, that would be infringed by some manner, permitted
|
||||
by this License, of making, using, or selling its contributor version,
|
||||
but do not include claims that would be infringed only as a
|
||||
consequence of further modification of the contributor version. For
|
||||
purposes of this definition, "control" includes the right to grant
|
||||
patent sublicenses in a manner consistent with the requirements of
|
||||
this License.
|
||||
|
||||
Each contributor grants you a non-exclusive, worldwide, royalty-free
|
||||
patent license under the contributor's essential patent claims, to
|
||||
make, use, sell, offer for sale, import and otherwise run, modify and
|
||||
propagate the contents of its contributor version.
|
||||
|
||||
In the following three paragraphs, a "patent license" is any express
|
||||
agreement or commitment, however denominated, not to enforce a patent
|
||||
(such as an express permission to practice a patent or covenant not to
|
||||
sue for patent infringement). To "grant" such a patent license to a
|
||||
party means to make such an agreement or commitment not to enforce a
|
||||
patent against the party.
|
||||
|
||||
If you convey a covered work, knowingly relying on a patent license,
|
||||
and the Corresponding Source of the work is not available for anyone
|
||||
to copy, free of charge and under the terms of this License, through a
|
||||
publicly available network server or other readily accessible means,
|
||||
then you must either (1) cause the Corresponding Source to be so
|
||||
available, or (2) arrange to deprive yourself of the benefit of the
|
||||
patent license for this particular work, or (3) arrange, in a manner
|
||||
consistent with the requirements of this License, to extend the patent
|
||||
license to downstream recipients. "Knowingly relying" means you have
|
||||
actual knowledge that, but for the patent license, your conveying the
|
||||
covered work in a country, or your recipient's use of the covered work
|
||||
in a country, would infringe one or more identifiable patents in that
|
||||
country that you have reason to believe are valid.
|
||||
|
||||
If, pursuant to or in connection with a single transaction or
|
||||
arrangement, you convey, or propagate by procuring conveyance of, a
|
||||
covered work, and grant a patent license to some of the parties
|
||||
receiving the covered work authorizing them to use, propagate, modify
|
||||
or convey a specific copy of the covered work, then the patent license
|
||||
you grant is automatically extended to all recipients of the covered
|
||||
work and works based on it.
|
||||
|
||||
A patent license is "discriminatory" if it does not include within
|
||||
the scope of its coverage, prohibits the exercise of, or is
|
||||
conditioned on the non-exercise of one or more of the rights that are
|
||||
specifically granted under this License. You may not convey a covered
|
||||
work if you are a party to an arrangement with a third party that is
|
||||
in the business of distributing software, under which you make payment
|
||||
to the third party based on the extent of your activity of conveying
|
||||
the work, and under which the third party grants, to any of the
|
||||
parties who would receive the covered work from you, a discriminatory
|
||||
patent license (a) in connection with copies of the covered work
|
||||
conveyed by you (or copies made from those copies), or (b) primarily
|
||||
for and in connection with specific products or compilations that
|
||||
contain the covered work, unless you entered into that arrangement,
|
||||
or that patent license was granted, prior to 28 March 2007.
|
||||
|
||||
Nothing in this License shall be construed as excluding or limiting
|
||||
any implied license or other defenses to infringement that may
|
||||
otherwise be available to you under applicable patent law.
|
||||
|
||||
12. No Surrender of Others' Freedom.
|
||||
|
||||
If conditions are imposed on you (whether by court order, agreement or
|
||||
otherwise) that contradict the conditions of this License, they do not
|
||||
excuse you from the conditions of this License. If you cannot convey a
|
||||
covered work so as to satisfy simultaneously your obligations under this
|
||||
License and any other pertinent obligations, then as a consequence you may
|
||||
not convey it at all. For example, if you agree to terms that obligate you
|
||||
to collect a royalty for further conveying from those to whom you convey
|
||||
the Program, the only way you could satisfy both those terms and this
|
||||
License would be to refrain entirely from conveying the Program.
|
||||
|
||||
13. Use with the GNU Affero General Public License.
|
||||
|
||||
Notwithstanding any other provision of this License, you have
|
||||
permission to link or combine any covered work with a work licensed
|
||||
under version 3 of the GNU Affero General Public License into a single
|
||||
combined work, and to convey the resulting work. The terms of this
|
||||
License will continue to apply to the part which is the covered work,
|
||||
but the special requirements of the GNU Affero General Public License,
|
||||
section 13, concerning interaction through a network will apply to the
|
||||
combination as such.
|
||||
|
||||
14. Revised Versions of this License.
|
||||
|
||||
The Free Software Foundation may publish revised and/or new versions of
|
||||
the GNU General Public License from time to time. Such new versions will
|
||||
be similar in spirit to the present version, but may differ in detail to
|
||||
address new problems or concerns.
|
||||
|
||||
Each version is given a distinguishing version number. If the
|
||||
Program specifies that a certain numbered version of the GNU General
|
||||
Public License "or any later version" applies to it, you have the
|
||||
option of following the terms and conditions either of that numbered
|
||||
version or of any later version published by the Free Software
|
||||
Foundation. If the Program does not specify a version number of the
|
||||
GNU General Public License, you may choose any version ever published
|
||||
by the Free Software Foundation.
|
||||
|
||||
If the Program specifies that a proxy can decide which future
|
||||
versions of the GNU General Public License can be used, that proxy's
|
||||
public statement of acceptance of a version permanently authorizes you
|
||||
to choose that version for the Program.
|
||||
|
||||
Later license versions may give you additional or different
|
||||
permissions. However, no additional obligations are imposed on any
|
||||
author or copyright holder as a result of your choosing to follow a
|
||||
later version.
|
||||
|
||||
15. Disclaimer of Warranty.
|
||||
|
||||
THERE IS NO WARRANTY FOR THE PROGRAM, TO THE EXTENT PERMITTED BY
|
||||
APPLICABLE LAW. EXCEPT WHEN OTHERWISE STATED IN WRITING THE COPYRIGHT
|
||||
HOLDERS AND/OR OTHER PARTIES PROVIDE THE PROGRAM "AS IS" WITHOUT WARRANTY
|
||||
OF ANY KIND, EITHER EXPRESSED OR IMPLIED, INCLUDING, BUT NOT LIMITED TO,
|
||||
THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR
|
||||
PURPOSE. THE ENTIRE RISK AS TO THE QUALITY AND PERFORMANCE OF THE PROGRAM
|
||||
IS WITH YOU. SHOULD THE PROGRAM PROVE DEFECTIVE, YOU ASSUME THE COST OF
|
||||
ALL NECESSARY SERVICING, REPAIR OR CORRECTION.
|
||||
|
||||
16. Limitation of Liability.
|
||||
|
||||
IN NO EVENT UNLESS REQUIRED BY APPLICABLE LAW OR AGREED TO IN WRITING
|
||||
WILL ANY COPYRIGHT HOLDER, OR ANY OTHER PARTY WHO MODIFIES AND/OR CONVEYS
|
||||
THE PROGRAM AS PERMITTED ABOVE, BE LIABLE TO YOU FOR DAMAGES, INCLUDING ANY
|
||||
GENERAL, SPECIAL, INCIDENTAL OR CONSEQUENTIAL DAMAGES ARISING OUT OF THE
|
||||
USE OR INABILITY TO USE THE PROGRAM (INCLUDING BUT NOT LIMITED TO LOSS OF
|
||||
DATA OR DATA BEING RENDERED INACCURATE OR LOSSES SUSTAINED BY YOU OR THIRD
|
||||
PARTIES OR A FAILURE OF THE PROGRAM TO OPERATE WITH ANY OTHER PROGRAMS),
|
||||
EVEN IF SUCH HOLDER OR OTHER PARTY HAS BEEN ADVISED OF THE POSSIBILITY OF
|
||||
SUCH DAMAGES.
|
||||
|
||||
17. Interpretation of Sections 15 and 16.
|
||||
|
||||
If the disclaimer of warranty and limitation of liability provided
|
||||
above cannot be given local legal effect according to their terms,
|
||||
reviewing courts shall apply local law that most closely approximates
|
||||
an absolute waiver of all civil liability in connection with the
|
||||
Program, unless a warranty or assumption of liability accompanies a
|
||||
copy of the Program in return for a fee.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
How to Apply These Terms to Your New Programs
|
||||
|
||||
If you develop a new program, and you want it to be of the greatest
|
||||
possible use to the public, the best way to achieve this is to make it
|
||||
free software which everyone can redistribute and change under these terms.
|
||||
|
||||
To do so, attach the following notices to the program. It is safest
|
||||
to attach them to the start of each source file to most effectively
|
||||
state the exclusion of warranty; and each file should have at least
|
||||
the "copyright" line and a pointer to where the full notice is found.
|
||||
|
||||
<one line to give the program's name and a brief idea of what it does.>
|
||||
Copyright (C) <year> <name of author>
|
||||
|
||||
This program is free software: you can redistribute it and/or modify
|
||||
it under the terms of the GNU General Public License as published by
|
||||
the Free Software Foundation, either version 3 of the License, or
|
||||
(at your option) any later version.
|
||||
|
||||
This program is distributed in the hope that it will be useful,
|
||||
but WITHOUT ANY WARRANTY; without even the implied warranty of
|
||||
MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
|
||||
GNU General Public License for more details.
|
||||
|
||||
You should have received a copy of the GNU General Public License
|
||||
along with this program. If not, see <https://www.gnu.org/licenses/>.
|
||||
|
||||
Also add information on how to contact you by electronic and paper mail.
|
||||
|
||||
If the program does terminal interaction, make it output a short
|
||||
notice like this when it starts in an interactive mode:
|
||||
|
||||
<program> Copyright (C) <year> <name of author>
|
||||
This program comes with ABSOLUTELY NO WARRANTY; for details type `show w'.
|
||||
This is free software, and you are welcome to redistribute it
|
||||
under certain conditions; type `show c' for details.
|
||||
|
||||
The hypothetical commands `show w' and `show c' should show the appropriate
|
||||
parts of the General Public License. Of course, your program's commands
|
||||
might be different; for a GUI interface, you would use an "about box".
|
||||
|
||||
You should also get your employer (if you work as a programmer) or school,
|
||||
if any, to sign a "copyright disclaimer" for the program, if necessary.
|
||||
For more information on this, and how to apply and follow the GNU GPL, see
|
||||
<https://www.gnu.org/licenses/>.
|
||||
|
||||
The GNU General Public License does not permit incorporating your program
|
||||
into proprietary programs. If your program is a subroutine library, you
|
||||
may consider it more useful to permit linking proprietary applications with
|
||||
the library. If this is what you want to do, use the GNU Lesser General
|
||||
Public License instead of this License. But first, please read
|
||||
<https://www.gnu.org/licenses/why-not-lgpl.html>.
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
"Licensor" shall mean the copyright owner or entity authorized by
|
||||
the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all
|
||||
other entities that control, are controlled by, or are under common
|
||||
control with that entity. For the purposes of this definition,
|
||||
"control" means (i) the power, direct or indirect, to cause the
|
||||
direction or management of such entity, whether by contract or
|
||||
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity
|
||||
exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications,
|
||||
including but not limited to software source code, documentation
|
||||
source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical
|
||||
transformation or translation of a Source form, including but
|
||||
not limited to compiled object code, generated documentation,
|
||||
and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or
|
||||
Object form, made available under the License, as indicated by a
|
||||
copyright notice that is included in or attached to the work
|
||||
(an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object
|
||||
form, that is based on (or derived from) the Work and for which the
|
||||
editorial revisions, annotations, elaborations, or other modifications
|
||||
represent, as a whole, an original work of authorship. For the purposes
|
||||
of this License, Derivative Works shall not include works that remain
|
||||
separable from, or merely link (or bind by name) to the interfaces of,
|
||||
the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including
|
||||
the original version of the Work and any modifications or additions
|
||||
to that Work or Derivative Works thereof, that is intentionally
|
||||
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||
or by an individual or Legal Entity authorized to submit on behalf of
|
||||
the copyright owner. For the purposes of this definition, "submitted"
|
||||
means any form of electronic, verbal, or written communication sent
|
||||
to the Licensor or its representatives, including but not limited to
|
||||
communication on electronic mailing lists, source code control systems,
|
||||
and issue tracking systems that are managed by, or on behalf of, the
|
||||
Licensor for the purpose of discussing and improving the Work, but
|
||||
excluding communication that is conspicuously marked or otherwise
|
||||
designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||
on behalf of whom a Contribution has been received by Licensor and
|
||||
subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
copyright license to reproduce, prepare Derivative Works of,
|
||||
publicly display, publicly perform, sublicense, and distribute the
|
||||
Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
(except as stated in this section) patent license to make, have made,
|
||||
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||
where such license applies only to those patent claims licensable
|
||||
by such Contributor that are necessarily infringed by their
|
||||
Contribution(s) alone or by combination of their Contribution(s)
|
||||
with the Work to which such Contribution(s) was submitted. If You
|
||||
institute patent litigation against any entity (including a
|
||||
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||
or a Contribution incorporated within the Work constitutes direct
|
||||
or contributory patent infringement, then any patent licenses
|
||||
granted to You under this License for that Work shall terminate
|
||||
as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the
|
||||
Work or Derivative Works thereof in any medium, with or without
|
||||
modifications, and in Source or Object form, provided that You
|
||||
meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or
|
||||
Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices
|
||||
stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works
|
||||
that You distribute, all copyright, patent, trademark, and
|
||||
attribution notices from the Source form of the Work,
|
||||
excluding those notices that do not pertain to any part of
|
||||
the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its
|
||||
distribution, then any Derivative Works that You distribute must
|
||||
include a readable copy of the attribution notices contained
|
||||
within such NOTICE file, excluding those notices that do not
|
||||
pertain to any part of the Derivative Works, in at least one
|
||||
of the following places: within a NOTICE text file distributed
|
||||
as part of the Derivative Works; within the Source form or
|
||||
documentation, if provided along with the Derivative Works; or,
|
||||
within a display generated by the Derivative Works, if and
|
||||
wherever such third-party notices normally appear. The contents
|
||||
of the NOTICE file are for informational purposes only and
|
||||
do not modify the License. You may add Your own attribution
|
||||
notices within Derivative Works that You distribute, alongside
|
||||
or as an addendum to the NOTICE text from the Work, provided
|
||||
that such additional attribution notices cannot be construed
|
||||
as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and
|
||||
may provide additional or different license terms and conditions
|
||||
for use, reproduction, or distribution of Your modifications, or
|
||||
for any such Derivative Works as a whole, provided Your use,
|
||||
reproduction, and distribution of the Work otherwise complies with
|
||||
the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||
any Contribution intentionally submitted for inclusion in the Work
|
||||
by You to the Licensor shall be under the terms and conditions of
|
||||
this License, without any additional terms or conditions.
|
||||
Notwithstanding the above, nothing herein shall supersede or modify
|
||||
the terms of any separate license agreement you may have executed
|
||||
with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade
|
||||
names, trademarks, service marks, or product names of the Licensor,
|
||||
except as required for reasonable and customary use in describing the
|
||||
origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides the Work (and each
|
||||
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or conditions
|
||||
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||
appropriateness of using or redistributing the Work and assume any
|
||||
risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall any Contributor be
|
||||
liable to You for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as a
|
||||
result of this License or out of the use or inability to use the
|
||||
Work (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if such Contributor
|
||||
has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing
|
||||
the Work or Derivative Works thereof, You may choose to offer,
|
||||
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||
or other liability obligations and/or rights consistent with this
|
||||
License. However, in accepting such obligations, You may act only
|
||||
on Your own behalf and on Your sole responsibility, not on behalf
|
||||
of any other Contributor, and only if You agree to indemnify,
|
||||
defend, and hold each Contributor harmless for any liability
|
||||
incurred by, or claims asserted against, such Contributor by reason
|
||||
of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
APPENDIX: How to apply the Apache License to your work.
|
||||
|
||||
To apply the Apache License to your work, attach the following
|
||||
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||
replaced with your own identifying information. (Don't include
|
||||
the brackets!) The text should be enclosed in the appropriate
|
||||
comment syntax for the file format. We also recommend that a
|
||||
file or class name and description of purpose be included on the
|
||||
same "printed page" as the copyright notice for easier
|
||||
identification within third-party archives.
|
||||
|
||||
Copyright [yyyy] [name of copyright owner]
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
<div align="center">
|
||||
|
||||
<img src="assets/images/logo.png" width="auto" alt="Logo">
|
||||
<img src="docs/images/logo.png" width="auto" alt="Logo">
|
||||
<p>
|
||||
<strong>A lightweight Transformer training & inference framework</strong>
|
||||
</p>
|
||||
@@ -8,27 +8,28 @@
|
||||
|
||||
<div align="center">
|
||||
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
||||
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license">
|
||||
<img src="https://img.shields.io/github/v/release/ViperEkura/AstrAI?color=76bad9" alt="release">
|
||||
<img src="https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fapi.github.com%2Frepos%2FViperEkura%2FAstrAI&query=%24.stargazers_count&label=stars&suffix=%20stars&color=76bad9" alt="stars">
|
||||
<img src="https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fapi.github.com%2Frepos%2FViperEkura%2FAstrAI&query=%24.forks_count&label=forks&suffix=%20forks&color=76bad9" alt="forks">
|
||||
<img src="https://img.shields.io/badge/license-Apache--2.0-blue.svg" alt="license">
|
||||
<img src="https://img.shields.io/github/v/tag/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release">
|
||||
<img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars">
|
||||
<img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks">
|
||||
</div>
|
||||
<br>
|
||||
|
||||
<div align="center">
|
||||
<a href="#english">English</a> •
|
||||
<a href="assets/docs/README-zh-CN.md">中文</a> •
|
||||
<a href="docs/README-zh-CN.md">中文</a> •
|
||||
<a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> •
|
||||
<a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> •
|
||||
<a href="https://huggingface.co/ViperEk/">HuggingFace</a>
|
||||
<a href="https://huggingface.co/ViperEkura">HuggingFace</a>
|
||||
</div>
|
||||
|
||||
<br>
|
||||
|
||||
## 📖 Table of Contents
|
||||
|
||||
- [Features](#features)
|
||||
- [Quick Start](#quick-start)
|
||||
- [Overview](#overview)
|
||||
- [Getting Started](#getting-started)
|
||||
- [Demo](#demo)
|
||||
- [Documentation](#documentation)
|
||||
- [Contributing](#contributing)
|
||||
- [Community](#community)
|
||||
@@ -39,55 +40,132 @@
|
||||
<a id="english"></a>
|
||||
## English
|
||||
|
||||
### Features
|
||||
### Overview
|
||||
|
||||
- 🚀 **High Performance**: Optimized for both training and inference with efficient parallelization.
|
||||
- 🔧 **Flexible**: Support for seq/sft/dpo/grpo training, customizable model architectures.
|
||||
- 💡 **Easy to Use**: Simple API with comprehensive examples and demos.
|
||||
- 📦 **Lightweight**: Minimal dependencies, easy to deploy.
|
||||
- 🔬 **Research‑Friendly**: Modular design, easy to experiment with new ideas.
|
||||
- 🤗 **HuggingFace-Style API**: AutoModel/AutoTokenizer APIs inspired by HuggingFace for easy model and tokenizer loading.
|
||||
- 🔌 **Dual API Compatibility**: Supports both OpenAI and Anthropic chat completion APIs out of the box.
|
||||
AstrAI is an end-to-end Transformer framework for building, training, evaluating, and serving models. It provides a compact PyTorch codebase for the complete model lifecycle, from declarative data preprocessing and distributed training to continuous-batching inference and OpenAI/Anthropic-compatible APIs.
|
||||
|
||||
### Quick Start
|
||||
| Area | Capabilities |
|
||||
|---|---|
|
||||
| **Models** | Autoregressive language models and embedding models with GQA, MLA, MoE, RoPE, and extensible attention/FFN components |
|
||||
| **Training** | Pre-training (`seq`), supervised fine-tuning (`sft`), DPO, and GRPO with gradient accumulation, checkpointing, DDP, and FSDP |
|
||||
| **Data** | Declarative JSON preprocessing, configurable masking and packing, binary/JSONL storage, and streaming datasets |
|
||||
| **Inference** | Continuous batching, paged KV cache, radix prefix caching, streaming generation, and Torch/CUDA/FlashAttention backends |
|
||||
| **Serving** | FastAPI server with OpenAI and Anthropic chat completion protocols, including SSE streaming and tool calls |
|
||||
| **Evaluation** | Perplexity, MMLU, HumanEval, IFEval, IFD, and ROUGE evaluation tools |
|
||||
| **Extensibility** | Factory and registry architecture for models, datasets, training strategies, callbacks, kernels, and protocol components |
|
||||
|
||||
#### Installation
|
||||
### Getting Started
|
||||
|
||||
End-to-end walkthrough in 5 steps:
|
||||
|
||||
**1. Install**
|
||||
|
||||
AstrAI requires Python 3.12+ and pins PyTorch exactly to `2.11.0`. Training, `scripts/tools/generate.py`, generation evaluations, and the generation demos require CUDA; CPU support is limited to components with an explicit CPU device path, such as the HTTP server and direct-scoring evaluations.
|
||||
|
||||
```bash
|
||||
git clone https://github.com/ViperEkura/AstrAI.git
|
||||
cd AstrAI
|
||||
pip install -e .
|
||||
pip install -e . # pure PyTorch (no CUDA kernels)
|
||||
# CSRC_KERNELS=true pip install -e . --no-build-isolation # optional: fused CUDA kernels
|
||||
# pip install -e ".[dev]" # dev dependencies (pytest, ruff)
|
||||
```
|
||||
|
||||
For development dependencies:
|
||||
**2. Download model**
|
||||
|
||||
```bash
|
||||
pip install -e ".[dev]"
|
||||
python scripts/demo/download.py # downloads 1B checkpoint to params/
|
||||
```
|
||||
|
||||
#### Train a Model
|
||||
**3. Preprocess data**
|
||||
|
||||
Create `pretrain.json` (preprocessing config for `seq` strategy):
|
||||
|
||||
```json
|
||||
{
|
||||
"version": 1,
|
||||
"input": {"sections": [{"field": "text", "action": "train"}]},
|
||||
"preprocessing": {"max_seq_len": 2048},
|
||||
"output": {"storage_format": "bin"}
|
||||
}
|
||||
```
|
||||
|
||||
```bash
|
||||
CUDA_VISIBLE_DEVICES=0,1,2,3 python scripts/tools/train.py \
|
||||
--train_type seq \
|
||||
--data_root_path /path/to/dataset \
|
||||
--param_path /path/to/model \
|
||||
--batch_size 4 \
|
||||
--accumulation_steps 8 \
|
||||
--max_lr 3e-4 \
|
||||
--warmup_steps 1000 \
|
||||
--n_epoch 1
|
||||
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c pretrain.json
|
||||
```
|
||||
|
||||
Full reference at [Parameter Guide](assets/docs/params.md).
|
||||
**4. Train**
|
||||
|
||||
#### Generate Text
|
||||
```bash
|
||||
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||
|
||||
nohup python scripts/tools/train.py \
|
||||
--nprocs=4 \
|
||||
--parallel_mode=ddp \
|
||||
--train_type=seq \
|
||||
--data_root_path=/path/to/dataset \
|
||||
--param_path=/path/to/model \
|
||||
--batch_per_device=4 \
|
||||
--grad_accum_steps=8 \
|
||||
--warmup_ratio=0.05 \
|
||||
--max_lr=1e-4 \
|
||||
--max_grad_norm=1.0 \
|
||||
--weight_decay=0.1 \
|
||||
--window_size=2048 \
|
||||
--ckpt_interval=10000 \
|
||||
--ckpt_dir=./checkpoint \
|
||||
--random_seed=3407 \
|
||||
--label_smoothing=0.05 \
|
||||
> out.log 2> err.log &
|
||||
```
|
||||
|
||||
**5. Serve & query**
|
||||
|
||||
```bash
|
||||
# Terminal 1: start server
|
||||
python scripts/tools/server.py --param_path ./params --device cuda
|
||||
|
||||
# Terminal 2: query
|
||||
curl http://localhost:8000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||
```
|
||||
|
||||
### Demo
|
||||
|
||||
Check out the demos in the `scripts/demo/` folder:
|
||||
|
||||
```bash
|
||||
# Download model weights (required before running demos)
|
||||
python scripts/demo/download.py # model → params/
|
||||
|
||||
# Single-turn interactive streaming prompt loop (no conversation history)
|
||||
python scripts/demo/stream_chat.py
|
||||
# Type your message after >>, type !exit to quit
|
||||
|
||||
# Batch generation (5 hardcoded prompts, non-streaming)
|
||||
python scripts/demo/generate_batch.py
|
||||
|
||||
# Single-prompt autoregressive streaming
|
||||
python scripts/demo/generate_ar.py
|
||||
```
|
||||
|
||||
All generation demos use `temperature=0.8`, `top_p=0.95`, `top_k=50`, `max_tokens=2048` by default and require `params/` to contain model weights (run `download.py` first).
|
||||
|
||||
Watch a video walkthrough on [bilibili](https://www.bilibili.com/video/BV1fuLB6yEj6).
|
||||
|
||||
---
|
||||
|
||||
See [Documentation](#documentation) for full references beyond the examples above.
|
||||
|
||||
#### Text Generation
|
||||
|
||||
Batch generation from a JSONL file:
|
||||
|
||||
```bash
|
||||
python scripts/tools/generate.py \
|
||||
--param_path /path/to/model \
|
||||
--input_json_file /path/to/input.json \
|
||||
--output_json_file /path/to/output.json
|
||||
--param_path ./params \
|
||||
--input_json_file input.jsonl \
|
||||
--output_json_file output.jsonl
|
||||
```
|
||||
|
||||
#### Docker
|
||||
@@ -101,9 +179,6 @@ docker build -t astrai:latest .
|
||||
# Run with GPU support
|
||||
docker run --gpus all -it astrai:latest
|
||||
|
||||
# Run with specific GPUs
|
||||
docker run --gpus '"device=0,1"' -it astrai:latest
|
||||
|
||||
# Run inference server
|
||||
docker run --gpus all -p 8000:8000 astrai:latest \
|
||||
python -m scripts.tools.server --port 8000 --device cuda
|
||||
@@ -114,93 +189,53 @@ docker run --gpus all -v /path/to/data:/data -it astrai:latest
|
||||
# Docker Compose (GPU, default)
|
||||
docker compose up -d
|
||||
|
||||
# Docker Compose (CPU only)
|
||||
# Docker Compose CPU server profile (CUDA-only generation scripts/demos are unavailable)
|
||||
docker compose --profile cpu up -d
|
||||
```
|
||||
|
||||
> **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`.
|
||||
|
||||
#### Start HTTP Server
|
||||
#### HTTP API Examples
|
||||
|
||||
Start the inference server with OpenAI and Anthropic-compatible HTTP API:
|
||||
Additional request examples beyond the [Getting Started](#getting-started) flow:
|
||||
|
||||
```bash
|
||||
python -m scripts.tools.server --port 8000 --device cuda
|
||||
```
|
||||
|
||||
Make requests:
|
||||
|
||||
```bash
|
||||
# OpenAI-compatible
|
||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"max_tokens": 512
|
||||
}'
|
||||
|
||||
# OpenAI-compatible streaming
|
||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"messages": [{"role": "user", "content": "Tell a story"}],
|
||||
"stream": true,
|
||||
"max_tokens": 500
|
||||
}'
|
||||
-d '{"messages":[{"role":"user","content":"Tell a story"}],"stream":true,"max_tokens":500}'
|
||||
|
||||
# Anthropic-compatible
|
||||
curl -X POST http://localhost:8000/v1/messages \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "astrai",
|
||||
"system": "You are a helpful assistant.",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"max_tokens": 512
|
||||
}'
|
||||
-d '{"model":"astrai","system":"You are a helpful assistant.","messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||
|
||||
# Anthropic-compatible streaming with stop sequences
|
||||
curl -X POST http://localhost:8000/v1/messages \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "astrai",
|
||||
"messages": [{"role": "user", "content": "Write a story"}],
|
||||
"max_tokens": 500,
|
||||
"stream": true,
|
||||
"stop_sequences": ["The end"]
|
||||
}'
|
||||
-d '{"model":"astrai","messages":[{"role":"user","content":"Write a story"}],"max_tokens":500,"stream":true,"stop_sequences":["The end"]}'
|
||||
|
||||
# Health check
|
||||
curl http://localhost:8000/health
|
||||
```
|
||||
|
||||
#### Demo
|
||||
|
||||
Check out the demos in the `scripts/demo/` folder:
|
||||
|
||||
```bash
|
||||
# Download pre‑processed data (required before running demos)
|
||||
python scripts/demo/download.py
|
||||
|
||||
# Interactive streaming chat
|
||||
python scripts/demo/stream_chat.py
|
||||
|
||||
# Batch generation
|
||||
python scripts/demo/generate_batch.py
|
||||
|
||||
# Auto‑regressive generation
|
||||
python scripts/demo/generate_ar.py
|
||||
```
|
||||
|
||||
Watch a video walkthrough on [bilibili](https://www.bilibili.com/video/BV1z5RPYHEkd).
|
||||
See [Inference Guide](docs/guides/inference.md) for SSE streaming format, error codes, and stats endpoint.
|
||||
|
||||
### Documentation
|
||||
|
||||
| Document | Description |
|
||||
|----------|-------------|
|
||||
| [Parameter Guide](./assets/docs/params.md) | Training & inference parameters |
|
||||
| [Design Document](./assets/docs/design.md) | Framework architecture & module design |
|
||||
| [Data Flow](./assets/docs/dataflow.md) | Data processing pipeline details |
|
||||
| [Model Introduction](./assets/docs/introduction.md) | Model architecture & technical details |
|
||||
| [Get Started](./docs/get-started.md) | Installation and quickstart |
|
||||
| [CLI Reference](./docs/guides/params.md) | Parameters for all CLI tools (train, server, generate, preprocess) |
|
||||
| [Preprocessing](./docs/guides/preprocessing.md) | Declarative JSON-driven data preprocessing |
|
||||
| [Training](./docs/guides/training.md) | Training loop, strategies & formulas |
|
||||
| [Inference](./docs/guides/inference.md) | KVCache, continuous batching, sampling & HTTP API |
|
||||
| [Evaluation](./docs/guides/evaluation.md) | HumanEval, MMLU, PPL, ROUGE, IFD, IFEval |
|
||||
| [Distributed](./docs/guides/distributed.md) | Multi-GPU DDP / FSDP training |
|
||||
| [Architecture](./docs/developer/architecture.md) | System architecture, class diagram & design patterns |
|
||||
| [Data Flow](./docs/developer/dataflow.md) | Data pipeline, storage backends & dataset architecture |
|
||||
| [Internals](./docs/developer/internals.md) | Training internals: loss formulas, callback lifecycle, KV cache |
|
||||
| [CUDA Kernels](./docs/developer/cuda_kernels.md) | Custom CUDA attention kernels & benchmarks |
|
||||
|
||||
### Contributing
|
||||
|
||||
@@ -217,14 +252,14 @@ For major changes, please open an issue first to discuss what you would like to
|
||||
|
||||
- **GitHub Issues**: [Issue Tracker](https://github.com/ViperEkura/AstrAI/issues)
|
||||
- **Discussions**: [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions)
|
||||
- **HuggingFace**: [Model Hub](https://huggingface.co/ViperEk)
|
||||
- **HuggingFace**: [Model Hub](https://huggingface.co/ViperEkura)
|
||||
|
||||
### License
|
||||
|
||||
This project is licensed under the [GPL-3.0 License](LICENSE).
|
||||
This project is licensed under the [Apache-2.0 License](LICENSE).
|
||||
|
||||
---
|
||||
|
||||
<div align="center">
|
||||
<em>A lightweight Transformer framework designed for both high performance and ease of use.</em>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
@@ -1,236 +0,0 @@
|
||||
<div align="center">
|
||||
|
||||
<img src="../images/logo.png" width="auto" alt="Logo">
|
||||
|
||||
<div>
|
||||
<a href="../../README.md">English</a> •
|
||||
<a href="#chinese">中文</a>
|
||||
</div>
|
||||
|
||||
<p>
|
||||
<strong>轻量级 Transformer 训练与推理框架</strong>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div align="center">
|
||||
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
||||
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license">
|
||||
<img src="https://img.shields.io/github/v/release/ViperEkura/AstrAI?color=76bad9" alt="release">
|
||||
<img src="https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fapi.github.com%2Frepos%2FViperEkura%2FAstrAI&query=%24.stargazers_count&label=stars&suffix=%20stars&color=76bad9" alt="stars">
|
||||
<img src="https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fapi.github.com%2Frepos%2FViperEkura%2FAstrAI&query=%24.forks_count&label=forks&suffix=%20forks&color=76bad9" alt="forks">
|
||||
</div>
|
||||
|
||||
<br>
|
||||
|
||||
<div align="center">
|
||||
<a href="../../README.md">English</a> •
|
||||
<a href="#chinese">中文</a> •
|
||||
<a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> •
|
||||
<a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> •
|
||||
<a href="https://huggingface.co/ViperEk">HuggingFace</a>
|
||||
</div>
|
||||
<br>
|
||||
|
||||
## 📖 目录
|
||||
|
||||
- [特性](#特性)
|
||||
- [快速开始](#快速开始)
|
||||
- [文档](#文档)
|
||||
- [贡献](#贡献)
|
||||
- [社区](#社区)
|
||||
- [许可证](#许可证)
|
||||
|
||||
---
|
||||
|
||||
<a id="chinese"></a>
|
||||
## 中文
|
||||
|
||||
### 特性
|
||||
|
||||
- 🚀 **高性能**: 训练与推理双向优化,高效并行。
|
||||
- 🔧 **灵活**: 支持 seq/sft/dpo/grpo 多种训练方式,可定制模型架构。
|
||||
- 💡 **易用**: 简洁的 API 与丰富的示例、演示。
|
||||
- 📦 **轻量**: 依赖少,部署简单。
|
||||
- 🔬 **研究友好**: 模块化设计,便于实验新想法。
|
||||
- 🤗 **HuggingFace 风格 API**: 类 HuggingFace 的 AutoModel/AutoTokenizer 接口,方便加载模型和分词器。
|
||||
- 🔌 **双 API 兼容**: 同时支持 OpenAI 和 Anthropic 聊天补全 API,开箱即用。
|
||||
|
||||
### 快速开始
|
||||
|
||||
#### 安装
|
||||
|
||||
```bash
|
||||
git clone https://github.com/ViperEkura/AstrAI.git
|
||||
cd AstrAI
|
||||
pip install -e .
|
||||
```
|
||||
|
||||
安装开发依赖:
|
||||
|
||||
```bash
|
||||
pip install -e ".[dev]"
|
||||
```
|
||||
|
||||
#### 训练模型
|
||||
|
||||
```bash
|
||||
CUDA_VISIBLE_DEVICES=0,1,2,3 python scripts/tools/train.py \
|
||||
--train_type seq \
|
||||
--data_root_path /path/to/dataset \
|
||||
--param_path /path/to/model \
|
||||
--batch_size 4 \
|
||||
--accumulation_steps 8 \
|
||||
--max_lr 3e-4 \
|
||||
--warmup_steps 1000 \
|
||||
--n_epoch 1
|
||||
```
|
||||
|
||||
完整参数列表见[参数说明](./params.md)。
|
||||
|
||||
#### 文本生成
|
||||
|
||||
```bash
|
||||
python scripts/tools/generate.py \
|
||||
--param_path /path/to/model \
|
||||
--input_json_file /path/to/input.json \
|
||||
--output_json_file /path/to/output.json
|
||||
```
|
||||
|
||||
#### Docker
|
||||
|
||||
使用 Docker 构建和运行(推荐用于 GPU 环境):
|
||||
|
||||
```bash
|
||||
# 构建镜像
|
||||
docker build -t astrai:latest .
|
||||
|
||||
# 启用 GPU 运行
|
||||
docker run --gpus all -it astrai:latest
|
||||
|
||||
# 指定特定 GPU
|
||||
docker run --gpus '"device=0,1"' -it astrai:latest
|
||||
|
||||
# 运行推理服务
|
||||
docker run --gpus all -p 8000:8000 astrai:latest \
|
||||
python -m scripts.tools.server --port 8000 --device cuda
|
||||
|
||||
# 挂载数据卷
|
||||
docker run --gpus all -v /path/to/data:/data -it astrai:latest
|
||||
|
||||
# Docker Compose(GPU,默认)
|
||||
docker compose up -d
|
||||
|
||||
# Docker Compose(仅 CPU)
|
||||
docker compose --profile cpu up -d
|
||||
```
|
||||
|
||||
> **注意**: 必须使用 `--gpus all` 才能启用 CUDA 支持,否则 `torch.cuda.is_available()` 将返回 `False`。
|
||||
|
||||
#### 启动 HTTP 服务
|
||||
|
||||
启动推理服务器,支持 OpenAI 和 Anthropic 兼容的 HTTP API:
|
||||
|
||||
```bash
|
||||
python -m scripts.tools.server --port 8000 --device cuda
|
||||
```
|
||||
|
||||
发起请求:
|
||||
|
||||
```bash
|
||||
# OpenAI 兼容
|
||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"messages": [{"role": "user", "content": "你好"}],
|
||||
"max_tokens": 512
|
||||
}'
|
||||
|
||||
# OpenAI 兼容流式
|
||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"messages": [{"role": "user", "content": "讲个故事"}],
|
||||
"stream": true,
|
||||
"max_tokens": 500
|
||||
}'
|
||||
|
||||
# Anthropic 兼容
|
||||
curl -X POST http://localhost:8000/v1/messages \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "astrai",
|
||||
"system": "你是一个乐于助人的助手。",
|
||||
"messages": [{"role": "user", "content": "你好"}],
|
||||
"max_tokens": 512
|
||||
}'
|
||||
|
||||
# Anthropic 兼容流式并设置停止序列
|
||||
curl -X POST http://localhost:8000/v1/messages \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "astrai",
|
||||
"messages": [{"role": "user", "content": "写个故事"}],
|
||||
"max_tokens": 500,
|
||||
"stream": true,
|
||||
"stop_sequences": ["结束"]
|
||||
}'
|
||||
|
||||
# 健康检查
|
||||
curl http://localhost:8000/health
|
||||
```
|
||||
|
||||
#### 演示
|
||||
|
||||
查看 `scripts/demo/` 文件夹中的演示:
|
||||
|
||||
```bash
|
||||
# 下载预处理数据(运行演示前必需)
|
||||
python scripts/demo/download.py
|
||||
|
||||
# 交互式流式聊天
|
||||
python scripts/demo/stream_chat.py
|
||||
|
||||
# 批量生成
|
||||
python scripts/demo/generate_batch.py
|
||||
|
||||
# 自回归生成
|
||||
python scripts/demo/generate_ar.py
|
||||
```
|
||||
|
||||
观看 [bilibili](https://www.bilibili.com/video/BV1z5RPYHEkd) 上的视频演示。
|
||||
|
||||
### 文档
|
||||
|
||||
| 文档 | 说明 |
|
||||
|------|------|
|
||||
| [参数说明](./params.md) | 训练与推理参数配置 |
|
||||
| [设计文档](./design.md) | 系统架构与模块设计 |
|
||||
| [数据流程](./dataflow.md) | 数据处理管道详解 |
|
||||
| [模型介绍](./introduction.md) | 模型架构与技术细节 |
|
||||
|
||||
### 贡献
|
||||
|
||||
我们欢迎贡献!请参阅[贡献指南](../../CONTRIBUTING.md)了解详情。
|
||||
|
||||
1. Fork 本仓库。
|
||||
2. 创建功能分支。
|
||||
3. 提交更改。
|
||||
4. 发起 Pull Request。
|
||||
|
||||
重大更改请先开 issue 讨论。
|
||||
|
||||
### 社区
|
||||
|
||||
- **GitHub Issues**: [问题追踪](https://github.com/ViperEkura/AstrAI/issues)
|
||||
- **Discussions**: [GitHub 讨论区](https://github.com/ViperEkura/AstrAI/discussions)
|
||||
- **HuggingFace**: [模型中心](https://huggingface.co/ViperEk)
|
||||
|
||||
### 许可证
|
||||
|
||||
本项目采用 [GPL-3.0 许可证](../../LICENSE)。
|
||||
|
||||
---
|
||||
|
||||
<div align="center">
|
||||
<em>专为高性能与易用性设计的轻量级 Transformer 框架。</em>
|
||||
</div>
|
||||
@@ -1,237 +0,0 @@
|
||||
# AstrAI Data Flow Documentation
|
||||
|
||||
This document describes the data flow of the AstrAI project (a training and inference framework for autoregressive Transformer language models). It covers the complete flow from raw data to model training and inference.
|
||||
|
||||
## Overview
|
||||
|
||||
AstrAI adopts a modular design with the following main components:
|
||||
- **Dataset Module** (`astrai/dataset/`): Dataset, sampler, serialization tools
|
||||
- **Model Module** (`astrai/model/`): AutoModel, Transformer model and its submodules
|
||||
- **Training Module** (`astrai/trainer/`): Trainer, training context, strategies, schedulers, callbacks, metric utilities
|
||||
- **Inference Module** (`astrai/inference/`): Inference engine with continuous batching, streaming generation
|
||||
- **Config Module** (`astrai/config/`): ModelConfig, TrainConfig
|
||||
- **Factory Module** (`astrai/factory/`): Registry, BaseFactory for component registration
|
||||
- **Parallel Module** (`astrai/parallel/`): Distributed training support
|
||||
- **Serialization** (`astrai/serialization.py`): HDF5 data loading, checkpoint management
|
||||
|
||||
## Data Flow Diagram
|
||||
|
||||
```mermaid
|
||||
flowchart LR
|
||||
subgraph A[Data Preparation]
|
||||
direction TB
|
||||
A1[Raw Text] --> A2[AutoTokenizer]
|
||||
A2 --> A3[Tokenized .h5 files]
|
||||
A3 --> A4[BaseDataset]
|
||||
A4 --> A5[ResumableDistributedSampler]
|
||||
A5 --> A6[DataLoader]
|
||||
end
|
||||
|
||||
subgraph B[Training]
|
||||
direction TB
|
||||
B1[DataLoader] --> B2[BaseStrategy]
|
||||
B2 --> B3[Transformer Forward]
|
||||
B3 --> B4[Loss + Backward]
|
||||
B4 --> B5[Gradient Accumulation]
|
||||
B5 -->|every accum_steps| B6[Optimizer Step]
|
||||
B6 --> B7[LR Scheduler]
|
||||
B7 -->|next batch| B2
|
||||
B6 --> B8[CheckpointCallback]
|
||||
end
|
||||
|
||||
subgraph C[Inference]
|
||||
direction TB
|
||||
C1[Checkpoint] --> C2[AutoModel]
|
||||
C1 --> C3[AutoTokenizer]
|
||||
C2 --> C4[InferenceEngine]
|
||||
C3 --> C4
|
||||
C4 --> C5[InferenceScheduler]
|
||||
C5 --> C6[Transformer Forward]
|
||||
C6 --> C7[sample]
|
||||
C7 --> C8{End?}
|
||||
C8 -->|No| C6
|
||||
C8 -->|Yes| C9[Generated Text]
|
||||
end
|
||||
|
||||
A --> B
|
||||
B --> C
|
||||
```
|
||||
|
||||
## Detailed Module Descriptions
|
||||
|
||||
### 1. Serialization (`astrai/serialization.py`)
|
||||
|
||||
- **`save_h5`**: Saves tensors by groups as HDF5 files (`.h5`), each key maps to a list of tensors
|
||||
- **`load_h5`**: Loads `.h5` files, returns `Dict[str, List[Tensor]]`, supports shared memory
|
||||
- **`Checkpoint`**: Encapsulates model state dict + epoch + iteration; uses safetensors
|
||||
|
||||
### 2. Dataset Module
|
||||
|
||||
#### 2.1 Dataset (`dataset.py`)
|
||||
- **`BaseDataset`**: Abstract base class for windowed sequence sampling
|
||||
- **`BaseSegmentFetcher` / `MultiSegmentFetcher`**: Fetch tensor segments by index range
|
||||
- **`DatasetFactory`**: Creates dataset instances by `train_type` (`seq`, `sft`, `dpo`, `grpo`)
|
||||
- Data keys: `"sequence"` (SEQ), `"loss_mask"` (SFT), `"chosen_mask"/"rejected_mask"` (DPO), `"masks"` (GRPO)
|
||||
|
||||
#### 2.2 Sampler (`sampler.py`)
|
||||
- **`ResumableDistributedSampler`**: Tracks `epoch` and `iter` for breakpoint resume; supports shuffle and drop_last
|
||||
|
||||
### 3. Model Module
|
||||
|
||||
#### 3.1 Transformer / AutoModel
|
||||
- **`AutoModel`**: Base class with `from_pretrained()` / `save_pretrained()`
|
||||
- **`Transformer`**: Decoder-only architecture, registered via `@AutoModel.register('transformer')`
|
||||
- Embedding → N×DecoderBlock → RMSNorm → Linear lm_head
|
||||
- RoPE position encoding, optional weight tying
|
||||
|
||||
#### 3.2 Submodules (`module.py`)
|
||||
- **`DecoderBlock`**: GQA attention + residual + MLP + RMSNorm
|
||||
- **`GQA`**: Grouped Query Attention (also `MLA` for multi-latent attention)
|
||||
- **`MLP`**: `SiLU(gate(x)) * up(x)` → down projection
|
||||
- **`RotaryEmbedding`**: RoPE cos/sin cache
|
||||
- **`RMSNorm`**: Layer normalization
|
||||
|
||||
### 4. Training Module
|
||||
|
||||
#### 4.1 Training Context (`train_context.py`)
|
||||
- **`TrainContext`**: Dataclass holding model, optimizer, dataloader, strategy, scheduler, checkpoint state
|
||||
- **`TrainContextBuilder`**: Builder pattern — takes checkpoint for resume, builds all components
|
||||
|
||||
#### 4.2 Trainer (`trainer.py`)
|
||||
|
||||
The training loop is nested: **epoch** → **batch** (with step phase interspersed):
|
||||
|
||||
```
|
||||
on_train_begin
|
||||
on_epoch_begin
|
||||
for each batch:
|
||||
if iteration % accumulation_steps == 0: ← step phase
|
||||
on_step_begin → optimizer.step() → zero_grad → on_step_end
|
||||
← batch phase
|
||||
on_batch_begin → strategy(batch) → loss → backward → on_batch_end
|
||||
iteration += 1
|
||||
|
||||
on_epoch_end
|
||||
on_train_end
|
||||
```
|
||||
|
||||
Key points:
|
||||
- `on_step_*` wraps optimizer step (fires every `accumulation_steps` batches)
|
||||
- `on_batch_*` wraps loss computation (fires every batch)
|
||||
- `SchedulerCallback` fires on `on_batch_end` — LR scheduler steps every batch
|
||||
- `GradientClippingCallback` fires on `on_step_begin`
|
||||
|
||||
#### 4.3 Strategy (`strategy.py`)
|
||||
- **`SEQStrategy`**: Next-token prediction, cross-entropy with label smoothing
|
||||
- **`SFTStrategy`**: Supervised fine-tuning with loss masking
|
||||
- **`DPOStrategy`**: Direct Preference Optimization with reference model
|
||||
- **`GRPOStrategy`**: Group Relative Policy Optimization with clipped ratio
|
||||
|
||||
#### 4.4 Scheduler (`schedule.py`)
|
||||
- **`CosineScheduler`**: Cosine decay + linear warmup
|
||||
- **`SGDRScheduler`**: Cosine annealing with warm restarts
|
||||
- Created by `SchedulerFactory` and bound to optimizer
|
||||
|
||||
#### 4.5 Callbacks
|
||||
- **`CheckpointCallback`**: Saves safetensors at `ckpt_interval` iterations
|
||||
- **`ProgressBarCallback`**: tqdm progress display
|
||||
- **`MetricLoggerCallback`**: Writes JSONL metrics to `{ckpt_dir}/logs/`
|
||||
- **`GradientClippingCallback`**: `clip_grad_norm_` on `on_step_begin`
|
||||
- **`SchedulerCallback`**: `scheduler.step()` on `on_batch_end`
|
||||
|
||||
### 5. Inference Module
|
||||
|
||||
#### 5.1 Inference Engine (`engine.py`)
|
||||
- **`InferenceEngine`**: Facade over scheduler; provides `generate()`, `generate_with_request()`, `generate_async()`
|
||||
- Accepts `prompt: str | List[str]`, returns generator (stream) or string (non-stream)
|
||||
|
||||
#### 5.2 Scheduler 4-Phase Loop (`scheduler.py`)
|
||||
|
||||
Background thread runs continuously:
|
||||
|
||||
```
|
||||
1. Cleanup → Remove finished tasks, free KV cache pages
|
||||
2. Refill → Pop from waiting_queue, alloc pages, add to active
|
||||
3. Prefill → Group active tasks by prompt_len, run full forward pass
|
||||
4. Decode → Pick largest same-position group, run single-token forward
|
||||
```
|
||||
|
||||
- **`Task`**: Tracks prompt_ids, output_ids, page_table, status (PENDING/RUNNING/FINISHED/ABORTED)
|
||||
- **`PagedCache`**: Bitmask-based page allocator with page-table-indirected read/write
|
||||
- **`CacheView`**: Batch view bundling cache + page table for attention layers
|
||||
- **`sample()`**: Temperature → top-k → top-p → multinomial
|
||||
|
||||
#### 5.3 Server (`server.py`)
|
||||
- FastAPI with OpenAI `/v1/chat/completions` and Anthropic `/v1/messages` endpoints
|
||||
- Streaming via SSE, health check at `/health`, stats at `/stats`
|
||||
|
||||
### 6. Tokenizer Module
|
||||
|
||||
- **`AutoTokenizer`**: Wraps HuggingFace tokenizers (BBPE); `encode`/`decode`/`apply_chat_template`
|
||||
- **`ChatTemplate`**: Jinja2-based template rendering for multi-turn chat
|
||||
|
||||
### 7. Factory & Parallel
|
||||
|
||||
- **`Registry` / `BaseFactory`**: Decorator-based component registration
|
||||
- **`spawn_parallel_fn`**: Multi-process DDP launcher with NCCL backend
|
||||
- **`ParallelModel` / `ColumnParallelLinear` / `RowParallelLinear`**: Tensor model parallelism
|
||||
|
||||
## Training Data Flow — Detailed Steps
|
||||
|
||||
1. **Data Preparation**
|
||||
- Raw text → token IDs via `AutoTokenizer.encode()`
|
||||
- Save as `.h5` files (groups of tensor lists per data key)
|
||||
|
||||
2. **Dataset Loading**
|
||||
- `BaseDataset.load()` calls `load_h5()`, builds `MultiSegmentFetcher`
|
||||
- Sliding window of `window_size` with `stride` determines sample boundaries
|
||||
|
||||
3. **Sampling & Batching**
|
||||
- `ResumableDistributedSampler` produces shuffled index sequences
|
||||
- `DataLoader` fetches `[batch_size, window_size]` tensors via `__getitem__`
|
||||
|
||||
4. **Strategy Forward**
|
||||
- Strategy receives batch, calls `Transformer.forward()` for logits
|
||||
- Computes task-specific loss (cross-entropy, DPO, GRPO)
|
||||
|
||||
5. **Backward & Accumulation**
|
||||
- `loss = raw_loss / accumulation_steps`
|
||||
- `loss.backward()` accumulates gradients
|
||||
- Every `accumulation_steps` batches: `optimizer.step()` → `zero_grad()`
|
||||
- Every batch: `scheduler.step()` updates learning rate
|
||||
|
||||
6. **Checkpoint**
|
||||
- `CheckpointCallback` saves `model.state_dict()` + metadata to safetensors at `ckpt_interval` iterations
|
||||
- Does NOT save optimizer/scheduler state (resume resets those)
|
||||
|
||||
## Inference Data Flow — Detailed Steps
|
||||
|
||||
1. **Model Loading**
|
||||
- `AutoModel.from_pretrained(path)` loads weights from safetensors
|
||||
- `torch.inference_mode()` wraps generation
|
||||
|
||||
2. **Prompt Construction**
|
||||
- Messages → `apply_chat_template(messages, tokenize=False)` → prompt string
|
||||
- `tokenizer.encode(prompt)` → token IDs (truncated to `max_prompt_len`)
|
||||
|
||||
3. **Continuous Batching Loop**
|
||||
- **Cleanup**: Finished tasks → `stream_callback(STOP)`, free KV pages
|
||||
- **Refill**: Pop from waiting queue, `PagedCache.alloc_n()` for prompt pages
|
||||
- **Prefill**: Group by prompt length, run full forward with `start_pos=0`
|
||||
- **Decode**: Pick position group with most tasks, single-token forward:
|
||||
- Model forward → `logits` → `sample()` → next token ID
|
||||
- Append to `output_ids`, update `output_tokens`
|
||||
- `_maybe_alloc_page()` grows page table as needed
|
||||
- `stream_callback(token)` for streaming clients
|
||||
|
||||
4. **Output**
|
||||
- `tokenizer.decode(output_ids)` → text
|
||||
- Return to caller (streaming: token-by-token; non-streaming: complete string)
|
||||
|
||||
## Checkpoint & Serialization
|
||||
|
||||
- **Training Checkpoint**: safetensors weights + epoch/iteration metadata. Optimizer/scheduler state is NOT persisted.
|
||||
- **Inference Loading**: `AutoModel.from_pretrained()` loads from the same safetensors format.
|
||||
- **Dataset Serialization**: HDF5 with shared memory support for large-scale pre-training data.
|
||||
|
||||
> Document Update Time: 2026-05-09
|
||||
@@ -1,719 +0,0 @@
|
||||
## 1. Why I Created This Project
|
||||
|
||||
There are many large language models on the market today, such as GPT, LLaMA, and others, with tens of billions or even hundreds of billions of parameters. But honestly, these models have extremely high hardware requirements, making them inaccessible for ordinary developers. I thought: **Can we create a model that is both useful and can run on ordinary computers?** This is also what most people currently hope for - a locally deployable AI project that achieves complete privatization while maintaining some level of intelligence.
|
||||
|
||||
Thus, the AstrAI project was born - 1B parameters, Chinese-English bilingual, supporting dialogue, text generation, and the training code is open source!
|
||||
|
||||
## 2. System Architecture
|
||||
|
||||
```mermaid
|
||||
classDiagram
|
||||
namespace config {
|
||||
class ModelConfig {
|
||||
+int vocab_size
|
||||
+int dim
|
||||
+int n_layers
|
||||
+float norm_eps
|
||||
+int dim_ffn
|
||||
+bool tie_weight
|
||||
+int max_len
|
||||
+float rope_theta
|
||||
+int n_heads
|
||||
+int n_kv_heads
|
||||
+bool use_qk_norm
|
||||
+bool use_gated_attention
|
||||
+load(config_path) ModelConfig
|
||||
+save(config_path)
|
||||
}
|
||||
|
||||
class TrainConfig {
|
||||
+nn.Module model
|
||||
+str strategy
|
||||
+Dataset dataset
|
||||
+Callable optimizer_fn
|
||||
+Callable scheduler_fn
|
||||
+int n_epoch
|
||||
+int batch_size
|
||||
+int accumulation_steps
|
||||
+float max_grad_norm
|
||||
+int start_epoch
|
||||
+int start_batch
|
||||
+str ckpt_dir
|
||||
+int ckpt_interval
|
||||
+int random_seed
|
||||
+int num_workers
|
||||
+int prefetch_factor
|
||||
+bool pin_memory
|
||||
+int nprocs
|
||||
+str backend
|
||||
+str master_addr
|
||||
+str master_port
|
||||
+Callable parallel_wrapper
|
||||
+Callable state_dict_fn
|
||||
+str device_type
|
||||
+dict extra_kwargs
|
||||
+validate()
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
namespace dataset {
|
||||
class BaseDataset {
|
||||
+int window_size
|
||||
+int stride
|
||||
+MultiSegmentFetcher fetcher
|
||||
+load(load_path)
|
||||
+__getitem__(index)
|
||||
+__len__()
|
||||
}
|
||||
|
||||
class SEQDataset {
|
||||
+__getitem__(index) Dict
|
||||
}
|
||||
|
||||
class SFTDataset {
|
||||
+__getitem__(index) Dict
|
||||
}
|
||||
|
||||
class DPODataset {
|
||||
+__getitem__(index) Dict
|
||||
}
|
||||
|
||||
class GRPODataset {
|
||||
+__getitem__(index) Dict
|
||||
}
|
||||
|
||||
class BaseSegmentFetcher {
|
||||
+List[Tensor] segments
|
||||
+List[int] cum_lengths
|
||||
+int total_length
|
||||
+fetch_data(begin_idx, end_idx) Tensor
|
||||
}
|
||||
|
||||
class MultiSegmentFetcher {
|
||||
+Dict multi_fetchers
|
||||
+List multi_keys
|
||||
+key_fetch(begin_idx, end_idx, keys) Dict
|
||||
+fetch_data(begin_idx, end_idx) Dict
|
||||
}
|
||||
|
||||
class ResumableDistributedSampler {
|
||||
+int epoch
|
||||
+int iter
|
||||
}
|
||||
|
||||
class DatasetFactory {
|
||||
+Registry _registry
|
||||
+register(name) decorator
|
||||
+create(train_type, window_size, stride) BaseDataset
|
||||
+load(train_type, load_path, window_size, stride) BaseDataset
|
||||
}
|
||||
}
|
||||
|
||||
namespace serialization {
|
||||
class Checkpoint {
|
||||
+dict state_dict
|
||||
+int epoch
|
||||
+int iteration
|
||||
+save(save_dir)
|
||||
+load(save_dir) Checkpoint
|
||||
}
|
||||
}
|
||||
|
||||
namespace model {
|
||||
class AutoModel {
|
||||
+ModelConfig config
|
||||
+Registry _registry
|
||||
+register(model_type) decorator
|
||||
+get_model_class(model_type) Type
|
||||
+from_pretrained(path, disable_random_init) nn.Module
|
||||
+save_pretrained(save_directory)
|
||||
+to(*args, **kwargs) Self
|
||||
}
|
||||
|
||||
class Transformer {
|
||||
+ModelConfig config
|
||||
+RotaryEmbedding rotary_embedding
|
||||
+Embedding embed_tokens
|
||||
+ModuleList layers
|
||||
+RMSNorm norm
|
||||
+Linear lm_head
|
||||
+forward(input_ids, input_mask, paged_cache, start_pos) Dict
|
||||
+load_state_dict(state_dict)
|
||||
+state_dict()
|
||||
}
|
||||
|
||||
class DecoderBlock {
|
||||
+GQA attention
|
||||
+RMSNorm input_norm
|
||||
+MLP mlp
|
||||
+RMSNorm post_attention_norm
|
||||
+forward(x, rotary_emb, attention_mask, paged_cache, start_pos) Tensor
|
||||
}
|
||||
|
||||
class GQA {
|
||||
+int n_heads
|
||||
+int n_kv_heads
|
||||
+int head_dim
|
||||
+Linear q_proj, k_proj, v_proj, o_proj
|
||||
+RMSNorm q_norm, k_norm
|
||||
+forward(x, rotary_emb, mask, paged_cache, start_pos) Tensor
|
||||
}
|
||||
|
||||
class MLA {
|
||||
+int n_heads
|
||||
+int n_kv_heads
|
||||
+int head_dim
|
||||
+int kv_lora_rank
|
||||
+int qk_nope_head_dim
|
||||
+int qk_rope_head_dim
|
||||
+Linear q_proj, kv_a_proj, kv_b_proj
|
||||
+Linear o_proj
|
||||
+RMSNorm kv_norm
|
||||
+forward(x, rotary_emb, mask, paged_cache, start_pos) Tensor
|
||||
}
|
||||
|
||||
class MLP {
|
||||
+Linear up, gate, down
|
||||
+forward(x) Tensor
|
||||
}
|
||||
|
||||
class RMSNorm {
|
||||
+Parameter weight
|
||||
+float norm_eps
|
||||
+forward(x) Tensor
|
||||
}
|
||||
|
||||
class Linear {
|
||||
+Parameter weight
|
||||
+Parameter bias
|
||||
+forward(x) Tensor
|
||||
}
|
||||
|
||||
class RotaryEmbedding {
|
||||
+int dim
|
||||
+int max_len
|
||||
+float base
|
||||
+forward(x, start_pos) Tuple[Tensor, Tensor]
|
||||
}
|
||||
|
||||
class Embedding {
|
||||
+Parameter weight
|
||||
+forward(x) Tensor
|
||||
}
|
||||
}
|
||||
|
||||
namespace tokenize {
|
||||
class AutoTokenizer {
|
||||
+List[int] stop_ids
|
||||
+int bos_id
|
||||
+int eos_id
|
||||
+int pad_id
|
||||
+vocab_size int
|
||||
+encode(tokens, out_ids, add_special_tokens) List[int]
|
||||
+decode(tokens, skip_special_tokens) str
|
||||
+apply_chat_template(messages, tokenize) Union[str, List[int]]
|
||||
+set_chat_template(template)
|
||||
+load(path)
|
||||
+from_pretrained(path) AutoTokenizer
|
||||
+save_pretrained(save_path)
|
||||
}
|
||||
|
||||
class ChatTemplate {
|
||||
+String template_str
|
||||
+render(messages, system_prompt, **extra_variables) str
|
||||
+from_string(template) ChatTemplate
|
||||
}
|
||||
}
|
||||
|
||||
namespace factory {
|
||||
class Registry {
|
||||
+Dict _entries
|
||||
+register(name, component_cls, category, priority)
|
||||
+get(name) Type
|
||||
+list_names() List[str]
|
||||
}
|
||||
|
||||
class BaseFactory {
|
||||
+Registry _registry
|
||||
+register(name, category, priority) decorator
|
||||
+create(name, *args, **kwargs) T
|
||||
+list_registered() list
|
||||
}
|
||||
}
|
||||
|
||||
namespace trainer {
|
||||
class Trainer {
|
||||
+TrainConfig train_config
|
||||
+List[TrainCallback] callbacks
|
||||
+train(checkpoint)
|
||||
+_build_context(checkpoint) TrainContext
|
||||
+_get_default_callbacks() List[TrainCallback]
|
||||
}
|
||||
|
||||
class TrainContext {
|
||||
+nn.Module model
|
||||
+BaseStrategy strategy
|
||||
+DataLoader dataloader
|
||||
+Optimizer optimizer
|
||||
+LRScheduler scheduler
|
||||
+Checkpoint checkpoint
|
||||
+int epoch
|
||||
+int iteration
|
||||
+float loss
|
||||
+int world_size
|
||||
+int rank
|
||||
}
|
||||
|
||||
class TrainContextBuilder {
|
||||
+TrainConfig config
|
||||
+with_checkpoint(checkpoint) TrainContextBuilder
|
||||
+build() TrainContext
|
||||
}
|
||||
|
||||
class BaseStrategy {
|
||||
+nn.Module model
|
||||
+str device
|
||||
+compute_loss(batch) Tensor
|
||||
}
|
||||
|
||||
class StrategyFactory {
|
||||
+Registry _registry
|
||||
+register(name) decorator
|
||||
+create(model, train_type, device, **kwargs) BaseStrategy
|
||||
}
|
||||
|
||||
class SEQStrategy {
|
||||
+float label_smoothing
|
||||
+compute_loss(batch) Tensor
|
||||
}
|
||||
|
||||
class SFTStrategy {
|
||||
+float label_smoothing
|
||||
+compute_loss(batch) Tensor
|
||||
}
|
||||
|
||||
class DPOStrategy {
|
||||
+nn.Module ref_model
|
||||
+float beta
|
||||
+str reduction
|
||||
+compute_loss(batch) Tensor
|
||||
}
|
||||
|
||||
class GRPOStrategy {
|
||||
+nn.Module ref_model
|
||||
+float clip_eps
|
||||
+float kl_coef
|
||||
+int group_size
|
||||
+compute_loss(batch) Tensor
|
||||
}
|
||||
|
||||
class BaseScheduler {
|
||||
+get_lr() List[float]
|
||||
+step()
|
||||
}
|
||||
|
||||
class SchedulerFactory {
|
||||
+Registry _registry
|
||||
+register(name) decorator
|
||||
+create(optimizer, schedule_type, **kwargs) BaseScheduler
|
||||
}
|
||||
|
||||
class CosineScheduler {
|
||||
+int warmup_steps
|
||||
+int lr_decay_steps
|
||||
+float min_rate
|
||||
}
|
||||
|
||||
class SGDRScheduler {
|
||||
+int warmup_steps
|
||||
+int cycle_length
|
||||
+float min_rate
|
||||
+int t_mult
|
||||
}
|
||||
|
||||
class TrainCallback {
|
||||
+on_train_begin(context)
|
||||
+on_train_end(context)
|
||||
+on_epoch_begin(context)
|
||||
+on_epoch_end(context)
|
||||
+on_step_begin(context)
|
||||
+on_step_end(context)
|
||||
+on_batch_begin(context)
|
||||
+on_batch_end(context)
|
||||
+on_error(context)
|
||||
}
|
||||
|
||||
class GradientClippingCallback {
|
||||
+float max_grad_norm
|
||||
+on_step_begin(context)
|
||||
}
|
||||
|
||||
class SchedulerCallback {
|
||||
+on_train_begin(context)
|
||||
+on_batch_end(context)
|
||||
}
|
||||
|
||||
class CheckpointCallback {
|
||||
+str save_dir
|
||||
+int interval
|
||||
+_save_checkpoint(context)
|
||||
+on_batch_end(context)
|
||||
+on_train_end(context)
|
||||
+on_error(context)
|
||||
}
|
||||
|
||||
class ProgressBarCallback {
|
||||
+int num_epoch
|
||||
+on_epoch_begin(context)
|
||||
+on_batch_end(context)
|
||||
+on_epoch_end(context)
|
||||
}
|
||||
|
||||
class MetricLoggerCallback {
|
||||
+str log_dir
|
||||
+int save_interval
|
||||
+on_batch_end(context)
|
||||
+on_train_end(context)
|
||||
}
|
||||
|
||||
class CallbackFactory {
|
||||
+Registry _registry
|
||||
+register(name) decorator
|
||||
+create(name, **kwargs) TrainCallback
|
||||
}
|
||||
}
|
||||
|
||||
namespace inference {
|
||||
class InferenceEngine {
|
||||
+nn.Module model
|
||||
+AutoTokenizer tokenizer
|
||||
+InferenceScheduler scheduler
|
||||
+int max_batch_size
|
||||
+Optional int max_seq_len
|
||||
+generate(prompt, stream, max_tokens, temperature, top_p, top_k) Union[Generator, str, List[str]]
|
||||
+generate_with_request(request) Union[Generator, str, List[str]]
|
||||
+generate_async(prompt, max_tokens, temperature, top_p, top_k) AsyncGenerator
|
||||
+get_stats() Dict
|
||||
+shutdown()
|
||||
}
|
||||
|
||||
class InferenceScheduler {
|
||||
+nn.Module model
|
||||
+AutoTokenizer tokenizer
|
||||
+PagedCache page_cache
|
||||
+int max_batch_size
|
||||
+int max_seq_len
|
||||
+int max_prompt_len
|
||||
+int page_size
|
||||
+List waiting_queue
|
||||
+List active_tasks
|
||||
+add_task(prompt, max_tokens, temperature, top_p, top_k, stream_callback) str
|
||||
+remove_task(task_id)
|
||||
+start()
|
||||
+stop()
|
||||
+get_stats() Dict
|
||||
}
|
||||
|
||||
class PagedCache {
|
||||
+int page_size
|
||||
+int _free_mask
|
||||
+List[int] _refs
|
||||
+Tensor k_cache
|
||||
+Tensor v_cache
|
||||
+alloc() int
|
||||
+alloc_n(n) List[int]
|
||||
+free(idx)
|
||||
+bind(page_table, total_len) CacheView
|
||||
+write(layer_id, page_table, start_pos, k, v)
|
||||
+gather(layer_id, page_table) Tuple[Tensor, Tensor]
|
||||
}
|
||||
|
||||
class CacheView {
|
||||
+PagedCache _cache
|
||||
+Tensor _page_table
|
||||
+int _total_len
|
||||
+write(layer_id, start_pos, k, v)
|
||||
+gather(layer_id) Tuple[Tensor, Tensor]
|
||||
}
|
||||
|
||||
class Task {
|
||||
+str task_id
|
||||
+List prompt_ids
|
||||
+int max_tokens
|
||||
+float temperature
|
||||
+float top_p
|
||||
+int top_k
|
||||
+TaskStatus status
|
||||
+List output_ids
|
||||
+int input_tokens
|
||||
+int output_tokens
|
||||
+List[int] page_table
|
||||
+int n_pages
|
||||
+float arrival_time
|
||||
+float finish_time
|
||||
+Callable stream_callback
|
||||
+int next_pos
|
||||
+is_finished(stop_ids) bool
|
||||
}
|
||||
|
||||
class TaskStatus {
|
||||
<<enumeration>>
|
||||
PENDING
|
||||
RUNNING
|
||||
FINISHED
|
||||
ABORTED
|
||||
}
|
||||
|
||||
class GenerationRequest {
|
||||
+List[Dict] messages
|
||||
+GenerationParams params
|
||||
+bool stream
|
||||
}
|
||||
|
||||
class GenerationParams {
|
||||
<<value object>>
|
||||
+int top_k
|
||||
+float top_p
|
||||
+float temperature
|
||||
+int max_tokens
|
||||
}
|
||||
|
||||
class BaseSamplingStrategy {
|
||||
<<abstract>>
|
||||
+apply(logits, filter_value) Tensor
|
||||
}
|
||||
|
||||
class TemperatureStrategy {
|
||||
+float temperature
|
||||
+apply(logits, filter_value) Tensor
|
||||
}
|
||||
|
||||
class TopKStrategy {
|
||||
+int top_k
|
||||
+apply(logits, filter_value) Tensor
|
||||
}
|
||||
|
||||
class TopPStrategy {
|
||||
+float top_p
|
||||
+apply(logits, filter_value) Tensor
|
||||
}
|
||||
|
||||
class SamplingPipeline {
|
||||
+List strategies
|
||||
+apply(logits, filter_value) Tensor
|
||||
+sample(logits, filter_value) Tensor
|
||||
}
|
||||
|
||||
class _Result {
|
||||
+List[str] tokens
|
||||
+List[str] results
|
||||
+List[bool] _done
|
||||
+append(token, idx)
|
||||
+get_results() List[str]
|
||||
+pop_all() List[str]
|
||||
+wait(timeout) bool
|
||||
}
|
||||
|
||||
class ChatMessage {
|
||||
+str role
|
||||
+str content
|
||||
}
|
||||
|
||||
class ChatCompletionRequest {
|
||||
+List[ChatMessage] messages
|
||||
+float temperature
|
||||
+float top_p
|
||||
+int top_k
|
||||
+int max_tokens
|
||||
+bool stream
|
||||
+Optional[str] stop
|
||||
+Optional[int] n
|
||||
}
|
||||
}
|
||||
|
||||
namespace parallel {
|
||||
class ParallelFunctions {
|
||||
+spawn_parallel_fn(fn, nprocs)
|
||||
+setup_parallel(rank, world_size, backend, master_addr, master_port, device_type)
|
||||
}
|
||||
|
||||
class ParallelModel {
|
||||
+dist.ProcessGroup process_group
|
||||
+int rank
|
||||
+int world_size
|
||||
}
|
||||
|
||||
class ColumnParallelLinear {
|
||||
+forward(x) Tensor
|
||||
}
|
||||
|
||||
class RowParallelLinear {
|
||||
+forward(x) Tensor
|
||||
}
|
||||
}
|
||||
|
||||
%% Relationships
|
||||
TrainConfig --> ModelConfig : uses
|
||||
TrainConfig --> BaseDataset : uses
|
||||
TrainConfig --> StrategyFactory : selects
|
||||
StrategyFactory ..> BaseStrategy : creates
|
||||
BaseStrategy <|-- SEQStrategy
|
||||
BaseStrategy <|-- SFTStrategy
|
||||
BaseStrategy <|-- DPOStrategy
|
||||
BaseStrategy <|-- GRPOStrategy
|
||||
DPOStrategy --> Transformer : uses
|
||||
GRPOStrategy --> Transformer : uses
|
||||
Trainer --> TrainConfig : configures
|
||||
Trainer --> TrainContextBuilder : builds
|
||||
Trainer --> TrainCallback : manages
|
||||
TrainContextBuilder --> TrainContext : creates
|
||||
Checkpoint ..> Checkpoint : saves/loads
|
||||
TrainContext --> Checkpoint : manages
|
||||
TrainContext --> BaseStrategy : uses
|
||||
TrainContext --> BaseScheduler : uses
|
||||
SchedulerFactory ..> BaseScheduler : creates
|
||||
BaseScheduler <|-- CosineScheduler
|
||||
BaseScheduler <|-- SGDRScheduler
|
||||
CallbackFactory ..> TrainCallback : creates
|
||||
TrainCallback <|-- GradientClippingCallback
|
||||
TrainCallback <|-- SchedulerCallback
|
||||
TrainCallback <|-- CheckpointCallback
|
||||
TrainCallback <|-- ProgressBarCallback
|
||||
TrainCallback <|-- MetricLoggerCallback
|
||||
InferenceEngine --> InferenceScheduler : uses
|
||||
InferenceEngine --> GenerationRequest : uses
|
||||
GenerationRequest --> GenerationParams : contains
|
||||
InferenceScheduler --> Task : manages
|
||||
Task --> TaskStatus : uses
|
||||
InferenceScheduler --> TaskStatus : uses
|
||||
InferenceScheduler --> PagedCache : uses
|
||||
InferenceScheduler --> Transformer : uses
|
||||
InferenceEngine --> Transformer : uses
|
||||
InferenceEngine --> _Result : uses
|
||||
BaseSamplingStrategy <|-- TemperatureStrategy
|
||||
BaseSamplingStrategy <|-- TopKStrategy
|
||||
BaseSamplingStrategy <|-- TopPStrategy
|
||||
SamplingPipeline --> BaseSamplingStrategy : composes
|
||||
BaseDataset <|-- SEQDataset
|
||||
BaseDataset <|-- SFTDataset
|
||||
BaseDataset <|-- DPODataset
|
||||
BaseDataset <|-- GRPODataset
|
||||
DatasetFactory ..> BaseDataset : creates
|
||||
MultiSegmentFetcher --> BaseSegmentFetcher : uses
|
||||
BaseDataset --> MultiSegmentFetcher : uses
|
||||
AutoModel <|-- Transformer
|
||||
AutoModel --> ModelConfig : contains
|
||||
Transformer --> DecoderBlock : uses
|
||||
Transformer --> RotaryEmbedding : uses
|
||||
Transformer --> Embedding : uses
|
||||
DecoderBlock --> GQA : uses
|
||||
DecoderBlock --> MLP : uses
|
||||
DecoderBlock --> RMSNorm : uses
|
||||
TrainContextBuilder --> ResumableDistributedSampler : creates
|
||||
ResumableDistributedSampler --> BaseDataset : samples
|
||||
ParallelModel <|-- RowParallelLinear
|
||||
ParallelModel <|-- ColumnParallelLinear
|
||||
AutoTokenizer --> ChatTemplate : uses
|
||||
TrainConfig --> DatasetFactory : selects
|
||||
TrainConfig --> SchedulerFactory : selects
|
||||
TrainConfig --> CallbackFactory : selects
|
||||
AutoModel ..> AutoTokenizer : loads with
|
||||
BaseFactory <|-- DatasetFactory
|
||||
BaseFactory <|-- StrategyFactory
|
||||
BaseFactory <|-- SchedulerFactory
|
||||
BaseFactory <|-- CallbackFactory
|
||||
```
|
||||
|
||||
### Module Overview
|
||||
|
||||
| Module | Components | Description |
|
||||
|--------|------------|-------------|
|
||||
| **astrai.config** | ModelConfig, TrainConfig | Configuration management |
|
||||
| **astrai.dataset** | BaseDataset, SEQDataset, SFTDataset, DPODataset, GRPODataset, BaseSegmentFetcher, MultiSegmentFetcher, ResumableDistributedSampler, DatasetFactory | Dataset loading and management |
|
||||
| **astrai.serialization** | Checkpoint, save_h5, load_h5 | Model serialization and checkpoint management |
|
||||
| **astrai.model** | AutoModel, Transformer, DecoderBlock, GQA, MLA, MLP, RMSNorm, Linear, RotaryEmbedding, Embedding | Neural network model |
|
||||
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
||||
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy, StrategyFactory, BaseScheduler, SchedulerFactory, TrainCallback, CallbackFactory | Training workflow management |
|
||||
| **astrai.inference** | InferenceEngine, InferenceScheduler, PagedCache, CacheView, Task, TaskStatus, GenerationParams, GenerationRequest, BaseSamplingStrategy, TemperatureStrategy, TopKStrategy, TopPStrategy, SamplingPipeline, ChatMessage, ChatCompletionRequest | Inference service with continuous batching and paged KV cache |
|
||||
| **astrai.parallel** | ParallelFunctions, ParallelModel, ColumnParallelLinear, RowParallelLinear | Distributed parallel |
|
||||
| **astrai.factory** | Registry, BaseFactory | Generic component registration |
|
||||
|
||||
### Design Patterns
|
||||
|
||||
| Pattern | Classes | Purpose |
|
||||
|---------|---------|---------|
|
||||
| **Strategy** | `BaseStrategy`, `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy`, `StrategyFactory` | Flexible training strategy switching, supports SEQ/SFT/DPO/GRPO |
|
||||
| **Builder** | `TrainContextBuilder` | Chain-building training context, step-by-step initialization of components |
|
||||
| **Factory** | `StrategyFactory`, `SchedulerFactory`, `DatasetFactory`, `CallbackFactory`, `BaseFactory` | Decorator registration mechanism, dynamically create training strategies, schedulers, datasets, and callbacks |
|
||||
| **Observer** | `TrainCallback`, `CallbackFactory` | Callback mechanism for training process monitoring (checkpoint, early stopping, metrics) |
|
||||
| **Context** | `TrainContext` | Training process state container with model, optimizer, scheduler and checkpoint |
|
||||
| **Registry** | `BaseFactory`, `Registry` | Generic component registration with category and priority support |
|
||||
| **Object Pool** | `PagedCache` | Page-based KV cache with O(1) alloc/free via bitmask |
|
||||
| **Strategy (Sampling)** | `BaseSamplingStrategy`, `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations with temperature, top-k, top-p |
|
||||
| **Producer-Consumer** | `InferenceScheduler`, `Task`, `waiting_queue`, `active_tasks` | Continuous batching with dynamic task queue management |
|
||||
| **Event-Driven** | `threading.Event`, `_task_event` | Non-blocking wait mechanism for task scheduling using Python's `threading` module |
|
||||
| **AutoModel Registry** | `AutoModel`, `Transformer` | Model type registration and dynamic loading via decorator pattern |
|
||||
| **Generator Pattern** | `_Result`, `GenerationRequest` | Event-based result notification for streaming/non-streaming generation |
|
||||
|
||||
### Core Relationships
|
||||
|
||||
1. **Configuration → Training**: `TrainConfig` contains `ModelConfig`, holds model, dataset, optimizer and other references
|
||||
2. **Training Flow**: `Trainer` → `TrainContextBuilder` → `TrainContext`, uses `BaseStrategy` to compute loss
|
||||
3. **Strategy Selection**: `StrategyFactory` creates corresponding strategy instance based on `train_type`
|
||||
4. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `Transformer`, uses `PagedCache` for paged KV cache management and `SamplingPipeline` for efficient continuous batching with streaming/non-streaming
|
||||
5. **Distributed Support**: `spawn_parallel_fn` and `setup_parallel` provide multi-process training capability for `Trainer`
|
||||
6. **Dataset Loading**: `DatasetFactory` creates datasets (SEQDataset, SFTDataset, DPODataset, GRPODataset), supports HDF5 loading via `BaseSegmentFetcher` and `MultiSegmentFetcher`
|
||||
7. **Checkpoint Management**: `Checkpoint` handles model state serialization/deserialization with safetensors
|
||||
8. **Scheduler Support**: `SchedulerFactory` creates learning rate schedulers (CosineScheduler, SGDRScheduler)
|
||||
9. **AutoModel Loading**: `AutoModel.from_pretrained()` dynamically loads model based on `config.json` model_type, uses `Registry` pattern for model type registration
|
||||
|
||||
## 3. Training Process
|
||||
|
||||
The common training process for large language models (LLM) typically includes three stages: **Pre-training (SEQ)**, **Supervised Fine-Tuning (SFT)**, and **Reinforcement Learning from Human Feedback (DPO/GRPO)**. This system is designed to support seamless end-to-end flow, achieving efficient switching and state management of different training stages through modular strategies.
|
||||
|
||||
### Core Formulas
|
||||
|
||||
**Pre-training (SEQ):**
|
||||
|
||||
$$
|
||||
L_{\text{PT}} = - \sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta)
|
||||
$$
|
||||
|
||||
**SFT:**
|
||||
|
||||
$$
|
||||
L_{\text{SFT}} = - \sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
|
||||
$$
|
||||
|
||||
**DPO:**
|
||||
|
||||
$$
|
||||
L_{\text{DPO}} = -\mathbb{E}_{(x, y_w, y_l) \sim D} \left[ \log \sigma\left( \beta \log \frac{\pi_\theta(y_w \mid x)}{\pi_{\text{ref}}(y_w \mid x)} - \beta \log \frac{\pi_\theta(y_l \mid x)}{\pi_{\text{ref}}(y_l \mid x)} \right) \right]
|
||||
$$
|
||||
|
||||
**GRPO:**
|
||||
|
||||
GRPO (Group Relative Policy Optimization) computes advantages from multiple responses to the same prompt, then optimizes using a PPO-style clipped objective:
|
||||
|
||||
$$
|
||||
\text{Advantage}_i = \frac{r_i - \mu}{\sigma + \epsilon}
|
||||
$$
|
||||
|
||||
Where $r_i$ is the reward for the $i$-th response, $\mu$ and $\sigma$ are the mean and standard deviation of group rewards.
|
||||
|
||||
$$
|
||||
L_{\text{GRPO}} = -\mathbb{E} \left[ \min\left( \frac{\pi_\theta(a|s)}{\pi_{\text{ref}}(a|s)} \cdot A, \text{clip}\left(\frac{\pi_\theta(a|s)}{\pi_{\text{ref}}(a|s)}, 1-\epsilon, 1+\epsilon\right) \cdot A \right) \right] + \lambda \cdot D_{KL}
|
||||
$$
|
||||
|
||||
The KL divergence term uses mean squared error approximation:
|
||||
|
||||
$$
|
||||
L_{KL} = \lambda \cdot \mathbb{E} \left[ (\log \pi_\theta - \log \pi_{\text{ref}})^2 \right]
|
||||
$$
|
||||
|
||||
The final loss is the sum of both: $L = L_{\text{policy}} + L_{KL}$
|
||||
|
||||
Through the above three-stage progressive training, the model completes its evolution from a general language foundation to a specialized, highly-aligned dialogue intelligence.
|
||||
|
||||
> Document Update Time: 2026-04-09
|
||||
@@ -1,334 +0,0 @@
|
||||
## Model Introduction
|
||||
|
||||
### 1. Model Architecture
|
||||
|
||||
This model uses the Transformer architecture with GQA mechanism (q_head=24, kv_head=4), which saves KV cache memory compared to traditional MHA. The model is built by stacking 24 layers of Transformer blocks, with 1.0 billion parameters. Transformer is an autoregressive model that calculates the relationship between all previous tokens to obtain the probability distribution of the next token.
|
||||
|
||||
The model now uses the **AutoModel** base class for flexible loading and saving:
|
||||
|
||||
```python
|
||||
from astrai.model import AutoModel
|
||||
|
||||
# Load model from checkpoint
|
||||
model = AutoModel.from_pretrained("path/to/model")
|
||||
|
||||
# Save model to new directory
|
||||
model.save_pretrained("path/to/save")
|
||||
```
|
||||
|
||||
The Transformer model is registered via `@AutoModel.register('transformer')` decorator, allowing easy extension for new model types.
|
||||
|
||||
```mermaid
|
||||
flowchart TB
|
||||
subgraph Layers["Transformer Layers"]
|
||||
direction TB
|
||||
A[Input Embedding] --> B[Transformer Block\nLayer 1]
|
||||
B --> C[Transformer Block\nLayer ...]
|
||||
C --> D[Transformer Block\nLayer 32]
|
||||
D --> E[RMSNorm]
|
||||
E --> F[Linear]
|
||||
F --> G[SoftMax]
|
||||
end
|
||||
|
||||
subgraph TransformerBlock["Transformer Block"]
|
||||
direction TB
|
||||
H[x] --> I[RMSNorm]
|
||||
I --> J[Linear → Q/K/V]
|
||||
J --> K[Q]
|
||||
J --> L[K]
|
||||
J --> M[V]
|
||||
K --> N[RoPE]
|
||||
L --> O[RoPE]
|
||||
N --> P["Q @ K^T / sqrt(d)"]
|
||||
O --> P
|
||||
P --> Q[Masked SoftMax]
|
||||
Q --> R[S @ V]
|
||||
M --> R
|
||||
R --> S[Linear]
|
||||
S --> T[+]
|
||||
H --> T
|
||||
T --> U[RMSNorm]
|
||||
U --> V["Linear (gate)"]
|
||||
U --> W["Linear (up)"]
|
||||
V --> X[SiLU]
|
||||
X --> Y[×]
|
||||
W --> Y
|
||||
Y --> Z["Linear (down)"]
|
||||
Z --> AA[+]
|
||||
T --> AA
|
||||
AA --> BB[x']
|
||||
end
|
||||
|
||||
classDef main fill:#e6f3ff,stroke:#0066cc;
|
||||
classDef block fill:#fff2e6,stroke:#cc6600;
|
||||
class Layers main;
|
||||
class TransformerBlock block;
|
||||
```
|
||||
|
||||
What is an autoregressive model? After splitting a sentence into tokens, the model predicts the probability distribution of the next token. This means the model calculates the probability of the next possible token and its corresponding probability based on the given context (the sequence of tokens that have already appeared).
|
||||
|
||||
#### 1. Autoregression
|
||||
|
||||
In autoregressive modeling, when a sentence is tokenized into a sequence of tokens, the model learns to predict what comes next. Given a sequence of tokens as input, the model calculates a probability distribution over all possible next tokens. This distribution tells us how likely each potential next token is, given the current context.
|
||||
|
||||
For instance, if the input sequence contains tokens representing a question, the model might predict that certain response tokens have higher probabilities than others. The sampling process then selects one token from this distribution—controlled by parameters like top_k, top_p, and temperature—to serve as the next token in the sequence.
|
||||
|
||||
Once a token is selected, it is appended to the input sequence, and the model repeats this process. The updated sequence is then fed back into the model to predict the next token. This iterative process continues until either a special end-of-sequence token is generated, or the maximum sequence length is reached. These control tokens are essential because without them, the model would continue generating tokens indefinitely, eventually exhausting available memory.
|
||||
|
||||
#### 2. Causal Mask
|
||||
|
||||
Transformers use attention mechanism. The input shape is generally [bsz, seq_len], and the output is [bsz, seq_len, n_dim]. To predict the next token, the model's input and output must be offset by one position. The target predicted by the model must be offset by one position, and during training we also use the offset-by-one method:
|
||||
|
||||
```
|
||||
sequence : [[1, 2, 3, 4, 5, 6]]
|
||||
input_ids: [[1, 2, 3, 4, 5]]
|
||||
target_ids: [[2, 3, 4, 5, 6]]
|
||||
```
|
||||
|
||||
The attention score calculation formula is:
|
||||
|
||||
$$ s_{ij} = softmax(\frac{q_i^Tk_j}{\sqrt{d_k}}) $$
|
||||
$$ s_{ij} := s_{ij} + mask_{ij} $$
|
||||
|
||||
Here, the attention score represents the degree to which the model attends to the similarity between two tokens.
|
||||
|
||||
For decoder-only structure models, to prevent the model from "stealing" information from future positions, a mask needs to be added during attention calculation. We need to apply a mask before attention score calculation. This mask is typically a lower triangular matrix, and for a sequence of length n, its shape is [n, n]. Below is an example of how to create such a causal mask matrix for a sequence of length 5:
|
||||
|
||||
```
|
||||
[[0, -inf, -inf, -inf, -inf],
|
||||
[0, 0, -inf, -inf, -inf],
|
||||
[0, 0, 0, -inf, -inf],
|
||||
[0, 0, 0, 0, -inf],
|
||||
[0, 0, 0, 0, 0]]
|
||||
```
|
||||
|
||||
In this matrix, 0 represents positions that can be attended to, while -inf represents positions that should be masked (i.e., should not be attended to). Because this matrix ensures that after the softmax, the parts of the attention scores where $j > i$ change from `inf` to 0, meaning the model cannot see future information.
|
||||
|
||||
#### 3. Rotary Position Embedding
|
||||
|
||||
Rotary Position Embedding (RoPE) is a position encoding method designed to solve the problem of lacking direct modeling of sequence position information in Transformer models. Unlike traditional position encodings (such as sine and cosine function position encodings), RoPE embeds position information directly into the Query (Q) and Key (K) vectors, allowing the model to more naturally handle relative position relationships in sequences.
|
||||
|
||||
$$ q_i = R_i W_q x_i $$
|
||||
$$ k_j = R_j W_k x_j $$
|
||||
$$ q_i^T k_j = (R_i W_q x_i)^T( R_j W_k x_j) = x_i^T W_q^T R_{i-j} W_k x_j $$
|
||||
|
||||
The $R_{i-j}$ controls the attenuation of attention for different tokens at different relative distances. When the absolute value of $i - j$ is larger, the degree of attenuation is stronger. This approach allows the model to learn relative position relationships, enabling the model to scale and adapt to longer sequences.
|
||||
|
||||
## KV Cache Implementation
|
||||
|
||||
According to the attention calculation formula:
|
||||
|
||||
$$
|
||||
\begin{align*}
|
||||
o_i &= \sum_j s_{ij} v_{j} \newline
|
||||
s_{ij} &= \text{softmax}\left( \frac{q_{i} k_{j}}{\sqrt{d_k}} \right)
|
||||
\end{align*}
|
||||
$$
|
||||
|
||||
Since the model is an autoregressive model, we only need to calculate for the last part of the sequence, meaning the index $i$ is fixed as the last element of the sequence, and we compute $o_{n}$:
|
||||
|
||||
$$
|
||||
\begin{align*}
|
||||
o_n &= \sum_j s_{j}v_{j} \newline
|
||||
s_j &= \text{softmax}\left(\frac{q_n k_{j}}{\sqrt{d_k}} \right)
|
||||
\end{align*}
|
||||
$$
|
||||
|
||||
If we expand the expression:
|
||||
|
||||
$$
|
||||
o_n = \sum_j \text{softmax}\left(\frac{q_n k_{j}}{\sqrt{d_k}}\right)v_{j}
|
||||
$$
|
||||
|
||||
In the above expression, only k and v have length indices, while $q$ does not. Therefore, during the calculation process, the input of $q$ is fixed as the last token from the previous input, while $k$ and $v$ need to be cached for parts of different lengths. Also, when caching, note that position encoding calculation should be performed before KV cache computation, otherwise there will be position encoding calculation errors.
|
||||
|
||||
### 4. AutoModel Loading
|
||||
|
||||
The project now uses the **AutoModel** base class for flexible model loading and saving:
|
||||
|
||||
```python
|
||||
from astrai.model import AutoModel
|
||||
|
||||
# Load model from checkpoint
|
||||
model = AutoModel.from_pretrained("path/to/model")
|
||||
|
||||
# Save model to new directory
|
||||
model.save_pretrained("path/to/save")
|
||||
```
|
||||
|
||||
The Transformer model is registered via `@AutoModel.register('transformer')` decorator, allowing easy extension for new model types. The `from_pretrained` method automatically loads the `config.json` to determine the model type and uses safetensors format for weights.
|
||||
|
||||
### 5. Continuous Batching Inference
|
||||
|
||||
The inference engine supports **continuous batching** for efficient batch processing:
|
||||
|
||||
```python
|
||||
from astrai.inference import InferenceEngine, GenerationRequest
|
||||
|
||||
# Create inference engine with continuous batching
|
||||
engine = InferenceEngine(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
)
|
||||
|
||||
# Use GenerationRequest with messages format
|
||||
request = GenerationRequest(
|
||||
messages=[
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": "Hello"},
|
||||
],
|
||||
temperature=0.8,
|
||||
top_p=0.95,
|
||||
top_k=50,
|
||||
max_len=1024,
|
||||
stream=True,
|
||||
)
|
||||
|
||||
# Generate with streaming
|
||||
for token in engine.generate_with_request(request):
|
||||
print(token, end="", flush=True)
|
||||
```
|
||||
|
||||
The continuous batching feature allows dynamic batch composition where new requests can join at any time and completed requests are released immediately.
|
||||
|
||||
## HTTP API Usage
|
||||
|
||||
The inference server provides HTTP endpoints for remote inference. Start the server first:
|
||||
|
||||
```bash
|
||||
python -m scripts.tools.server --port 8000
|
||||
```
|
||||
|
||||
### OpenAI-Compatible Endpoint
|
||||
|
||||
The server provides an OpenAI-compatible chat completion endpoint at `/v1/chat/completions`:
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": "Hello, how are you?"}
|
||||
],
|
||||
"temperature": 0.8,
|
||||
"max_tokens": 2048,
|
||||
"stream": false
|
||||
}'
|
||||
```
|
||||
|
||||
**Request Parameters:**
|
||||
| Parameter | Type | Default | Description |
|
||||
|-----------|------|---------|-------------|
|
||||
| `messages` | List[dict] | Required | Chat messages with role and content |
|
||||
| `temperature` | float | 1.0 | Sampling temperature (0.0-2.0) |
|
||||
| `top_p` | float | 1.0 | Nucleus sampling threshold |
|
||||
| `top_k` | int | 50 | Top-k sampling parameter |
|
||||
| `max_tokens` | int | 1024 | Maximum tokens to generate |
|
||||
| `stream` | bool | false | Enable streaming response |
|
||||
|
||||
**Response (non-streaming):**
|
||||
```json
|
||||
{
|
||||
"id": "chatcmpl-1234567890",
|
||||
"object": "chat.completion",
|
||||
"created": 1234567890,
|
||||
"model": "astrai",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "Hello! I'm doing well..."},
|
||||
"finish_reason": "stop"
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 20,
|
||||
"completion_tokens": 15,
|
||||
"total_tokens": 35
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### Streaming Response
|
||||
|
||||
Enable streaming for real-time token-by-token output:
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"messages": [{"role": "user", "content": "Write a story"}],
|
||||
"stream": true,
|
||||
"max_tokens": 500
|
||||
}'
|
||||
```
|
||||
|
||||
The server uses Server-Sent Events (SSE) with content type `text/event-stream`.
|
||||
|
||||
### Anthropic-Compatible Endpoint
|
||||
|
||||
The server also provides an Anthropic-compatible endpoint at `/v1/messages`:
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:8000/v1/messages \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "astrai",
|
||||
"system": "You are a helpful assistant.",
|
||||
"messages": [{"role": "user", "content": "Hello, how are you?"}],
|
||||
"max_tokens": 2048
|
||||
}'
|
||||
```
|
||||
|
||||
Response:
|
||||
```json
|
||||
{
|
||||
"id": "msg_abc123...",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "astrai",
|
||||
"content": [{"type": "text", "text": "Hello! I am doing well..."}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": null,
|
||||
"usage": {"input_tokens": 20, "output_tokens": 15}
|
||||
}
|
||||
```
|
||||
|
||||
Streaming:
|
||||
```bash
|
||||
curl -X POST http://localhost:8000/v1/messages \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "astrai",
|
||||
"system": "You are a helpful assistant.",
|
||||
"messages": [{"role": "user", "content": "Write a short poem"}],
|
||||
"max_tokens": 500,
|
||||
"stream": true
|
||||
}'
|
||||
```
|
||||
|
||||
Supports `stop_sequences` for early termination:
|
||||
```bash
|
||||
curl -X POST http://localhost:8000/v1/messages \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "astrai",
|
||||
"messages": [{"role": "user", "content": "Write a story"}],
|
||||
"max_tokens": 500,
|
||||
"stop_sequences": ["The end", "THE END"]
|
||||
}'
|
||||
```
|
||||
|
||||
### Health Check
|
||||
|
||||
Monitor server and model status:
|
||||
|
||||
```bash
|
||||
curl http://localhost:8000/health
|
||||
# {"status": "ok", "model_loaded": true}
|
||||
|
||||
curl http://localhost:8000/stats
|
||||
# {"total_tasks": 10, "total_tokens": 5000, "active_tasks": 1, "waiting_queue": 0}
|
||||
```
|
||||
|
||||
> Document Update Time: 2026-04-09
|
||||
@@ -1,158 +0,0 @@
|
||||
# Parameter Documentation
|
||||
|
||||
## Training Parameters
|
||||
|
||||
### Basic Parameters
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`) | required |
|
||||
| `--data_root_path` | Dataset root directory | required |
|
||||
| `--param_path` | Model parameters or checkpoint path | required |
|
||||
| `--n_epoch` | Total training epochs | 1 |
|
||||
| `--batch_size` | Batch size | 1 |
|
||||
| `--accumulation_steps` | Gradient accumulation steps between optimizer steps | 1 |
|
||||
|
||||
### Learning Rate Scheduling
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--warmup_steps` | Warmup steps | 1000 |
|
||||
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
||||
| `--max_grad_norm` | Maximum gradient norm for clipping | 1.0 |
|
||||
|
||||
### Optimizer (AdamW)
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--adamw_beta1` | AdamW beta1 | 0.9 |
|
||||
| `--adamw_beta2` | AdamW beta2 | 0.95 |
|
||||
| `--adamw_weight_decay` | AdamW weight decay | 0.01 |
|
||||
|
||||
### Data Loading
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--window_size` | Max input sequence length | model config `max_len` |
|
||||
| `--stride` | Stride for sliding window over sequences | None |
|
||||
| `--random_seed` | Random seed for reproducibility | 3407 |
|
||||
| `--num_workers` | DataLoader worker processes | 4 |
|
||||
| `--no_pin_memory` | Disable pin_memory (enabled by default) | (flag) |
|
||||
|
||||
### Checkpoint & Resume
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--ckpt_interval` | Iterations between checkpoints | 5000 |
|
||||
| `--ckpt_dir` | Checkpoint save directory | checkpoint |
|
||||
| `--start_epoch` | Resume from epoch (0 = from scratch) | 0 |
|
||||
| `--start_batch` | Resume from batch iteration | 0 |
|
||||
|
||||
### Distributed Training
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--nprocs` | Number of GPUs / processes | 1 |
|
||||
| `--device_type` | Device type | cuda |
|
||||
|
||||
### Strategy-specific
|
||||
|
||||
| Parameter | Description | Default | Used by |
|
||||
|-----------|-------------|---------|---------|
|
||||
| `--dpo_beta` | DPO beta value | 0.1 | `dpo` |
|
||||
| `--label_smoothing` | Label smoothing for cross-entropy loss | 0.1 | `seq`, `sft` |
|
||||
| `--group_size` | GRPO group size | 4 | `grpo` |
|
||||
| `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo` |
|
||||
| `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo` |
|
||||
| `--grpo_sync_interval` | GRPO ref_model sync interval (steps) | 200 | `grpo` |
|
||||
|
||||
### Usage Example
|
||||
|
||||
```bash
|
||||
python scripts/tools/train.py \
|
||||
--train_type seq \
|
||||
--data_root_path /path/to/dataset \
|
||||
--param_path /path/to/model \
|
||||
--n_epoch 3 \
|
||||
--batch_size 4 \
|
||||
--accumulation_steps 8 \
|
||||
--max_lr 3e-4 \
|
||||
--warmup_steps 2000 \
|
||||
--max_grad_norm 1.0 \
|
||||
--ckpt_interval 5000 \
|
||||
--ckpt_dir ./checkpoints \
|
||||
--num_workers 4 \
|
||||
--nprocs 1 \
|
||||
--device_type cuda
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Generation Parameters
|
||||
|
||||
### GenerationRequest Parameters
|
||||
|
||||
| Parameter | Description | Default Value |
|
||||
|-----------|-------------|---------------|
|
||||
| `messages` | List of message dictionaries (role, content) | required |
|
||||
| `temperature` | Sampling temperature (higher = more random) | 1.0 |
|
||||
| `top_p` | Nucleus sampling threshold | 1.0 |
|
||||
| `top_k` | Top-k sampling count | 50 |
|
||||
| `max_len` | Maximum generation length | 1024 |
|
||||
| `stream` | Whether to stream output | False |
|
||||
|
||||
### Usage Example
|
||||
|
||||
```python
|
||||
import torch
|
||||
from astrai.model import AutoModel
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
from astrai.inference import InferenceEngine, GenerationRequest
|
||||
|
||||
# Load model using AutoModel
|
||||
model = AutoModel.from_pretrained("your_model_dir")
|
||||
|
||||
# Load tokenizer
|
||||
tokenizer = AutoTokenizer.from_pretrained("your_model_dir")
|
||||
|
||||
# Create engine with separate model and tokenizer
|
||||
engine = InferenceEngine(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
)
|
||||
|
||||
# Build request with messages format
|
||||
request = GenerationRequest(
|
||||
messages=[
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": "Hello"},
|
||||
],
|
||||
temperature=0.8,
|
||||
top_p=0.95,
|
||||
top_k=50,
|
||||
max_len=1024,
|
||||
)
|
||||
|
||||
# Generate (streaming)
|
||||
for token in engine.generate_with_request(request):
|
||||
print(token, end="", flush=True)
|
||||
|
||||
# Or use simple generate interface
|
||||
result = engine.generate(
|
||||
prompt="Hello",
|
||||
stream=False,
|
||||
max_tokens=1024,
|
||||
temperature=0.8,
|
||||
top_p=0.95,
|
||||
top_k=50,
|
||||
)
|
||||
```
|
||||
|
||||
### Generation Modes
|
||||
|
||||
| Mode | Description |
|
||||
|------|-------------|
|
||||
| `stream=True` | Streaming output, yields token by token |
|
||||
| `stream=False` | Non-streaming output, returns complete result |
|
||||
|
||||
> Document Update Time: 2026-04-09
|
||||
+87
-19
@@ -1,32 +1,100 @@
|
||||
__version__ = "1.3.4"
|
||||
__version__ = "1.3.13"
|
||||
__author__ = "ViperEkura"
|
||||
|
||||
from astrai.config import (
|
||||
ModelConfig,
|
||||
AutoRegressiveLMConfig,
|
||||
BaseModelConfig,
|
||||
ConfigFactory,
|
||||
EncoderConfig,
|
||||
PipelineConfig,
|
||||
TrainConfig,
|
||||
)
|
||||
from astrai.dataset import DatasetFactory
|
||||
from astrai.dataset import (
|
||||
BaseDataset,
|
||||
DatasetFactory,
|
||||
RDSampler,
|
||||
Store,
|
||||
StoreFactory,
|
||||
)
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.inference import (
|
||||
GenerationRequest,
|
||||
InferenceEngine,
|
||||
ProtocolHandler,
|
||||
SamplingPipeline,
|
||||
get_app,
|
||||
run_server,
|
||||
sample,
|
||||
)
|
||||
from astrai.logging import setup_logging
|
||||
from astrai.model import (
|
||||
AutoModel,
|
||||
AutoRegressiveLM,
|
||||
EmbeddingEncoder,
|
||||
LoRAConfig,
|
||||
inject_lora,
|
||||
)
|
||||
from astrai.parallel import (
|
||||
ExecutorFactory,
|
||||
get_rank,
|
||||
get_world_size,
|
||||
only_on_rank,
|
||||
spawn_parallel_fn,
|
||||
)
|
||||
from astrai.preprocessing import Pipeline, filter_by_length
|
||||
from astrai.serialization import Checkpoint
|
||||
from astrai.tokenize import AutoTokenizer, ChatTemplate
|
||||
from astrai.trainer import (
|
||||
BaseScheduler,
|
||||
BaseStrategy,
|
||||
CallbackFactory,
|
||||
SchedulerFactory,
|
||||
StrategyFactory,
|
||||
TrainCallback,
|
||||
Trainer,
|
||||
)
|
||||
from astrai.model import AutoModel, Transformer
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
from astrai.trainer import CallbackFactory, SchedulerFactory, StrategyFactory, Trainer
|
||||
|
||||
__all__ = [
|
||||
"Transformer",
|
||||
"ModelConfig",
|
||||
"TrainConfig",
|
||||
"DatasetFactory",
|
||||
"AutoTokenizer",
|
||||
"GenerationRequest",
|
||||
"InferenceEngine",
|
||||
"Trainer",
|
||||
"CallbackFactory",
|
||||
"StrategyFactory",
|
||||
"SchedulerFactory",
|
||||
"BaseFactory",
|
||||
"AutoRegressiveLM",
|
||||
"AutoRegressiveLMConfig",
|
||||
"AutoModel",
|
||||
"AutoTokenizer",
|
||||
"BaseDataset",
|
||||
"BaseFactory",
|
||||
"BaseModelConfig",
|
||||
"BaseScheduler",
|
||||
"BaseStrategy",
|
||||
"CallbackFactory",
|
||||
"ChatTemplate",
|
||||
"Checkpoint",
|
||||
"ConfigFactory",
|
||||
"DatasetFactory",
|
||||
"EmbeddingEncoder",
|
||||
"EncoderConfig",
|
||||
"ExecutorFactory",
|
||||
"InferenceEngine",
|
||||
"LoRAConfig",
|
||||
"Pipeline",
|
||||
"PipelineConfig",
|
||||
"ProtocolHandler",
|
||||
"RDSampler",
|
||||
"SamplingPipeline",
|
||||
"SchedulerFactory",
|
||||
"Store",
|
||||
"StoreFactory",
|
||||
"StrategyFactory",
|
||||
"TrainCallback",
|
||||
"TrainConfig",
|
||||
"Trainer",
|
||||
"filter_by_length",
|
||||
"get_app",
|
||||
"get_rank",
|
||||
"get_world_size",
|
||||
"inject_lora",
|
||||
"only_on_rank",
|
||||
"run_server",
|
||||
"sample",
|
||||
"setup_logging",
|
||||
"spawn_parallel_fn",
|
||||
]
|
||||
|
||||
setup_logging()
|
||||
|
||||
@@ -1,8 +1,25 @@
|
||||
from astrai.config.model_config import ModelConfig
|
||||
from astrai.config.model_config import (
|
||||
AutoRegressiveLMConfig,
|
||||
BaseModelConfig,
|
||||
ConfigFactory,
|
||||
EncoderConfig,
|
||||
)
|
||||
from astrai.config.preprocess_config import (
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
PipelineConfig,
|
||||
ProcessingConfig,
|
||||
)
|
||||
from astrai.config.train_config import TrainConfig
|
||||
|
||||
__all__ = [
|
||||
# Model configuration
|
||||
"ModelConfig",
|
||||
"BaseModelConfig",
|
||||
"AutoRegressiveLMConfig",
|
||||
"EncoderConfig",
|
||||
"ConfigFactory",
|
||||
"TrainConfig",
|
||||
"InputConfig",
|
||||
"OutputConfig",
|
||||
"PipelineConfig",
|
||||
"ProcessingConfig",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
import json
|
||||
from dataclasses import asdict
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Self, Union
|
||||
|
||||
from pydantic import ConfigDict
|
||||
from pydantic.dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass(config=ConfigDict(use_attribute_docstrings=True))
|
||||
class BaseConfig:
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
result = {}
|
||||
for k, v in asdict(self).items():
|
||||
if isinstance(v, tuple):
|
||||
v = list(v)
|
||||
try:
|
||||
json.dumps(v)
|
||||
result[k] = v
|
||||
except (TypeError, ValueError):
|
||||
# Skip non-serializable runtime objects (e.g. model_fn, dataset).
|
||||
# TrainConfig mixes hyperparams with callables/datasets; only the
|
||||
# JSON-serializable subset is written to checkpoint meta.
|
||||
pass
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, d: Dict[str, Any]) -> Self:
|
||||
return cls(**d)
|
||||
|
||||
@classmethod
|
||||
def from_file(cls, path: Union[str, Path]) -> Self:
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
return cls.from_dict(json.load(f))
|
||||
|
||||
def to_file(self, path: Union[str, Path]):
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
json.dump(self.to_dict(), f, indent=2, ensure_ascii=False)
|
||||
+166
-30
@@ -1,42 +1,178 @@
|
||||
import json
|
||||
from dataclasses import asdict, dataclass
|
||||
from typing import Optional, Self
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from pydantic import field_validator
|
||||
from pydantic.dataclasses import dataclass
|
||||
|
||||
from astrai.config.base import BaseConfig
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
_ATTN_TYPES = frozenset({"gqa", "mla"})
|
||||
_FFN_TYPES = frozenset({"mlp", "moe"})
|
||||
|
||||
|
||||
class ConfigFactory(BaseFactory[BaseConfig]):
|
||||
"""Factory that dispatches config classes by ``model_type``."""
|
||||
|
||||
@classmethod
|
||||
def load(cls, raw: Dict[str, Any]) -> BaseConfig:
|
||||
model_type = raw.get("model_type") or "autoregressive_lm"
|
||||
config_cls = cls.get_component_class(model_type)
|
||||
return config_cls.from_dict(raw)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelConfig:
|
||||
# basic config
|
||||
class BaseModelConfig(BaseConfig):
|
||||
"""Base config with ``model_type`` dispatch and file I/O.
|
||||
|
||||
Args:
|
||||
model_type (Optional[str]): Model type identifier for AutoModel dispatch. Defaults to None.
|
||||
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
|
||||
"""
|
||||
|
||||
model_type: Optional[str] = None
|
||||
neftune_alpha: float = 0.0
|
||||
|
||||
|
||||
@dataclass
|
||||
@ConfigFactory.register("autoregressive_lm")
|
||||
class AutoRegressiveLMConfig(BaseModelConfig):
|
||||
"""Configuration for autoregressive language model.
|
||||
|
||||
Args:
|
||||
model_type (Optional[str]): Model type identifier for AutoModel dispatch. Defaults to None.
|
||||
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
|
||||
vocab_size (Optional[int]): Vocabulary size. Defaults to None.
|
||||
hidden_size (Optional[int]): Hidden dimension size. Defaults to None.
|
||||
num_hidden_layers (Optional[int]): Number of transformer layers. Defaults to None.
|
||||
rms_norm_eps (Optional[float]): Epsilon for RMSNorm. Defaults to None.
|
||||
intermediate_size (Optional[int]): Intermediate size in FFN. Defaults to None.
|
||||
tie_word_embeddings (Optional[bool]): Whether to tie embedding and lm_head weights. Defaults to None.
|
||||
max_position_embeddings (Optional[int]): Maximum sequence length the model was trained with. Defaults to None.
|
||||
rope_theta (Optional[float]): Base frequency for RoPE. Defaults to None.
|
||||
rope_scaling (Optional[dict]): RoPE scaling config, e.g. {"type": "linear", "factor": 4.0}. Defaults to None.
|
||||
attn_type (str): Attention type: 'gqa' or 'mla'. Defaults to "gqa".
|
||||
num_attention_heads (Optional[int]): Number of query attention heads. Defaults to None.
|
||||
num_key_value_heads (Optional[int]): Number of key/value heads for GQA. Defaults to None.
|
||||
use_qk_norm (Optional[bool]): Whether to apply RMSNorm to Q/K. Defaults to None.
|
||||
use_gated_attention (Optional[bool]): Whether to use gated attention. Defaults to None.
|
||||
kv_lora_rank (Optional[int]): KV compression rank, MLA only. Defaults to None.
|
||||
qk_nope_head_dim (Optional[int]): Non-RoPE head dimension, MLA only. Defaults to None.
|
||||
qk_rope_head_dim (Optional[int]): RoPE head dimension, MLA only. Defaults to None.
|
||||
ffn_type (str): FFN type: 'mlp' or 'moe'. Defaults to "mlp".
|
||||
n_routed_experts (Optional[int]): Number of routed experts, MoE only. Defaults to None.
|
||||
n_shared_experts (Optional[int]): Number of shared experts, MoE only. Defaults to None.
|
||||
n_activated_experts (Optional[int]): Number of activated experts per token, MoE only. Defaults to None.
|
||||
topk_method (Optional[str]): Top-k routing method, MoE only. Defaults to None.
|
||||
moe_intermediate_size (Optional[int]): Expert hidden dim, defaults to intermediate_size if None. MoE only.
|
||||
shared_expert_intermediate_size (Optional[int]): Shared expert hidden dim, defaults to intermediate_size if None. MoE only.
|
||||
norm_topk_prob (bool): Normalize top-k routing probabilities. Defaults to True.
|
||||
decoder_sparse_step (int): Frequency of MoE layers, 1=every layer. Defaults to 1.
|
||||
mlp_only_layers (Optional[list[int]]): Layer indices using dense MLP instead of MoE. Defaults to None.
|
||||
"""
|
||||
|
||||
vocab_size: Optional[int] = None
|
||||
dim: Optional[int] = None
|
||||
|
||||
n_layers: Optional[int] = None
|
||||
norm_eps: Optional[float] = None
|
||||
dim_ffn: Optional[int] = None
|
||||
tie_weight: Optional[bool] = None
|
||||
|
||||
# RoPE
|
||||
max_len: Optional[int] = None
|
||||
hidden_size: Optional[int] = None
|
||||
num_hidden_layers: Optional[int] = None
|
||||
rms_norm_eps: Optional[float] = None
|
||||
intermediate_size: Optional[int] = None
|
||||
tie_word_embeddings: Optional[bool] = None
|
||||
max_position_embeddings: Optional[int] = None
|
||||
rope_theta: Optional[float] = None
|
||||
|
||||
# GQA
|
||||
n_heads: Optional[int] = None
|
||||
n_kv_heads: Optional[int] = None
|
||||
rope_scaling: Optional[dict] = None
|
||||
attn_type: str = "gqa"
|
||||
num_attention_heads: Optional[int] = None
|
||||
num_key_value_heads: Optional[int] = None
|
||||
use_qk_norm: Optional[bool] = None
|
||||
use_gated_attention: Optional[bool] = None
|
||||
kv_lora_rank: Optional[int] = None
|
||||
qk_nope_head_dim: Optional[int] = None
|
||||
qk_rope_head_dim: Optional[int] = None
|
||||
ffn_type: str = "mlp"
|
||||
n_routed_experts: Optional[int] = None
|
||||
n_shared_experts: Optional[int] = None
|
||||
n_activated_experts: Optional[int] = None
|
||||
topk_method: Optional[str] = None
|
||||
moe_intermediate_size: Optional[int] = None
|
||||
shared_expert_intermediate_size: Optional[int] = None
|
||||
norm_topk_prob: bool = True
|
||||
decoder_sparse_step: int = 1
|
||||
mlp_only_layers: Optional[list[int]] = None
|
||||
moe_aux_loss_coef: float = 0.01
|
||||
|
||||
def load(self, config_path: str) -> Self:
|
||||
config = {}
|
||||
with open(config_path, "r") as f:
|
||||
config.update(json.load(f))
|
||||
@field_validator("attn_type")
|
||||
def _validate_attn_type(cls, v: str) -> str:
|
||||
if v not in _ATTN_TYPES:
|
||||
raise ValueError(
|
||||
f"attn_type must be one of {sorted(_ATTN_TYPES)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
for key, value in config.items():
|
||||
if hasattr(self, key):
|
||||
setattr(self, key, value)
|
||||
@field_validator("ffn_type")
|
||||
def _validate_ffn_type(cls, v: str) -> str:
|
||||
if v not in _FFN_TYPES:
|
||||
raise ValueError(f"ffn_type must be one of {sorted(_FFN_TYPES)}, got {v!r}")
|
||||
return v
|
||||
|
||||
return self
|
||||
@field_validator("decoder_sparse_step")
|
||||
def _validate_decoder_sparse_step(cls, v: int) -> int:
|
||||
if v < 1:
|
||||
raise ValueError(f"decoder_sparse_step must be at least 1, got {v}")
|
||||
return v
|
||||
|
||||
def save(self, config_path: str):
|
||||
config_dict = {k: v for k, v in asdict(self).items() if v is not None}
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
|
||||
@dataclass
|
||||
@ConfigFactory.register("embedding")
|
||||
class EncoderConfig(BaseModelConfig):
|
||||
"""Configuration for embedding encoder model.
|
||||
|
||||
Args:
|
||||
model_type (Optional[str]): Model type identifier for AutoModel dispatch. Defaults to None.
|
||||
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
|
||||
vocab_size (Optional[int]): Vocabulary size. Defaults to None.
|
||||
hidden_size (Optional[int]): Hidden dimension size. Defaults to None.
|
||||
num_hidden_layers (Optional[int]): Number of transformer layers. Defaults to None.
|
||||
rms_norm_eps (Optional[float]): Epsilon for RMSNorm. Defaults to None.
|
||||
intermediate_size (Optional[int]): Intermediate size in FFN. Defaults to None.
|
||||
max_position_embeddings (Optional[int]): Maximum sequence length the model was trained with. Defaults to None.
|
||||
rope_theta (Optional[float]): Base frequency for RoPE. Defaults to None.
|
||||
rope_scaling (Optional[dict]): RoPE scaling config, e.g. {"type": "linear", "factor": 4.0}. Defaults to None.
|
||||
attn_type (str): Attention type: 'gqa' or 'mla'. Defaults to "gqa".
|
||||
num_attention_heads (Optional[int]): Number of query attention heads. Defaults to None.
|
||||
num_key_value_heads (Optional[int]): Number of key/value heads for GQA. Defaults to None.
|
||||
use_qk_norm (Optional[bool]): Whether to apply RMSNorm to Q/K. Defaults to None.
|
||||
use_gated_attention (Optional[bool]): Whether to use gated attention. Defaults to None.
|
||||
ffn_type (str): FFN type: 'mlp' or 'moe'. Defaults to "mlp".
|
||||
pooling_type (Optional[str]): Pooling strategy for embedding, e.g. 'mean', 'cls'. Defaults to None.
|
||||
normalize_embeddings (Optional[bool]): Whether to L2-normalize output embeddings. Defaults to None.
|
||||
"""
|
||||
|
||||
vocab_size: Optional[int] = None
|
||||
hidden_size: Optional[int] = None
|
||||
num_hidden_layers: Optional[int] = None
|
||||
rms_norm_eps: Optional[float] = None
|
||||
intermediate_size: Optional[int] = None
|
||||
max_position_embeddings: Optional[int] = None
|
||||
rope_theta: Optional[float] = None
|
||||
rope_scaling: Optional[dict] = None
|
||||
attn_type: str = "gqa"
|
||||
num_attention_heads: Optional[int] = None
|
||||
num_key_value_heads: Optional[int] = None
|
||||
use_qk_norm: Optional[bool] = None
|
||||
use_gated_attention: Optional[bool] = None
|
||||
ffn_type: str = "mlp"
|
||||
pooling_type: Optional[str] = None
|
||||
normalize_embeddings: Optional[bool] = None
|
||||
|
||||
@field_validator("attn_type")
|
||||
def _validate_attn_type(cls, v: str) -> str:
|
||||
if v not in _ATTN_TYPES:
|
||||
raise ValueError(
|
||||
f"attn_type must be one of {sorted(_ATTN_TYPES)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator("ffn_type")
|
||||
def _validate_ffn_type(cls, v: str) -> str:
|
||||
if v not in _FFN_TYPES:
|
||||
raise ValueError(f"ffn_type must be one of {sorted(_FFN_TYPES)}, got {v!r}")
|
||||
return v
|
||||
|
||||
@@ -0,0 +1,152 @@
|
||||
"""Pipeline configuration for JSONL preprocessing.
|
||||
|
||||
Supports single-sequence (SFT/pretrain) and multi-output (DPO/GRPO)
|
||||
modes, both driven declaratively through ``input.sections`` or
|
||||
``input.sources``.
|
||||
"""
|
||||
|
||||
from dataclasses import field
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from pydantic import field_validator
|
||||
from pydantic.dataclasses import dataclass
|
||||
|
||||
from astrai.config.base import BaseConfig
|
||||
|
||||
_PACKING_STRATEGIES = frozenset({"simple", "bfd", "bfd_split"})
|
||||
_TRUNCATION_MODES = frozenset({"keep_start", "keep_end"})
|
||||
_STORAGE_FORMATS = frozenset({"bin", "jsonl"})
|
||||
_POSITION_IDS_MODES = frozenset({"none", "doc_reset", "continuous"})
|
||||
|
||||
|
||||
@dataclass
|
||||
class InputConfig(BaseConfig):
|
||||
"""Declarative input mapping.
|
||||
|
||||
Single-output mode (backward-compatible)::
|
||||
|
||||
{"input": {"sections": [{"field": "messages", ...}]}}
|
||||
|
||||
Multi-output mode (DPO / GRPO)::
|
||||
|
||||
{"input": {"sources": {
|
||||
"chosen": {"sections": [{"field": "chosen", ...}]},
|
||||
"rejected": {"sections": [{"field": "rejected", ...}]},
|
||||
}}}
|
||||
|
||||
Args:
|
||||
sections (Optional[List[Dict]]): Section list for single-output mode. Defaults to None.
|
||||
sources (Optional[Dict[str, Dict]]): Source map for multi-output mode, DPO/GRPO. Defaults to None.
|
||||
"""
|
||||
|
||||
sections: Optional[List[Dict]] = None
|
||||
sources: Optional[Dict[str, Dict]] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProcessingConfig(BaseConfig):
|
||||
"""Processing configuration for tokenization and packing.
|
||||
|
||||
Args:
|
||||
max_seq_len (int): Maximum sequence length. Defaults to 2048.
|
||||
min_chars (int): Minimum number of characters to keep. Defaults to 50.
|
||||
max_chars (int): Maximum number of characters to keep. Defaults to 2_000_000.
|
||||
max_items (Optional[int]): Maximum number of items to process, None=unlimited. Defaults to None.
|
||||
batch_size (int): Number of records tokenized together. Defaults to 256.
|
||||
packing_strategy (str): How to pack sequences: 'simple', 'bfd', or 'bfd_split'. Defaults to "simple".
|
||||
max_packed_len (int): Maximum length of a packed bin. Defaults to 8192.
|
||||
truncation_mode (str): How to truncate over-length sequences: 'keep_start' or 'keep_end'. Defaults to "keep_start".
|
||||
"""
|
||||
|
||||
max_seq_len: int = 2048
|
||||
min_chars: int = 50
|
||||
max_chars: int = 2_000_000
|
||||
max_items: Optional[int] = None
|
||||
batch_size: int = 256
|
||||
packing_strategy: str = "simple"
|
||||
max_packed_len: int = 8192
|
||||
truncation_mode: str = "keep_start"
|
||||
|
||||
@field_validator("packing_strategy")
|
||||
def _validate_packing_strategy(cls, v: str) -> str:
|
||||
if v not in _PACKING_STRATEGIES:
|
||||
raise ValueError(
|
||||
f"packing_strategy must be one of {sorted(_PACKING_STRATEGIES)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator("truncation_mode")
|
||||
def _validate_truncation_mode(cls, v: str) -> str:
|
||||
if v not in _TRUNCATION_MODES:
|
||||
raise ValueError(
|
||||
f"truncation_mode must be one of {sorted(_TRUNCATION_MODES)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator("max_seq_len", "batch_size", "max_packed_len")
|
||||
def _validate_positive_int(cls, v: int) -> int:
|
||||
if v <= 0:
|
||||
raise ValueError(f"must be positive, got {v}")
|
||||
return v
|
||||
|
||||
@field_validator("min_chars")
|
||||
def _validate_non_negative(cls, v: int) -> int:
|
||||
if v < 0:
|
||||
raise ValueError(f"min_chars must be non-negative, got {v}")
|
||||
return v
|
||||
|
||||
|
||||
@dataclass
|
||||
class OutputConfig(BaseConfig):
|
||||
"""Output configuration for storage.
|
||||
|
||||
Args:
|
||||
domain_key (Optional[str]): Domain key for the output store. Defaults to None.
|
||||
storage_format (str): Storage format: 'bin' or 'jsonl'. Defaults to "bin".
|
||||
max_tokens_per_shard (int): Maximum tokens per shard before splitting. Defaults to 100_000_000.
|
||||
dtype (Dict[str, str]): Per-key dtype overrides, e.g. {"input_ids": "int32"}. Defaults to {}.
|
||||
position_ids_mode (str): Position ids mode: 'none', 'doc_reset', or 'continuous'. Defaults to "doc_reset".
|
||||
"""
|
||||
|
||||
domain_key: Optional[str] = None
|
||||
storage_format: str = "bin"
|
||||
max_tokens_per_shard: int = 100_000_000
|
||||
dtype: Dict[str, str] = field(default_factory=dict)
|
||||
position_ids_mode: str = "doc_reset"
|
||||
|
||||
@field_validator("storage_format")
|
||||
def _validate_storage_format(cls, v: str) -> str:
|
||||
if v not in _STORAGE_FORMATS:
|
||||
raise ValueError(
|
||||
f"storage_format must be one of {sorted(_STORAGE_FORMATS)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator("position_ids_mode")
|
||||
def _validate_position_ids_mode(cls, v: str) -> str:
|
||||
if v not in _POSITION_IDS_MODES:
|
||||
raise ValueError(
|
||||
f"position_ids_mode must be one of {sorted(_POSITION_IDS_MODES)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
|
||||
@dataclass
|
||||
class PipelineConfig(BaseConfig):
|
||||
"""Top-level preprocessing pipeline config.
|
||||
|
||||
Args:
|
||||
version (int): Config schema version. Defaults to 1.
|
||||
input (InputConfig): Input mapping config.
|
||||
mask (Dict[str, str]): Per-field mask labels, e.g. {"system": "mask", "assistant": "train"}. Defaults to {}.
|
||||
mask_default (str): Default mask label for unlisted fields. Defaults to "mask".
|
||||
preprocessing (ProcessingConfig): Processing config.
|
||||
output (OutputConfig): Output config.
|
||||
"""
|
||||
|
||||
version: int = 1
|
||||
input: InputConfig = field(default_factory=InputConfig)
|
||||
mask: Dict[str, str] = field(default_factory=dict)
|
||||
mask_default: str = "mask"
|
||||
preprocessing: ProcessingConfig = field(default_factory=ProcessingConfig)
|
||||
output: OutputConfig = field(default_factory=OutputConfig)
|
||||
+204
-82
@@ -1,98 +1,220 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, Optional
|
||||
from dataclasses import field
|
||||
from typing import Any, Callable, Dict, List, Optional
|
||||
|
||||
import torch.nn as nn
|
||||
from pydantic import ConfigDict, field_validator, model_validator
|
||||
from pydantic.dataclasses import dataclass
|
||||
from torch.optim import Optimizer
|
||||
from torch.optim.lr_scheduler import LRScheduler
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from astrai.config.base import BaseConfig
|
||||
from astrai.model.components.lora import LoRAConfig
|
||||
|
||||
@dataclass
|
||||
class TrainConfig:
|
||||
# basic setting
|
||||
model: nn.Module = field(default=None, metadata={"help": "Model for training."})
|
||||
strategy: str = field(default=None, metadata={"help": "Training strategy."})
|
||||
dataset: Dataset = field(default=None, metadata={"help": "Dataset for training."})
|
||||
optimizer_fn: Callable[[nn.Module], Optimizer] = field(
|
||||
default=None, metadata={"help": "Optimizer factory for training."}
|
||||
)
|
||||
scheduler_fn: Callable[[Optimizer], LRScheduler] = field(
|
||||
default=None, metadata={"help": "Scheduler factory for training."}
|
||||
)
|
||||
n_epoch: int = field(default=1, metadata={"help": "Number of epochs for training."})
|
||||
batch_size: int = field(default=4, metadata={"help": "Batch size for training."})
|
||||
accumulation_steps: int = field(
|
||||
default=1, metadata={"help": "Number of iterations between steps."}
|
||||
)
|
||||
max_grad_norm: float = field(
|
||||
default=1.0, metadata={"help": "Maximum gradient norm."}
|
||||
)
|
||||
_TRAIN_TYPES = frozenset({"seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"})
|
||||
_PARALLEL_MODES = frozenset({"none", "ddp", "fsdp"})
|
||||
_BACKENDS = frozenset({"nccl", "gloo"})
|
||||
_START_METHODS = frozenset({"spawn", "fork", "forkserver"})
|
||||
_COMPILE_MODES = frozenset({"default", "reduce-overhead", "max-autotune"})
|
||||
|
||||
# checkpoint setting
|
||||
start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."})
|
||||
start_batch: int = field(
|
||||
default=0, metadata={"help": "Start batch iteration for training."}
|
||||
)
|
||||
ckpt_dir: str = field(
|
||||
default="./checkpoint", metadata={"help": "Checkpoint directory."}
|
||||
)
|
||||
ckpt_interval: int = field(
|
||||
default=5000, metadata={"help": "Number of iterations between checkpoints."}
|
||||
)
|
||||
|
||||
# dataloader setting
|
||||
random_seed: int = field(default=3407, metadata={"help": "Random seed."})
|
||||
num_workers: int = field(
|
||||
default=0, metadata={"help": "Number of workers for dataloader."}
|
||||
)
|
||||
prefetch_factor: Optional[int] = field(
|
||||
default=None, metadata={"help": "Prefetch factor for dataloader."}
|
||||
)
|
||||
pin_memory: bool = field(
|
||||
default=False, metadata={"help": "Pin memory for dataloader."}
|
||||
)
|
||||
@dataclass(config=ConfigDict(arbitrary_types_allowed=True))
|
||||
class TrainConfig(BaseConfig):
|
||||
"""Training configuration.
|
||||
|
||||
# distributed training
|
||||
nprocs: int = field(
|
||||
default=1, metadata={"help": "Number of processes for distributed training."}
|
||||
)
|
||||
backend: str = field(
|
||||
default="nccl", metadata={"help": "Distributed training backend."}
|
||||
)
|
||||
master_addr: str = field(
|
||||
default="localhost",
|
||||
metadata={"help": "Master address for distributed training."},
|
||||
)
|
||||
master_port: str = field(
|
||||
default="29500", metadata={"help": "Master port for distributed training."}
|
||||
)
|
||||
parallel_wrapper: Optional[Callable] = field(
|
||||
default=None, metadata={"help": "Parallel function for training."}
|
||||
)
|
||||
state_dict_fn: Optional[Callable] = field(
|
||||
default=None, metadata={"help": "Parallel function for state dict saving."}
|
||||
)
|
||||
Combines hyperparameters with runtime objects (model_fn, dataset, etc.).
|
||||
Only JSON-serializable fields are written to checkpoint meta via to_dict().
|
||||
|
||||
# others
|
||||
device_type: str = field(
|
||||
default="cuda", metadata={"help": "Device type for distributed training."}
|
||||
Args:
|
||||
model_fn (Callable[[], nn.Module]): Model factory for training.
|
||||
strategy (str): Training strategy (seq, sft, dpo, grpo, online_*).
|
||||
dataset (Dataset): Dataset for training.
|
||||
optimizer_fn (Callable[[nn.Module], Optimizer]): Optimizer factory for training.
|
||||
optimizer_name (Optional[str]): Serializable built-in optimizer identifier. Defaults to None.
|
||||
optimizer_hyperparameters (Dict[str, Any]): Serializable optimizer settings. Defaults to {}.
|
||||
scheduler_fn (Callable[[Optimizer], LRScheduler]): Scheduler factory for training.
|
||||
n_epoch (int): Number of epochs for training. Defaults to 1.
|
||||
batch_per_device (int): Batch size per device. Defaults to 4.
|
||||
grad_accum_steps (int): Number of iterations between optimizer steps. Defaults to 1.
|
||||
max_grad_norm (Optional[float]): Maximum gradient norm. None disables clipping. Defaults to 1.0.
|
||||
gradient_checkpointing_modules (List[type]): Module types to enable activation checkpointing for. Defaults to [].
|
||||
compile_mode (Optional[str]): torch.compile mode: 'default', 'reduce-overhead', 'max-autotune', or None. Defaults to None.
|
||||
start_epoch (int): Start epoch for training. Defaults to 0.
|
||||
start_samples (int): Start samples count (per rank). Superseded by checkpoint consumed_samples. Defaults to 0.
|
||||
ckpt_dir (str): Checkpoint directory. Defaults to "./checkpoint".
|
||||
ckpt_interval (int): Number of optimizer steps between checkpoints. Defaults to 5000.
|
||||
lora (Optional[LoRAConfig]): LoRA config. None means full fine-tuning. Defaults to None.
|
||||
metrics (List[str]): Metrics to record during training. Defaults to ["loss", "lr", "grad_norm"].
|
||||
random_seed (int): Random seed. Defaults to 3407.
|
||||
num_workers (int): Number of workers for dataloader. Defaults to 0.
|
||||
prefetch_factor (Optional[int]): Prefetch factor for dataloader. Defaults to None.
|
||||
persistent_workers (bool): Keep DataLoader workers alive between epochs. Defaults to False.
|
||||
pin_memory (bool): Pin memory for dataloader. Defaults to False.
|
||||
collate_fn (Optional[Callable[[List[Any]], Any]]): Collate function for dataloader (e.g. dpo_collate_fn). Defaults to None.
|
||||
nprocs (int): Number of processes for distributed training. Defaults to 1.
|
||||
backend (str): Distributed training backend. Defaults to "nccl".
|
||||
master_addr (str): Master address for distributed training. Defaults to "localhost".
|
||||
master_port (str): Master port for distributed training. Defaults to "29500".
|
||||
parallel_mode (str): Parallel strategy: none, ddp, fsdp. Defaults to "none".
|
||||
start_method (str): Multiprocessing start method: spawn/fork/forkserver. Defaults to "spawn".
|
||||
device_type (str): Device type for distributed training. Defaults to "cuda".
|
||||
val_dataset (Optional[Dataset]): Dataset for validation. Defaults to None.
|
||||
val_split (Optional[float]): Ratio to split from training dataset for validation, e.g. 0.05. Defaults to None.
|
||||
val_step (int): Number of optimizer steps between validation runs. Defaults to 1000.
|
||||
neftune_alpha (float): NEFTune noise alpha, 0=disabled, typical: 5.0. Defaults to 0.0.
|
||||
moe_aux_loss_coef (float): Weight applied to the MoE load-balancing loss. Defaults to 0.01.
|
||||
rollout_interval (int): Number of optimizer steps between online rollouts. Defaults to 512.
|
||||
rollout_temperature (float): Sampling temperature for online rollout. Defaults to 0.7.
|
||||
rollout_top_k (int): Top-k filtering for online rollout, 0=disable. Defaults to 0.
|
||||
rollout_top_p (float): Top-p (nucleus) filtering for online rollout. Defaults to 0.9.
|
||||
rollout_max_tokens (int): Maximum generated tokens per response in rollout. Defaults to 1024.
|
||||
reward_model_fn (Optional[Callable]): Factory for reward model, required for online RL strategies. Defaults to None.
|
||||
executor_kwargs (Dict[str, Any]): Extra kwargs passed to ExecutorFactory.create(). Defaults to {}.
|
||||
extra_kwargs (Dict[str, Any]): Other arguments. Defaults to {}.
|
||||
"""
|
||||
|
||||
model_fn: Callable[[], nn.Module]
|
||||
strategy: str
|
||||
dataset: Dataset
|
||||
optimizer_fn: Callable[[nn.Module], Optimizer]
|
||||
scheduler_fn: Callable[[Optimizer], LRScheduler]
|
||||
optimizer_name: Optional[str] = None
|
||||
optimizer_hyperparameters: Dict[str, Any] = field(default_factory=dict)
|
||||
n_epoch: int = 1
|
||||
batch_per_device: int = 4
|
||||
grad_accum_steps: int = 1
|
||||
max_grad_norm: Optional[float] = 1.0
|
||||
gradient_checkpointing_modules: List[type] = field(default_factory=list)
|
||||
compile_mode: Optional[str] = None
|
||||
|
||||
start_epoch: int = 0
|
||||
start_samples: int = 0
|
||||
ckpt_dir: str = "./checkpoint"
|
||||
ckpt_interval: int = 5000
|
||||
|
||||
lora: Optional[LoRAConfig] = None
|
||||
|
||||
metrics: List[str] = field(default_factory=lambda: ["loss", "lr", "grad_norm"])
|
||||
|
||||
random_seed: int = 3407
|
||||
num_workers: int = 0
|
||||
prefetch_factor: Optional[int] = None
|
||||
persistent_workers: bool = False
|
||||
pin_memory: bool = False
|
||||
collate_fn: Optional[Callable[[List[Any]], Any]] = None
|
||||
|
||||
nprocs: int = 1
|
||||
backend: str = "nccl"
|
||||
master_addr: str = "localhost"
|
||||
master_port: str = "29500"
|
||||
parallel_mode: str = "none"
|
||||
start_method: str = "spawn"
|
||||
|
||||
device_type: str = "cuda"
|
||||
val_dataset: Optional[Dataset] = None
|
||||
val_split: Optional[float] = None
|
||||
val_step: int = 1000
|
||||
neftune_alpha: float = 0.0
|
||||
moe_aux_loss_coef: float = 0.01
|
||||
|
||||
rollout_interval: int = 512
|
||||
rollout_temperature: float = 0.7
|
||||
rollout_top_k: int = 0
|
||||
rollout_top_p: float = 0.9
|
||||
rollout_max_tokens: int = 1024
|
||||
reward_model_fn: Optional[Callable] = None
|
||||
|
||||
executor_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||
extra_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
@field_validator("strategy")
|
||||
def _validate_strategy(cls, v: str) -> str:
|
||||
if v not in _TRAIN_TYPES:
|
||||
raise ValueError(
|
||||
f"strategy must be one of {sorted(_TRAIN_TYPES)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator("parallel_mode")
|
||||
def _validate_parallel_mode(cls, v: str) -> str:
|
||||
if v not in _PARALLEL_MODES:
|
||||
raise ValueError(
|
||||
f"parallel_mode must be one of {sorted(_PARALLEL_MODES)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator("backend")
|
||||
def _validate_backend(cls, v: str) -> str:
|
||||
if v not in _BACKENDS:
|
||||
raise ValueError(f"backend must be one of {sorted(_BACKENDS)}, got {v!r}")
|
||||
return v
|
||||
|
||||
@field_validator("start_method")
|
||||
def _validate_start_method(cls, v: str) -> str:
|
||||
if v not in _START_METHODS:
|
||||
raise ValueError(
|
||||
f"start_method must be one of {sorted(_START_METHODS)}, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator("compile_mode")
|
||||
def _validate_compile_mode(cls, v: Optional[str]) -> Optional[str]:
|
||||
if v is not None and v not in _COMPILE_MODES:
|
||||
raise ValueError(
|
||||
f"compile_mode must be one of {sorted(_COMPILE_MODES)} or None, got {v!r}"
|
||||
)
|
||||
return v
|
||||
|
||||
@field_validator(
|
||||
"n_epoch",
|
||||
"batch_per_device",
|
||||
"grad_accum_steps",
|
||||
"ckpt_interval",
|
||||
"val_step",
|
||||
"rollout_interval",
|
||||
"rollout_max_tokens",
|
||||
)
|
||||
extra_kwargs: dict = field(
|
||||
default_factory=dict, metadata={"help": "Other arguments."}
|
||||
def _validate_positive_int(cls, v: int) -> int:
|
||||
if v <= 0:
|
||||
raise ValueError(f"must be positive, got {v}")
|
||||
return v
|
||||
|
||||
@field_validator("rollout_temperature")
|
||||
def _validate_positive_float(cls, v: float) -> float:
|
||||
if v <= 0:
|
||||
raise ValueError(f"must be positive, got {v}")
|
||||
return v
|
||||
|
||||
@field_validator("rollout_top_p")
|
||||
def _validate_top_p(cls, v: float) -> float:
|
||||
if not 0 < v <= 1:
|
||||
raise ValueError(f"rollout_top_p must be in (0, 1], got {v}")
|
||||
return v
|
||||
|
||||
@field_validator(
|
||||
"rollout_top_k", "num_workers", "neftune_alpha", "moe_aux_loss_coef"
|
||||
)
|
||||
def _validate_non_negative(cls, v):
|
||||
if v < 0:
|
||||
raise ValueError(f"must be non-negative, got {v}")
|
||||
return v
|
||||
|
||||
def __post_init__(self):
|
||||
self.validate()
|
||||
@field_validator("max_grad_norm")
|
||||
def _validate_max_grad_norm(cls, v: Optional[float]) -> Optional[float]:
|
||||
if v is not None and v <= 0:
|
||||
raise ValueError(f"max_grad_norm must be positive or None, got {v}")
|
||||
return v
|
||||
|
||||
def validate(self):
|
||||
required_fields = [
|
||||
"model",
|
||||
"strategy",
|
||||
"dataset",
|
||||
"optimizer_fn",
|
||||
"scheduler_fn",
|
||||
]
|
||||
@field_validator("val_split")
|
||||
def _validate_val_split(cls, v: Optional[float]) -> Optional[float]:
|
||||
if v is not None and not 0 < v < 1:
|
||||
raise ValueError(f"val_split must be in (0, 1) or None, got {v}")
|
||||
return v
|
||||
|
||||
for field_name in required_fields:
|
||||
if getattr(self, field_name) is None:
|
||||
raise ValueError(f"{field_name} is required.")
|
||||
@model_validator(mode="after")
|
||||
def _validate_online_strategy(self) -> "TrainConfig":
|
||||
if self.strategy.startswith("online_") and self.reward_model_fn is None:
|
||||
raise ValueError(
|
||||
f"reward_model_fn is required for online RL strategy {self.strategy!r}"
|
||||
)
|
||||
return self
|
||||
|
||||
+28
-10
@@ -1,19 +1,37 @@
|
||||
from astrai.dataset.dataset import (
|
||||
BaseDataset,
|
||||
BaseSegmentFetcher,
|
||||
DatasetFactory,
|
||||
MultiSegmentFetcher,
|
||||
dpo_collate_fn,
|
||||
grpo_collate_fn,
|
||||
)
|
||||
from astrai.dataset.sampler import RDSampler
|
||||
from astrai.dataset.storage import (
|
||||
JsonlStore,
|
||||
MmapStore,
|
||||
Recordable,
|
||||
Store,
|
||||
StoreFactory,
|
||||
Streamable,
|
||||
detect_format,
|
||||
)
|
||||
from astrai.serialization import (
|
||||
load_bin,
|
||||
save_bin,
|
||||
)
|
||||
from astrai.dataset.sampler import ResumableDistributedSampler
|
||||
|
||||
__all__ = [
|
||||
# Base classes
|
||||
"BaseDataset",
|
||||
# Factory
|
||||
"DatasetFactory",
|
||||
# Fetchers
|
||||
"BaseSegmentFetcher",
|
||||
"MultiSegmentFetcher",
|
||||
# Sampler
|
||||
"ResumableDistributedSampler",
|
||||
"dpo_collate_fn",
|
||||
"grpo_collate_fn",
|
||||
"Store",
|
||||
"Streamable",
|
||||
"Recordable",
|
||||
"StoreFactory",
|
||||
"MmapStore",
|
||||
"JsonlStore",
|
||||
"detect_format",
|
||||
"save_bin",
|
||||
"load_bin",
|
||||
"RDSampler",
|
||||
]
|
||||
|
||||
+440
-240
@@ -1,338 +1,538 @@
|
||||
"""Dataset implementations with factory pattern for training."""
|
||||
"""Dataset implementations for training.
|
||||
|
||||
Composition over inheritance — every dataset is a thin wrapper that
|
||||
binds a :class:`Store` to a particular train-type's key mapping. All
|
||||
sample-id → token/record indexing lives on the Store; datasets never
|
||||
know about window/stride math or segment layouts.
|
||||
|
||||
Class hierarchy:
|
||||
|
||||
BaseDataset (ABC) — holds a Store, exposes __len__/keys,
|
||||
overrides __getitem__
|
||||
├── SEQDataset — next-token prediction (stream)
|
||||
├── SFTDataset — loss-mask + position_ids (stream)
|
||||
├── DPODataset — chosen/rejected pairs (record)
|
||||
└── GRPODataset — prompt + response group (record)
|
||||
|
||||
``DatasetFactory.load(train_type, load_path, window_size, stride, …)``
|
||||
builds the Store (auto-detecting format) before constructing the
|
||||
matching dataset. Passing ``store=`` skips Store construction.
|
||||
|
||||
When a record dataset (DPO) reads from raw JSONL, a *processor*
|
||||
function (pure ``record -> Dict[str, Tensor]``) is forwarded to
|
||||
:class:`JsonlStore` so tokenisation happens on the fly.
|
||||
"""
|
||||
|
||||
import bisect
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Dict, List, Optional, Union
|
||||
from functools import partial
|
||||
from pathlib import Path
|
||||
from typing import Callable, Dict, List, Optional
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
from astrai.dataset.storage import (
|
||||
Store,
|
||||
StoreFactory,
|
||||
detect_format,
|
||||
)
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.serialization import load_h5
|
||||
from astrai.preprocessing.transform import TokenizeTransform
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
_DEFAULT_MESSAGES_CONFIG = {
|
||||
"version": 1,
|
||||
"input": {"sections": [{"field": "messages", "action": "$role", "template": True}]},
|
||||
"mask": {"system": "mask", "user": "mask", "assistant": "train"},
|
||||
"mask_default": "mask",
|
||||
"output": {"position_ids_mode": "doc_reset"},
|
||||
}
|
||||
|
||||
|
||||
class BaseSegmentFetcher:
|
||||
"""Fetches data segments across multiple tensor segments.
|
||||
def _build_jsonl_transform(
|
||||
path: str, tokenizer_path: Optional[str] = None
|
||||
) -> Optional["TokenizeTransform"]:
|
||||
"""Auto-build a TokenizeTransform for JSONL eager loading.
|
||||
|
||||
Maintains cumulative lengths for efficient range queries across
|
||||
multiple discontinuous segments.
|
||||
Reads ``dataset_config.json`` from the data dir if present, or
|
||||
falls back to the built-in chatml SFT config when *tokenizer_path*
|
||||
is provided.
|
||||
"""
|
||||
|
||||
def __init__(self, segments: List[Tensor]):
|
||||
self.segments = segments
|
||||
self.cum_lengths = []
|
||||
|
||||
total = 0
|
||||
for seg in segments:
|
||||
total += torch.numel(seg)
|
||||
self.cum_lengths.append(total)
|
||||
|
||||
self.total_length = total
|
||||
|
||||
def __len__(self) -> int:
|
||||
return self.total_length
|
||||
|
||||
def fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
||||
"""Fetch data in the range [begin_idx, end_idx).
|
||||
|
||||
Args:
|
||||
begin_idx: Starting index (inclusive)
|
||||
end_idx: Ending index (exclusive)
|
||||
|
||||
Returns:
|
||||
Concatenated tensor of data in the specified range
|
||||
"""
|
||||
if not (
|
||||
0 <= begin_idx < self.total_length and 0 <= end_idx <= self.total_length
|
||||
):
|
||||
raise ValueError("begin_idx or end_idx out of bounds")
|
||||
if begin_idx >= end_idx:
|
||||
return torch.tensor([], dtype=torch.long)
|
||||
|
||||
# Find segment boundaries for the range
|
||||
seg_start_idx = bisect.bisect_right(self.cum_lengths, begin_idx)
|
||||
seg_end_idx = bisect.bisect_left(self.cum_lengths, end_idx)
|
||||
|
||||
result_segments = []
|
||||
|
||||
for i in range(seg_start_idx, seg_end_idx + 1):
|
||||
prev_cum = self.cum_lengths[i - 1] if i > 0 else 0
|
||||
start = max(begin_idx - prev_cum, 0)
|
||||
end = min(end_idx - prev_cum, len(self.segments[i]))
|
||||
data = self.segments[i][start:end]
|
||||
result_segments.append(data)
|
||||
|
||||
return torch.cat(result_segments, dim=0)
|
||||
root = Path(path)
|
||||
config_path = root / "dataset_config.json" if root.is_dir() else None
|
||||
if config_path is not None and config_path.exists():
|
||||
return TokenizeTransform.from_config_file(str(config_path))
|
||||
if tokenizer_path:
|
||||
config = PipelineConfig.from_dict(_DEFAULT_MESSAGES_CONFIG)
|
||||
return TokenizeTransform(config, tokenizer_path)
|
||||
return None
|
||||
|
||||
|
||||
class MultiSegmentFetcher:
|
||||
"""Manages multiple segment fetchers for different data keys.
|
||||
def dpo_tokenize(
|
||||
record: dict,
|
||||
tokenizer,
|
||||
max_len: int = 2048,
|
||||
) -> Optional[dict]:
|
||||
"""Tokenize one DPO record into chosen/rejected + masks.
|
||||
|
||||
Each key corresponds to a different type of data (e.g., "sequence", "mask").
|
||||
Applies the tokenizer's chat template so token sequences match the
|
||||
SFT checkpoint's format. Prompt is rendered with
|
||||
``add_generation_prompt=True``; chosen/rejected are appended as a
|
||||
single assistant turn.
|
||||
|
||||
Accepts:
|
||||
|
||||
- Flat: ``{"prompt": str, "chosen": str, "rejected": str}``
|
||||
- Conv: ``{"prompt": [{role, content}, ...], "chosen": [...], ...}``
|
||||
- Legacy: ``{"input": str, "chosen": str, "rejected": str}``
|
||||
|
||||
No packing, no ``position_ids`` — DPO sequences are independent.
|
||||
"""
|
||||
prompt = record.get("prompt") or record.get("input")
|
||||
chosen = record.get("chosen")
|
||||
rejected = record.get("rejected")
|
||||
if prompt is None or chosen is None or rejected is None:
|
||||
return None
|
||||
|
||||
def __init__(self, multi_segments: Dict):
|
||||
self.multi_keys = list(multi_segments.keys())
|
||||
self.multi_fetchers = {
|
||||
key: BaseSegmentFetcher(segments)
|
||||
for key, segments in multi_segments.items()
|
||||
}
|
||||
prompt_messages = _to_messages(prompt)
|
||||
chosen_text = _extract_text(chosen)
|
||||
rejected_text = _extract_text(rejected)
|
||||
if chosen_text is None or rejected_text is None:
|
||||
return None
|
||||
chosen_messages = prompt_messages + [{"role": "assistant", "content": chosen_text}]
|
||||
rejected_messages = prompt_messages + [
|
||||
{"role": "assistant", "content": rejected_text}
|
||||
]
|
||||
|
||||
def __len__(self) -> int:
|
||||
"""Returns the minimum length across all fetchers."""
|
||||
len_list = [len(seg) for seg in self.multi_fetchers.values()]
|
||||
return min(len_list)
|
||||
prompt_ids = tokenizer.apply_chat_template(
|
||||
prompt_messages, tokenize=True, add_generation_prompt=True
|
||||
)
|
||||
ch_ids = tokenizer.apply_chat_template(
|
||||
chosen_messages, tokenize=True, add_generation_prompt=False
|
||||
)
|
||||
re_ids = tokenizer.apply_chat_template(
|
||||
rejected_messages, tokenize=True, add_generation_prompt=False
|
||||
)
|
||||
|
||||
def key_fetch(
|
||||
self, begin_idx: int, end_idx: int, keys: Union[str, List[str]]
|
||||
) -> Dict:
|
||||
"""Fetch data for specific keys.
|
||||
full_ch = ch_ids[:max_len]
|
||||
full_re = re_ids[:max_len]
|
||||
|
||||
Args:
|
||||
begin_idx: Starting index
|
||||
end_idx: Ending index
|
||||
keys: Single key or list of keys to fetch
|
||||
prompt_len = min(len(prompt_ids), max_len)
|
||||
ch_mask = [0] * prompt_len + [1] * max(0, len(full_ch) - prompt_len)
|
||||
ch_mask = ch_mask[:max_len]
|
||||
re_mask = [0] * prompt_len + [1] * max(0, len(full_re) - prompt_len)
|
||||
re_mask = re_mask[:max_len]
|
||||
|
||||
Returns:
|
||||
Dictionary of tensors if multiple keys, single tensor if one key
|
||||
"""
|
||||
fetch_dict = {}
|
||||
keys = [keys] if isinstance(keys, str) else keys
|
||||
return {
|
||||
"chosen": full_ch,
|
||||
"rejected": full_re,
|
||||
"chosen_mask": ch_mask,
|
||||
"rejected_mask": re_mask,
|
||||
}
|
||||
|
||||
for key in keys:
|
||||
fetcher = self.multi_fetchers[key]
|
||||
fetch_tensor = fetcher.fetch_data(begin_idx, end_idx)
|
||||
fetch_dict[key] = fetch_tensor
|
||||
|
||||
return fetch_dict if len(keys) > 1 else fetch_dict[keys[0]]
|
||||
def _to_messages(value) -> list:
|
||||
"""Accept str or conversation list; return message list."""
|
||||
if isinstance(value, str):
|
||||
return [{"role": "user", "content": value}]
|
||||
if isinstance(value, list):
|
||||
return value
|
||||
return [{"role": "user", "content": str(value)}]
|
||||
|
||||
def fetch_data(self, begin_idx: int, end_idx: int) -> Dict:
|
||||
"""Fetch all keys."""
|
||||
return self.key_fetch(begin_idx, end_idx, self.multi_keys)
|
||||
|
||||
def _extract_text(value) -> Optional[str]:
|
||||
"""Accept str or conversation list; return plain text."""
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
if isinstance(value, list):
|
||||
return "".join(m.get("content", "") for m in value if isinstance(m, dict))
|
||||
return None
|
||||
|
||||
|
||||
def dpo_processor(
|
||||
record: dict,
|
||||
tokenizer,
|
||||
max_len: int = 2048,
|
||||
) -> Dict[str, Tensor]:
|
||||
"""DPO processor: wraps :func:`dpo_tokenize` and returns tensors."""
|
||||
result = dpo_tokenize(record, tokenizer, max_len=max_len)
|
||||
if result is None:
|
||||
raise ValueError(f"Malformed DPO record: {list(record.keys())}")
|
||||
return {
|
||||
"chosen": torch.tensor(result["chosen"], dtype=torch.int32),
|
||||
"rejected": torch.tensor(result["rejected"], dtype=torch.int32),
|
||||
"chosen_mask": torch.tensor(result["chosen_mask"], dtype=torch.bool),
|
||||
"rejected_mask": torch.tensor(result["rejected_mask"], dtype=torch.bool),
|
||||
}
|
||||
|
||||
|
||||
def dpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
||||
"""Collate variable-length DPO samples into padded 2-D tensors.
|
||||
|
||||
Input: list of dicts, each with:
|
||||
- chosen: [C_i]
|
||||
- rejected: [R_i]
|
||||
- chosen_mask: [C_i]
|
||||
- rejected_mask: [R_i]
|
||||
|
||||
Output (padded to the max length across chosen/rejected within the batch):
|
||||
- chosen: [B, S_max]
|
||||
- rejected: [B, S_max]
|
||||
- chosen_mask: [B, S_max]
|
||||
- rejected_mask: [B, S_max]
|
||||
"""
|
||||
B = len(batch)
|
||||
S_max = max(b["chosen"].size(0) for b in batch)
|
||||
S_max = max(S_max, max(b["rejected"].size(0) for b in batch))
|
||||
|
||||
chosen = torch.zeros(B, S_max, dtype=torch.long)
|
||||
rejected = torch.zeros(B, S_max, dtype=torch.long)
|
||||
chosen_mask = torch.zeros(B, S_max, dtype=torch.bool)
|
||||
rejected_mask = torch.zeros(B, S_max, dtype=torch.bool)
|
||||
|
||||
for i, b in enumerate(batch):
|
||||
c_len = b["chosen"].size(0)
|
||||
r_len = b["rejected"].size(0)
|
||||
chosen[i, :c_len] = b["chosen"]
|
||||
rejected[i, :r_len] = b["rejected"]
|
||||
chosen_mask[i, :c_len] = b["chosen_mask"]
|
||||
rejected_mask[i, :r_len] = b["rejected_mask"]
|
||||
|
||||
return {
|
||||
"chosen": chosen,
|
||||
"rejected": rejected,
|
||||
"chosen_mask": chosen_mask,
|
||||
"rejected_mask": rejected_mask,
|
||||
}
|
||||
|
||||
|
||||
def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
||||
"""Collate variable-length GRPO samples into padded 3-D tensors.
|
||||
|
||||
Input: list of dicts, each with:
|
||||
- prompts: [P_i]
|
||||
- responses: list of G tensors, each [R_ij]
|
||||
- masks: list of G tensors, each [R_ij]
|
||||
- rewards: [G]
|
||||
|
||||
Output:
|
||||
- prompts: [B, P_max], left-padded
|
||||
- prompt_mask: [B, P_max]
|
||||
- responses: [B, G, R_max]
|
||||
- masks: [B, G, R_max]
|
||||
- rewards: [B, G]
|
||||
"""
|
||||
B = len(batch)
|
||||
G = len(batch[0]["responses"])
|
||||
P_max = max(b["prompts"].size(0) for b in batch)
|
||||
R_max = max(r.size(0) for b in batch for r in b["responses"])
|
||||
|
||||
prompts = torch.zeros(B, P_max, dtype=torch.long)
|
||||
prompt_mask = torch.zeros(B, P_max, dtype=torch.bool)
|
||||
responses = torch.zeros(B, G, R_max, dtype=torch.long)
|
||||
masks = torch.zeros(B, G, R_max, dtype=torch.bool)
|
||||
rewards = torch.zeros(B, G, dtype=torch.float32)
|
||||
|
||||
for i, b in enumerate(batch):
|
||||
p_len = b["prompts"].size(0)
|
||||
prompts[i, -p_len:] = b["prompts"]
|
||||
prompt_mask[i, -p_len:] = True
|
||||
rewards[i, : b["rewards"].size(0)] = b["rewards"]
|
||||
for g in range(min(G, len(b["responses"]))):
|
||||
r_len = b["responses"][g].size(0)
|
||||
responses[i, g, :r_len] = b["responses"][g]
|
||||
if g < len(b["masks"]):
|
||||
masks[i, g, :r_len] = b["masks"][g]
|
||||
|
||||
return {
|
||||
"prompts": prompts,
|
||||
"prompt_mask": prompt_mask,
|
||||
"responses": responses,
|
||||
"masks": masks,
|
||||
"rewards": rewards,
|
||||
}
|
||||
|
||||
|
||||
def validate_keys(store: Store, required: List[str]) -> None:
|
||||
"""Raise ``KeyError`` if *store* is missing any *required* key."""
|
||||
if not required:
|
||||
return
|
||||
actual = set(store.keys)
|
||||
missing = [k for k in required if k not in actual]
|
||||
if missing:
|
||||
raise KeyError(
|
||||
f"Store at {getattr(store, '_load_path', '?')} is missing required "
|
||||
f"keys {missing}; available keys are {sorted(actual)}."
|
||||
)
|
||||
|
||||
|
||||
class BaseDataset(Dataset, ABC):
|
||||
"""Abstract base class for all dataset types.
|
||||
"""Abstract base class for dataset types.
|
||||
|
||||
Implements common functionality for window-based data fetching.
|
||||
Holds a :class:`Store`. All sample-id indexing is delegated to the
|
||||
store — this class exposes ``__len__`` as ``len(store)`` and the
|
||||
``keys`` property as ``store.keys``. Subclasses implement
|
||||
``__getitem__`` with the train-type-specific key mapping and any
|
||||
training-only index arithmetic (e.g. the next-token ``+1`` shift).
|
||||
"""
|
||||
|
||||
def __init__(self, window_size: int, stride: int):
|
||||
required_keys: List[str] = []
|
||||
|
||||
def __init__(self, store: Store):
|
||||
super().__init__()
|
||||
self.segments = {}
|
||||
self.window_size = window_size
|
||||
self.stride = stride
|
||||
self.total_samples = None
|
||||
self.fetcher: Optional[MultiSegmentFetcher] = None
|
||||
self.store: Store = store
|
||||
validate_keys(store, self.required_keys)
|
||||
|
||||
def load(self, load_path: str):
|
||||
"""Load dataset from HDF5 file.
|
||||
def __len__(self) -> int:
|
||||
return len(self.store)
|
||||
|
||||
Args:
|
||||
load_path: Path to the HDF5 data file
|
||||
"""
|
||||
self.segments = load_h5(load_path)
|
||||
self.fetcher = MultiSegmentFetcher(self.segments)
|
||||
self.total_samples = len(self.fetcher)
|
||||
@property
|
||||
def keys(self) -> List[str]:
|
||||
return self.store.keys
|
||||
|
||||
def get_index(self, index: int) -> tuple:
|
||||
"""Calculate begin and end indices for a sample.
|
||||
|
||||
Args:
|
||||
index: Sample index
|
||||
|
||||
Returns:
|
||||
Tuple of (begin_idx, end_idx)
|
||||
"""
|
||||
assert self.total_samples > self.window_size
|
||||
|
||||
begin_idx = min(index * self.stride, self.total_samples - 1 - self.window_size)
|
||||
end_idx = min(begin_idx + self.window_size, self.total_samples - 1)
|
||||
|
||||
return begin_idx, end_idx
|
||||
@property
|
||||
def token_count(self) -> int:
|
||||
return self.store.token_count
|
||||
|
||||
@abstractmethod
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
"""Get a single sample by index.
|
||||
|
||||
Must be implemented by subclasses.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def __len__(self) -> int:
|
||||
assert self.total_samples is not None
|
||||
if self.total_samples <= self.window_size:
|
||||
return 0
|
||||
return (self.total_samples - 1 - self.window_size) // self.stride + 1
|
||||
|
||||
|
||||
class DatasetFactory(BaseFactory["BaseDataset"]):
|
||||
"""Factory class for creating dataset instances.
|
||||
"""Factory for creating dataset instances by train-type.
|
||||
|
||||
Supports decorator-based registration for extensible dataset types.
|
||||
All default dataset types (seq, sft, dpo, grpo) are registered automatically
|
||||
when their classes are defined with the decorator.
|
||||
|
||||
Example usage:
|
||||
@DatasetFactory.register("custom")
|
||||
class CustomDataset(BaseDataset):
|
||||
...
|
||||
|
||||
dataset = DatasetFactory.create("custom", window_size, stride)
|
||||
Use :meth:`DatasetFactory.register("custom")` to register new
|
||||
dataset classes; they must inherit from :class:`BaseDataset`.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def _validate_component(cls, dataset_cls: type) -> None:
|
||||
"""Validate that the dataset class inherits from BaseDataset."""
|
||||
if not issubclass(dataset_cls, BaseDataset):
|
||||
raise TypeError(f"{dataset_cls.__name__} must inherit from BaseDataset")
|
||||
|
||||
@classmethod
|
||||
def create(cls, train_type: str, window_size: int, stride: int) -> "BaseDataset":
|
||||
"""Create a dataset instance.
|
||||
|
||||
Args:
|
||||
train_type: Type of training ("seq", "sft", "dpo", "grpo")
|
||||
window_size: Window size for data sampling
|
||||
stride: Stride between consecutive samples
|
||||
|
||||
Returns:
|
||||
Dataset instance
|
||||
"""
|
||||
return super().create(train_type, window_size, stride)
|
||||
|
||||
@classmethod
|
||||
def load(
|
||||
cls,
|
||||
train_type: str,
|
||||
load_path: str,
|
||||
window_size: int,
|
||||
load_path: Optional[str] = None,
|
||||
window_size: int = 0,
|
||||
stride: Optional[int] = None,
|
||||
storage_type: Optional[str] = None,
|
||||
tokenizer_path: Optional[str] = None,
|
||||
max_len: int = 2048,
|
||||
store: Optional[Store] = None,
|
||||
**kwargs,
|
||||
) -> "BaseDataset":
|
||||
"""Create and load a dataset in one step.
|
||||
|
||||
Two entry points:
|
||||
|
||||
- **store given**: bind it directly — the caller fully controls
|
||||
Store construction and processor setup. *load_path*,
|
||||
*storage_type*, *tokenizer_path*, *window_size*, *stride* are
|
||||
ignored.
|
||||
- **store is None**: build a Store from *load_path*, auto-detecting
|
||||
format and constructing a processor when *tokenizer_path* is
|
||||
given for a record dataset on JSONL.
|
||||
|
||||
Args:
|
||||
train_type: Type of training dataset
|
||||
load_path: Path to the data file
|
||||
window_size: Window size for data sampling
|
||||
stride: Stride between consecutive samples (default: same as window_size)
|
||||
train_type: Registered dataset name ("seq", "sft", "dpo",
|
||||
"grpo", …).
|
||||
load_path: Path to the data file or directory (ignored if
|
||||
*store* is given).
|
||||
window_size: Stream window length — only meaningful for
|
||||
stream datasets (SEQ/SFT). Record datasets ignore it.
|
||||
stride: Stride between consecutive stream samples
|
||||
(default: same as *window_size*).
|
||||
storage_type: Storage backend ("bin", "jsonl") or
|
||||
None for auto-detection.
|
||||
tokenizer_path: Path to tokenizer for lazy JSONL
|
||||
tokenisation (record datasets only).
|
||||
max_len: Max sequence length forwarded to processors.
|
||||
store: Pre-built, already-loaded Store instance.
|
||||
**kwargs: Extra arguments forwarded to ``store.load()``.
|
||||
|
||||
Returns:
|
||||
Loaded dataset instance
|
||||
Loaded dataset instance.
|
||||
"""
|
||||
if store is not None:
|
||||
return cls.create(train_type, store=store)
|
||||
|
||||
if load_path is None:
|
||||
raise ValueError("Either load_path or store must be provided")
|
||||
|
||||
if storage_type is None:
|
||||
storage_type = detect_format(load_path)
|
||||
|
||||
if stride is None:
|
||||
stride = window_size
|
||||
|
||||
dataset = cls.create(train_type, window_size, stride)
|
||||
dataset.load(load_path)
|
||||
processor = cls._maybe_build_processor(
|
||||
train_type, storage_type, tokenizer_path, max_len
|
||||
)
|
||||
|
||||
return dataset
|
||||
store_window = cls._store_window_for(train_type, window_size)
|
||||
store = StoreFactory.create(
|
||||
storage_type,
|
||||
window_size=store_window,
|
||||
stride=stride if stride else store_window,
|
||||
)
|
||||
if processor is not None:
|
||||
store.load(load_path, processor=processor, **kwargs)
|
||||
elif storage_type == "jsonl":
|
||||
transform = _build_jsonl_transform(load_path, tokenizer_path)
|
||||
if transform is None:
|
||||
raise FileNotFoundError(
|
||||
f"JSONL dataset config not found. Expected "
|
||||
f"dataset_config.json alongside *.jsonl files, pass "
|
||||
f"tokenizer_path= for the built-in messages config, or "
|
||||
f"use processor= for lazy on-the-fly tokenisation."
|
||||
)
|
||||
store.load(load_path, transform=transform, **kwargs)
|
||||
else:
|
||||
store.load(load_path, **kwargs)
|
||||
|
||||
@classmethod
|
||||
def available_types(cls) -> list:
|
||||
"""Return list of registered dataset type names."""
|
||||
return cls.list_registered()
|
||||
return cls.create(train_type, store=store)
|
||||
|
||||
@staticmethod
|
||||
def _store_window_for(train_type: str, window_size: int) -> int:
|
||||
"""Stream datasets consume ``window_size``; record datasets ignore it.
|
||||
|
||||
# ============== Dataset Classes ==============
|
||||
# All dataset classes are registered at class definition time using the decorator
|
||||
Record datasets (dpo/grpo) treat each record as an independent
|
||||
training unit and never window, so the store is built with
|
||||
``window_size=0`` and ``len(store)`` returns the record count.
|
||||
"""
|
||||
if train_type in ("seq", "sft"):
|
||||
return window_size
|
||||
return 0
|
||||
|
||||
@staticmethod
|
||||
def _maybe_build_processor(
|
||||
train_type: str,
|
||||
storage_type: str,
|
||||
tokenizer_path: Optional[str],
|
||||
max_len: int,
|
||||
) -> Optional[Callable[[dict], Dict[str, Tensor]]]:
|
||||
"""Build an on-the-fly tokenisation processor if applicable.
|
||||
|
||||
Only raw JSONL + record datasets (DPO/GRPO) need a processor;
|
||||
pre-tokenised backends (bin) and stream datasets (SEQ/SFT)
|
||||
return ``None`` so no tokenizer is loaded.
|
||||
"""
|
||||
if tokenizer_path is None or storage_type != "jsonl":
|
||||
return None
|
||||
if train_type == "dpo":
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
||||
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
|
||||
return None
|
||||
|
||||
|
||||
@DatasetFactory.register("seq")
|
||||
class SEQDataset(BaseDataset):
|
||||
"""Dataset for sequential next-token prediction training."""
|
||||
"""Dataset for sequential next-token prediction training.
|
||||
|
||||
def __init__(self, window_size: int, stride: int):
|
||||
super().__init__(window_size, stride)
|
||||
Stream mode: ``store.fetch(begin, end, "sequence")`` returns the
|
||||
input window; the +1 shifted call returns the next-token target.
|
||||
"""
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
||||
return self.fetcher.key_fetch(begin_idx, end_idx, "sequence")
|
||||
required_keys = ["sequence"]
|
||||
|
||||
def __getitem__(self, index):
|
||||
begin_idx, end_idx = self.get_index(index)
|
||||
|
||||
x = self._fetch_data(begin_idx, end_idx).to(dtype=torch.long)
|
||||
y = self._fetch_data(begin_idx + 1, end_idx + 1).to(dtype=torch.long)
|
||||
|
||||
return {"input_ids": x, "target_ids": y}
|
||||
def __getitem__(self, index: int):
|
||||
begin, end = self.store.sample_window(index)
|
||||
x = self.store.fetch(begin, end, "sequence")
|
||||
y = self.store.fetch(begin + 1, end + 1, "sequence")
|
||||
return {
|
||||
"input_ids": x.to(dtype=torch.long),
|
||||
"target_ids": y.to(dtype=torch.long),
|
||||
}
|
||||
|
||||
|
||||
@DatasetFactory.register("sft")
|
||||
class SFTDataset(BaseDataset):
|
||||
"""Dataset for supervised fine-tuning with loss masking."""
|
||||
"""Dataset for supervised fine-tuning with loss masking.
|
||||
|
||||
def __init__(self, window_size: int, stride: int):
|
||||
super().__init__(window_size, stride)
|
||||
Stream mode: ``sequence``/``loss_mask``/``position_ids`` are sliced
|
||||
to the window. ``loss_mask`` and ``target_ids`` use the +1 shifted
|
||||
slice so they align with the predicted positions.
|
||||
"""
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
||||
required_keys = ["sequence", "loss_mask", "position_ids"]
|
||||
|
||||
def __getitem__(self, index):
|
||||
begin_idx, end_idx = self.get_index(index)
|
||||
|
||||
x = self._fetch_data(begin_idx, end_idx, "sequence").to(dtype=torch.long)
|
||||
y = self._fetch_data(begin_idx + 1, end_idx + 1, "sequence").to(
|
||||
dtype=torch.long
|
||||
)
|
||||
loss_mask = self._fetch_data(begin_idx + 1, end_idx + 1, "loss_mask").to(
|
||||
dtype=torch.bool
|
||||
)
|
||||
|
||||
return {"input_ids": x, "target_ids": y, "loss_mask": loss_mask}
|
||||
def __getitem__(self, index: int):
|
||||
begin, end = self.store.sample_window(index)
|
||||
x = self.store.fetch(begin, end, "sequence")
|
||||
y = self.store.fetch(begin + 1, end + 1, "sequence")
|
||||
position_ids = self.store.fetch(begin, end, "position_ids")
|
||||
loss_mask = self.store.fetch(begin + 1, end + 1, "loss_mask")
|
||||
return {
|
||||
"input_ids": x.to(dtype=torch.long),
|
||||
"target_ids": y.to(dtype=torch.long),
|
||||
"position_ids": position_ids.to(dtype=torch.long),
|
||||
"loss_mask": loss_mask.to(dtype=torch.bool),
|
||||
}
|
||||
|
||||
|
||||
@DatasetFactory.register("dpo")
|
||||
class DPODataset(BaseDataset):
|
||||
"""Dataset for Direct Preference Optimization training."""
|
||||
"""Record-structured dataset for Direct Preference Optimization.
|
||||
|
||||
def __init__(self, window_size: int, stride: int):
|
||||
super().__init__(window_size, stride)
|
||||
Each sample is one preference pair (chosen + rejected) and is an
|
||||
independent training unit — no windowing, stride, or cross-record
|
||||
concatenation. This keeps each sequence self-contained so attention
|
||||
never leaks across preference pairs.
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
||||
Two loading paths (handled by :class:`DatasetFactory`):
|
||||
|
||||
def __getitem__(self, index: int):
|
||||
begin_idx, end_idx = self.get_index(index)
|
||||
- **Pre-tokenized** (bin): ``store.load(path)`` reads per-record
|
||||
tensors; ``__getitem__`` returns them directly.
|
||||
- **Raw JSONL** (``tokenizer_path=...``): builds a lazy processor
|
||||
via :func:`dpo_processor` that tokenises on the fly — no packing,
|
||||
no ``position_ids``.
|
||||
"""
|
||||
|
||||
chosen = self._fetch_data(begin_idx, end_idx, "chosen").to(dtype=torch.long)
|
||||
rejected = self._fetch_data(begin_idx, end_idx, "rejected").to(dtype=torch.long)
|
||||
chosen_mask = self._fetch_data(begin_idx, end_idx, "chosen_mask").to(
|
||||
dtype=torch.bool
|
||||
)
|
||||
rejected_mask = self._fetch_data(begin_idx, end_idx, "rejected_mask").to(
|
||||
dtype=torch.bool
|
||||
)
|
||||
required_keys = ["chosen", "rejected", "chosen_mask", "rejected_mask"]
|
||||
|
||||
def make_processor(self, tokenizer, max_len: int):
|
||||
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
|
||||
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
return {
|
||||
"chosen": chosen,
|
||||
"rejected": rejected,
|
||||
"chosen_mask": chosen_mask,
|
||||
"rejected_mask": rejected_mask,
|
||||
"chosen": self.store.fetch_record(index, "chosen").to(dtype=torch.long),
|
||||
"rejected": self.store.fetch_record(index, "rejected").to(dtype=torch.long),
|
||||
"chosen_mask": self.store.fetch_record(index, "chosen_mask").to(
|
||||
dtype=torch.bool
|
||||
),
|
||||
"rejected_mask": self.store.fetch_record(index, "rejected_mask").to(
|
||||
dtype=torch.bool
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
@DatasetFactory.register("grpo")
|
||||
class GRPODataset(BaseDataset):
|
||||
"""Dataset for Group Relative Policy Optimization training."""
|
||||
"""Dataset for offline Group Relative Policy Optimization.
|
||||
|
||||
def __init__(self, window_size: int, stride: int):
|
||||
super().__init__(window_size, stride)
|
||||
Each sample is one prompt with its group of responses and scalar
|
||||
rewards — an independent training unit with no windowing or stride.
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
||||
Expected storage layout (produced by JsonlStore or pre-tokenized):
|
||||
|
||||
- ``prompts``: List[Tensor] — one 1-D token tensor per record
|
||||
- ``responses``: List[List[Tensor]] — G response tensors per record
|
||||
- ``masks``: List[List[Tensor]] — G mask tensors per record
|
||||
- ``rewards``: List[Tensor] — one 1-D float tensor (len G) per record
|
||||
"""
|
||||
|
||||
required_keys = ["prompts", "responses", "masks", "rewards"]
|
||||
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
begin_idx, end_idx = self.get_index(index)
|
||||
|
||||
prompts = self._fetch_data(begin_idx, end_idx, "prompts")
|
||||
responses = self._fetch_data(begin_idx, end_idx, "responses")
|
||||
masks = self._fetch_data(begin_idx, end_idx, "masks")
|
||||
rewards = self._fetch_data(begin_idx, end_idx, "rewards")
|
||||
|
||||
prompts = self.store.fetch_record(index, "prompts")
|
||||
responses = self.store.fetch_record(index, "responses")
|
||||
masks = self.store.fetch_record(index, "masks")
|
||||
rewards = self.store.fetch_record(index, "rewards")
|
||||
return {
|
||||
"prompts": prompts,
|
||||
"responses": responses,
|
||||
"masks": masks,
|
||||
"rewards": rewards,
|
||||
"prompts": prompts.to(dtype=torch.long),
|
||||
"responses": [r.to(dtype=torch.long) for r in responses],
|
||||
"masks": [m.to(dtype=torch.bool) for m in masks],
|
||||
"rewards": rewards.to(dtype=torch.float32),
|
||||
}
|
||||
|
||||
@@ -5,7 +5,15 @@ import torch.distributed as dist
|
||||
from torch.utils.data import Dataset, Sampler
|
||||
|
||||
|
||||
class ResumableDistributedSampler(Sampler[int]):
|
||||
class RDSampler(Sampler[int]):
|
||||
"""Resumable Distributed Sampler.
|
||||
|
||||
A distributed sampler that supports checkpoint-based resume: iteration
|
||||
state (epoch, position) is tracked so training can continue from the
|
||||
exact sample after a restart. Shards the dataset across
|
||||
``dist.world_size`` replicas with optional shuffling.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
data_source: Dataset,
|
||||
@@ -43,6 +51,7 @@ class ResumableDistributedSampler(Sampler[int]):
|
||||
offset = 0 if drop_last else self.num_replicas - 1
|
||||
self.num_samples_per_replica = (self.num_samples + offset) // self.num_replicas
|
||||
self.total_size = self.num_samples_per_replica * self.num_replicas
|
||||
self.iter = self.iter % self.num_samples_per_replica
|
||||
|
||||
self._indices = None
|
||||
|
||||
@@ -73,6 +82,12 @@ class ResumableDistributedSampler(Sampler[int]):
|
||||
|
||||
self.epoch += 1
|
||||
self._indices = None
|
||||
self.iter = self.iter % self.num_samples_per_replica
|
||||
|
||||
@property
|
||||
def _remaining(self):
|
||||
remaining = self.num_samples_per_replica - self.iter
|
||||
return max(remaining, 0)
|
||||
|
||||
def __len__(self):
|
||||
return self.num_samples_per_replica
|
||||
return self._remaining
|
||||
|
||||
@@ -0,0 +1,601 @@
|
||||
"""Storage backends for different data formats.
|
||||
|
||||
Architecture (composition over inheritance):
|
||||
|
||||
Store (ABC) — owns _data/_cum/_offsets bookkeeping
|
||||
+ window_size/stride for sample-id
|
||||
indexing. __getitem__/__len__ produce
|
||||
the smallest iterable unit so Dataset
|
||||
classes are pure delegators.
|
||||
Streamable (mixin) — raw token slice fetch(begin, end, keys)
|
||||
Recordable (mixin) — raw record slice fetch_record(idx, keys)
|
||||
|
||||
MmapStore(Store, Streamable, Recordable)
|
||||
JsonlStore(Store, Streamable, Recordable)
|
||||
|
||||
Each mixin is a stateless trait that relies on ``self._data`` etc.
|
||||
provided by :class:`Store`. Concrete stores mix in whichever access
|
||||
primitives they support — ``Store`` is the sole base class, so there is
|
||||
no diamond inheritance or MRO ambiguity.
|
||||
|
||||
Sample-id indexing lives on :class:`Store`, not on the dataset:
|
||||
|
||||
- **Stream mode** (``window_size > 0``): ``len(store)`` returns the number
|
||||
of ``(window_size, stride)`` windows that fit in the token river;
|
||||
``store[i]`` returns the *i*-th window as a dict of per-key tensors;
|
||||
``store.sample_window(i)`` exposes the underlying ``(begin, end)``
|
||||
token slice for callers (e.g. next-token trainers) that need a +1
|
||||
shifted companion window.
|
||||
- **Record mode** (``num_records > 0``): ``len(store)`` returns the
|
||||
record count; ``store[i]`` returns the *i*-th record dict.
|
||||
|
||||
Raw token/record access via :meth:`fetch` / :meth:`fetch_record`
|
||||
remains available for low-level callers that want explicit index
|
||||
control. ``store.token_count`` is the total stream token count (what
|
||||
``len(store)`` used to mean in the legacy stream-only API).
|
||||
|
||||
``segments_are_records`` (class attribute on each Store subclass)
|
||||
tells ``_normalize`` whether segments are inherently per-record (JSONL)
|
||||
or opaque shards (bin). Record access for bin relies on ``_offsets``
|
||||
instead.
|
||||
|
||||
:class:`JsonlStore` supports a lazy mode (``processor=fn``) that keeps
|
||||
raw records and defers tokenisation to ``fetch_record`` — used by DPO
|
||||
to train directly from a ``.jsonl`` file without a pre-tokenised copy.
|
||||
"""
|
||||
|
||||
import bisect
|
||||
import glob
|
||||
import json
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from pathlib import Path
|
||||
from typing import Callable, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.serialization import (
|
||||
load_bin,
|
||||
load_bin_offsets,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def detect_format(load_path: str) -> str:
|
||||
"""Auto-detect storage format from files in the directory.
|
||||
|
||||
Args:
|
||||
load_path: Directory or file path
|
||||
|
||||
Returns:
|
||||
Format string ("h5", "bin", "jsonl", or "processed")
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If no supported data files are found
|
||||
"""
|
||||
root = Path(load_path)
|
||||
if root.is_file():
|
||||
suffix = root.suffix.lower()
|
||||
if suffix == ".jsonl":
|
||||
return "jsonl"
|
||||
raise ValueError(f"Unsupported file format: {suffix}")
|
||||
|
||||
bin_files = [Path(p) for p in glob.glob(str(root / "**" / "*.bin"), recursive=True)]
|
||||
if bin_files:
|
||||
has_meta = (root / "meta.json").exists() or len(
|
||||
[Path(p) for p in glob.glob(str(root / "**" / "meta.json"), recursive=True)]
|
||||
) > 0
|
||||
if has_meta:
|
||||
return "bin"
|
||||
jsonl_files = [
|
||||
Path(p) for p in glob.glob(str(root / "**" / "*.jsonl"), recursive=True)
|
||||
]
|
||||
if jsonl_files:
|
||||
return "jsonl"
|
||||
raise FileNotFoundError(f"No supported data files found at {load_path}")
|
||||
|
||||
|
||||
class Store(ABC):
|
||||
"""Common base for all storage backends.
|
||||
|
||||
A Store owns both its data layout AND its sample-id → token/record
|
||||
index translation. Datasets are thin wrappers that bind a Store
|
||||
to a particular train-type's key mapping; they never know about
|
||||
window/stride math.
|
||||
|
||||
Two iteration modes:
|
||||
|
||||
- **Stream** (``window_size > 0``): data is treated as one long
|
||||
token river. ``len(store)`` returns the number of windows;
|
||||
``store[i]`` slices every stream-compatible key to window ``i``;
|
||||
``store.sample_window(i)`` returns the ``(begin, end)`` token
|
||||
slice for callers needing a +1 shifted companion window.
|
||||
- **Record** (``num_records > 0``): data is per-record.
|
||||
``len(store)`` returns ``num_records``; ``store[i]`` returns
|
||||
the *i*-th record as a dict.
|
||||
|
||||
Raw token slicing is still available via :meth:`fetch` (mixed in
|
||||
by :class:`Streamable`) when a store has stream support configured.
|
||||
Raw record slicing via :meth:`fetch_record` (mixed in by
|
||||
:class:`Recordable`) when a store has record support.
|
||||
|
||||
``token_count`` exposes the raw total stream length — this is what
|
||||
``len(store)`` returned in the legacy stream-only API and what
|
||||
stream-bound ``fetch`` uses for its bounds check.
|
||||
"""
|
||||
|
||||
segments_are_records: bool = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window_size: int = 0,
|
||||
stride: Optional[int] = None,
|
||||
):
|
||||
self._data: Dict[str, List[Tensor]] = {}
|
||||
self._cum: Dict[str, List[int]] = {}
|
||||
self._offsets: Dict[str, List[int]] = {}
|
||||
self._length: int = 0
|
||||
self._num_records: int = 0
|
||||
self._window_size: int = int(window_size)
|
||||
self._stride: int = int(stride) if stride is not None else int(window_size)
|
||||
|
||||
@abstractmethod
|
||||
def load(self, path: str, **kwargs) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
@property
|
||||
def keys(self) -> List[str]:
|
||||
return list(self._data.keys())
|
||||
|
||||
@property
|
||||
def window_size(self) -> int:
|
||||
return self._window_size
|
||||
|
||||
@property
|
||||
def stride(self) -> int:
|
||||
return self._stride
|
||||
|
||||
@property
|
||||
def token_count(self) -> int:
|
||||
"""Total tokens across all stream segments.
|
||||
|
||||
Useful for the bounds-checked raw :meth:`fetch` and as the
|
||||
legacy ``len(store)`` value.
|
||||
"""
|
||||
return self._length
|
||||
|
||||
@property
|
||||
def num_records(self) -> int:
|
||||
"""Number of records available via :meth:`fetch_record`.
|
||||
|
||||
Non-zero only when the backing layout provides per-record
|
||||
indexing (JSONL segments or bin ``_offsets``).
|
||||
"""
|
||||
return self._num_records
|
||||
|
||||
@property
|
||||
def num_samples(self) -> int:
|
||||
"""Number of items produced by ``__getitem__``.
|
||||
|
||||
Stream-mode wins when ``window_size > 0`` and there are tokens
|
||||
to slice; otherwise falls back to ``num_records``.
|
||||
"""
|
||||
if self._window_size > 0 and self._length > 0:
|
||||
total = self._length
|
||||
w = self._window_size
|
||||
if total <= w:
|
||||
return 0
|
||||
return (total - 1 - w) // self._stride + 1
|
||||
return self._num_records
|
||||
|
||||
def __len__(self) -> int:
|
||||
return self.num_samples
|
||||
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
if index < 0:
|
||||
index += self.num_samples
|
||||
if not 0 <= index < self.num_samples:
|
||||
raise IndexError(
|
||||
f"Store index out of range: {index}, num_samples={self.num_samples}"
|
||||
)
|
||||
if self._window_size > 0 and self._length > 0:
|
||||
begin, end = self.sample_window(index)
|
||||
keys = self._stream_keys()
|
||||
return {k: self.fetch(begin, end, k) for k in keys}
|
||||
return self.fetch_record(index, self._record_keys())
|
||||
|
||||
def sample_window(self, index: int) -> Tuple[int, int]:
|
||||
"""Return ``(begin, end)`` token positions for stream sample *index*.
|
||||
|
||||
The clipped tail keeps the last reachable window inside the
|
||||
token river instead of overshooting. Caller is responsible
|
||||
for staying within :attr:`num_samples`: an out-of-range index
|
||||
raises ``IndexError``.
|
||||
"""
|
||||
if self._window_size <= 0:
|
||||
raise RuntimeError("sample_window() requires window_size > 0 (stream mode)")
|
||||
if self._window_size <= 0 or self._length <= self._window_size:
|
||||
raise IndexError(
|
||||
f"Data too short for window: token_count={self._length}, "
|
||||
f"window_size={self._window_size}"
|
||||
)
|
||||
if not 0 <= index < self.num_samples:
|
||||
raise IndexError(
|
||||
f"Sample index out of range: {index}, num_samples={self.num_samples}"
|
||||
)
|
||||
total = self._length
|
||||
begin = min(index * self._stride, total - 1 - self._window_size)
|
||||
end = min(begin + self._window_size, total - 1)
|
||||
return begin, end
|
||||
|
||||
def _stream_keys(self) -> List[str]:
|
||||
out: List[str] = []
|
||||
for k, tensors in self._data.items():
|
||||
if tensors and isinstance(tensors[0], list):
|
||||
continue
|
||||
out.append(k)
|
||||
return out
|
||||
|
||||
def _record_keys(self) -> List[str]:
|
||||
return list(self._data.keys())
|
||||
|
||||
def _normalize(
|
||||
self,
|
||||
raw: Dict[str, list],
|
||||
offsets: Optional[Dict[str, List[int]]] = None,
|
||||
):
|
||||
"""Register segments and pre-compute indices for both access modes.
|
||||
|
||||
Stream mode: ``_cum[key]`` accumulates per-segment lengths so
|
||||
``Streamable._fetch_stream_key`` can bisect across segments
|
||||
without concatenation.
|
||||
|
||||
Record mode: if *offsets* is provided (bin layout),
|
||||
``_offsets[key]`` stores cumulative per-record offsets into the
|
||||
single concatenated segment. Otherwise, when
|
||||
``segments_are_records`` is True (JSONL), ``_data[key]`` is
|
||||
a per-record list and ``fetch_record`` indexes it directly.
|
||||
|
||||
Nested keys (GRPO ``responses``/``masks`` as
|
||||
``List[List[Tensor]]``) are stored as-is and excluded from both
|
||||
cumulative bookkeepings — they are only accessed record-by-record.
|
||||
"""
|
||||
flat_lengths = []
|
||||
for key, tensors in raw.items():
|
||||
self._data[key] = tensors
|
||||
if not tensors:
|
||||
self._cum[key] = []
|
||||
flat_lengths.append(0)
|
||||
continue
|
||||
if isinstance(tensors[0], list):
|
||||
self._cum[key] = []
|
||||
continue
|
||||
cum = []
|
||||
total = 0
|
||||
for t in tensors:
|
||||
total += t.shape[0]
|
||||
cum.append(total)
|
||||
self._cum[key] = cum
|
||||
flat_lengths.append(cum[-1] if cum else 0)
|
||||
self._length = min(flat_lengths) if flat_lengths else 0
|
||||
|
||||
valid_offsets: Dict[str, List[int]] = {}
|
||||
if offsets:
|
||||
for key, off in offsets.items():
|
||||
segs = self._data.get(key, [])
|
||||
if len(segs) == 1 and len(off) > 1:
|
||||
valid_offsets[key] = off
|
||||
elif len(segs) > 1:
|
||||
logger.warning(
|
||||
"Key '%s' has %d segments with offsets — record mode "
|
||||
"disabled for this key (multi-shard bin+offsets not "
|
||||
"supported). Merge shards or use JSONL.",
|
||||
key,
|
||||
len(segs),
|
||||
)
|
||||
self._offsets = valid_offsets
|
||||
if valid_offsets:
|
||||
record_counts = [len(v) - 1 for v in valid_offsets.values()]
|
||||
self._num_records = min(record_counts) if record_counts else 0
|
||||
elif self.segments_are_records:
|
||||
per_record_counts = []
|
||||
for key, tensors in self._data.items():
|
||||
if tensors and isinstance(tensors[0], list):
|
||||
continue
|
||||
per_record_counts.append(len(tensors))
|
||||
self._num_records = min(per_record_counts) if per_record_counts else 0
|
||||
else:
|
||||
self._num_records = 0
|
||||
|
||||
|
||||
class Streamable:
|
||||
"""Mixin granting raw token-stream access via :meth:`fetch`.
|
||||
|
||||
Stateless trait relying on ``self._data``, ``self._cum``,
|
||||
``self._length`` maintained by :class:`Store`. Stream mode is
|
||||
active when the owning store has ``window_size > 0``; for stores
|
||||
that can also serve record access (JSONL/bin+offsets), the
|
||||
``fetch_record`` API from :class:`Recordable` is used instead.
|
||||
"""
|
||||
|
||||
def fetch(
|
||||
self,
|
||||
begin: int,
|
||||
end: int,
|
||||
keys: Union[str, List[str]],
|
||||
):
|
||||
return _stream_fetch(self, begin, end, keys)
|
||||
|
||||
|
||||
def _stream_fetch(self, begin: int, end: int, keys: Union[str, List[str]]):
|
||||
if not getattr(self, "_data", None):
|
||||
raise RuntimeError("Store not loaded")
|
||||
if not (0 <= begin < self._length and 0 <= end <= self._length):
|
||||
raise ValueError(
|
||||
f"Index out of bounds: begin={begin}, end={end}, length={self._length}"
|
||||
)
|
||||
if isinstance(keys, str):
|
||||
return _fetch_stream_key(self, keys, begin, end)
|
||||
return {k: _fetch_stream_key(self, k, begin, end) for k in keys}
|
||||
|
||||
|
||||
def _fetch_stream_key(self, key: str, begin: int, end: int) -> Tensor:
|
||||
segments = self._data[key]
|
||||
cum = self._cum[key]
|
||||
seg_start = bisect.bisect_right(cum, begin)
|
||||
seg_end = bisect.bisect_left(cum, end)
|
||||
|
||||
results = []
|
||||
for i in range(seg_start, seg_end + 1):
|
||||
prev = cum[i - 1] if i > 0 else 0
|
||||
s = max(begin - prev, 0)
|
||||
e = min(end - prev, segments[i].shape[0])
|
||||
results.append(segments[i][s:e])
|
||||
|
||||
return results[0] if len(results) == 1 else torch.cat(results, dim=0)
|
||||
|
||||
|
||||
class Recordable:
|
||||
"""Mixin granting raw record access via :meth:`fetch_record`.
|
||||
|
||||
Stateless trait relying on ``self._data``, ``self._offsets``,
|
||||
``self._num_records`` maintained by :class:`Store`.
|
||||
"""
|
||||
|
||||
def fetch_record(
|
||||
self,
|
||||
index: int,
|
||||
keys: Union[str, List[str]],
|
||||
):
|
||||
return _record_fetch(self, index, keys)
|
||||
|
||||
|
||||
def _record_fetch(self, index: int, keys: Union[str, List[str]]):
|
||||
if not getattr(self, "_data", None) and self._num_records == 0:
|
||||
raise RuntimeError("Store not loaded")
|
||||
if not 0 <= index < self._num_records:
|
||||
raise ValueError(
|
||||
f"Record index out of bounds: {index}, num_records={self._num_records}"
|
||||
)
|
||||
if isinstance(keys, str):
|
||||
return _fetch_record_key(self, keys, index)
|
||||
return {k: _fetch_record_key(self, k, index) for k in keys}
|
||||
|
||||
|
||||
def _fetch_record_key(self, key: str, index: int):
|
||||
offsets = self._offsets.get(key)
|
||||
if offsets:
|
||||
start = offsets[index]
|
||||
end = (
|
||||
offsets[index + 1]
|
||||
if index + 1 < len(offsets)
|
||||
else self._data[key][0].shape[0]
|
||||
)
|
||||
return self._data[key][0][start:end]
|
||||
return self._data[key][index]
|
||||
|
||||
|
||||
class StoreFactory(BaseFactory["Store"]):
|
||||
"""Factory for creating Store instances by type name."""
|
||||
|
||||
|
||||
@StoreFactory.register("bin")
|
||||
class MmapStore(Store, Streamable, Recordable):
|
||||
"""Memory-mapped binary storage backend.
|
||||
|
||||
Each key is a single .bin file backed by ``np.memmap(mode="r")``.
|
||||
No per-process memory duplication — all DataLoader workers share the
|
||||
same OS page-cache pages.
|
||||
|
||||
Supports both access modes:
|
||||
|
||||
- **Stream**: always available via :meth:`fetch`.
|
||||
- **Record** (``fetch_record(i, key)``): only when ``meta.json``
|
||||
contains per-record ``offsets`` (written via
|
||||
``save_bin(..., record_keys=...)``). Legacy bin files without
|
||||
offsets have ``num_records == 0`` and ``len(store)`` reflects the
|
||||
windowed sample count when ``window_size > 0``.
|
||||
|
||||
``segments_are_records`` is ``False`` here (bin segments are
|
||||
contiguous streams, not per-record) — record access is driven
|
||||
purely by ``_offsets``.
|
||||
"""
|
||||
|
||||
segments_are_records = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window_size: int = 0,
|
||||
stride: Optional[int] = None,
|
||||
):
|
||||
super().__init__(window_size=window_size, stride=stride)
|
||||
self._mmap_refs: List[Tensor] = []
|
||||
|
||||
def load(self, path: str, **kwargs):
|
||||
self._mmap_refs = []
|
||||
root = Path(path)
|
||||
all_raw: Dict[str, List[Tensor]] = {}
|
||||
all_offsets: Dict[str, List[int]] = {}
|
||||
meta_paths = [
|
||||
Path(p) for p in glob.glob(str(root / "**" / "meta.json"), recursive=True)
|
||||
]
|
||||
for meta_path in meta_paths:
|
||||
raw = load_bin(str(meta_path.parent))
|
||||
off = load_bin_offsets(str(meta_path.parent))
|
||||
for key, tensors in raw.items():
|
||||
if key not in all_raw:
|
||||
all_raw[key] = []
|
||||
all_raw[key].extend(tensors)
|
||||
for key, o in off.items():
|
||||
if key not in all_offsets:
|
||||
all_offsets[key] = []
|
||||
all_offsets[key].extend(o)
|
||||
if not meta_paths:
|
||||
raise FileNotFoundError(f"No meta.json found under {path}")
|
||||
self._normalize(all_raw, offsets=all_offsets or None)
|
||||
for tensors in self._data.values():
|
||||
self._mmap_refs.extend(tensors)
|
||||
|
||||
|
||||
class JsonlSource:
|
||||
"""Read raw JSON records from a ``.jsonl`` file or directory.
|
||||
|
||||
A thin reader used by :class:`JsonlStore` in processor mode — holds
|
||||
no tokenizer, performs no tokenisation, just yields dicts.
|
||||
"""
|
||||
|
||||
def __init__(self, path: str):
|
||||
self.path = Path(path)
|
||||
self._records: Optional[List[dict]] = None
|
||||
|
||||
def load(self) -> List[dict]:
|
||||
if self._records is None:
|
||||
self._records = self._read(self.path)
|
||||
return self._records
|
||||
|
||||
@staticmethod
|
||||
def _read(root: Path) -> List[dict]:
|
||||
if root.is_file():
|
||||
return JsonlSource._read_file(root)
|
||||
return JsonlSource._read_dir(root)
|
||||
|
||||
@staticmethod
|
||||
def _read_file(path: Path) -> List[dict]:
|
||||
records: List[dict] = []
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
records.append(json.loads(line))
|
||||
except json.JSONDecodeError:
|
||||
logger.warning("Failed to parse JSON line in %s, skipping", path)
|
||||
return records
|
||||
|
||||
@staticmethod
|
||||
def _read_dir(root: Path) -> List[dict]:
|
||||
records: List[dict] = []
|
||||
for jsonl_path in sorted(root.glob("*.jsonl")):
|
||||
records.extend(JsonlSource._read_file(jsonl_path))
|
||||
return records
|
||||
|
||||
|
||||
@StoreFactory.register("jsonl")
|
||||
class JsonlStore(Store, Streamable, Recordable):
|
||||
"""JSONL reader with eager/lazy tokenisation modes.
|
||||
|
||||
A JSONL dataset is a ``.jsonl`` file or a directory of ``*.jsonl``
|
||||
files plus (optionally) a ``dataset_config.json`` describing the
|
||||
tokenization pipeline.
|
||||
|
||||
Three ways to supply an eager transform (first match wins):
|
||||
|
||||
- **Explicit** (``transform=``): caller-built
|
||||
:class:`TokenizeTransform` applied eagerly.
|
||||
- **Config file**: ``dataset_config.json`` alongside the ``*.jsonl``
|
||||
files — loaded via :meth:`TokenizeTransform.from_config_file`.
|
||||
- **Default messages** (``tokenizer_path=`` given, no config file):
|
||||
a built-in chatml config that tokenises the ``messages`` field,
|
||||
masking every role except ``assistant`` (loss on assistant only).
|
||||
Lets SFT/SEQ train straight from a chat-style JSONL directory
|
||||
without a hand-written config.
|
||||
|
||||
Two tokenisation modes, selected at :meth:`load` time:
|
||||
|
||||
- **Eager** (default): applies the transform to every record at load
|
||||
time and registers per-key tensors via ``_normalize``. Both
|
||||
``fetch`` (stream) and ``fetch_record`` (record) work.
|
||||
- **Lazy** (``processor=fn`` passed): keeps raw records and defers
|
||||
tokenisation to ``fetch_record``. Only record access works —
|
||||
``len(store)`` returns ``num_records``; stream primitives raise.
|
||||
"""
|
||||
|
||||
segments_are_records = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window_size: int = 0,
|
||||
stride: Optional[int] = None,
|
||||
):
|
||||
super().__init__(window_size=window_size, stride=stride)
|
||||
self._source: Optional[JsonlSource] = None
|
||||
self._processor: Optional[Callable[[dict], Dict[str, Tensor]]] = None
|
||||
self._keys_cache: Optional[List[str]] = None
|
||||
|
||||
def load(self, path: str, transform=None, processor=None, **kwargs):
|
||||
self._source = JsonlSource(path)
|
||||
records = self._source.load()
|
||||
|
||||
if processor is not None:
|
||||
self._processor = processor
|
||||
self._num_records = len(records)
|
||||
return
|
||||
|
||||
if transform is None:
|
||||
raise ValueError(
|
||||
"JsonlStore eager mode requires transform=. "
|
||||
"Use DatasetFactory.load() which auto-constructs it."
|
||||
)
|
||||
|
||||
transformed = transform.apply(records)
|
||||
self._normalize(transformed)
|
||||
|
||||
@property
|
||||
def keys(self) -> List[str]:
|
||||
if self._processor is not None:
|
||||
if self._keys_cache is None and self._num_records > 0:
|
||||
sample = self._processor(self._source.load()[0])
|
||||
self._keys_cache = list(sample.keys())
|
||||
return self._keys_cache or []
|
||||
return list(self._data.keys())
|
||||
|
||||
def fetch_record(self, index: int, keys: Union[str, List[str]]):
|
||||
if self._processor is not None:
|
||||
if not 0 <= index < self._num_records:
|
||||
raise ValueError(
|
||||
f"Record index out of bounds: {index}, "
|
||||
f"num_records={self._num_records}"
|
||||
)
|
||||
record = self._source.load()[index]
|
||||
data = self._processor(record)
|
||||
if isinstance(keys, str):
|
||||
return data[keys]
|
||||
return {k: data[k] for k in keys}
|
||||
return _record_fetch(self, index, keys)
|
||||
|
||||
def fetch(self, begin: int, end: int, keys: Union[str, List[str]]):
|
||||
if self._processor is not None:
|
||||
raise RuntimeError(
|
||||
"JsonlStore in lazy (processor) mode does not support "
|
||||
"stream fetch(); use fetch_record() instead."
|
||||
)
|
||||
return _stream_fetch(self, begin, end, keys)
|
||||
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
if self._processor is not None:
|
||||
return self.fetch_record(index, self._record_keys())
|
||||
return super().__getitem__(index)
|
||||
@@ -0,0 +1,55 @@
|
||||
"""CUDA attention kernel wrappers with torch fallback.
|
||||
|
||||
Public API:
|
||||
- ``attn_decode`` — single-query decode attention
|
||||
- ``attn_prefill`` — multi-query prefill attention
|
||||
- ``attn_paged_decode`` — paged decode attention (direct page-table access)
|
||||
- ``AttentionBackend`` — ABC for attention computation strategies
|
||||
- ``TorchNativeBackend`` — default SDPA backend with KV cache I/O
|
||||
- ``CudaBackend`` — CUDA kernel backend with paged decode + prefill
|
||||
|
||||
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
||||
(blhd). Scale is always ``1/sqrt(head_dim)``.
|
||||
|
||||
Each wrapper calls its compiled CUDA kernel directly. Fallback to torch
|
||||
SDPA is handled by the attention backend, not the wrapper functions.
|
||||
"""
|
||||
|
||||
from astrai.extension.backend import (
|
||||
ATTN_BACKEND,
|
||||
AttentionBackend,
|
||||
AttentionBackendFactory,
|
||||
CudaBackend,
|
||||
FlashAttnBackend,
|
||||
TorchNativeBackend,
|
||||
apply_rotary_emb,
|
||||
attention,
|
||||
attn_backend,
|
||||
get_backend,
|
||||
)
|
||||
from astrai.extension.loader import KERNEL_NAMES, is_available
|
||||
from astrai.extension.ops import (
|
||||
TensorLayout,
|
||||
attn_decode,
|
||||
attn_paged_decode,
|
||||
attn_prefill,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ATTN_BACKEND",
|
||||
"AttentionBackend",
|
||||
"AttentionBackendFactory",
|
||||
"CudaBackend",
|
||||
"TorchNativeBackend",
|
||||
"FlashAttnBackend",
|
||||
"TensorLayout",
|
||||
"attention",
|
||||
"attn_backend",
|
||||
"get_backend",
|
||||
"attn_decode",
|
||||
"attn_paged_decode",
|
||||
"attn_prefill",
|
||||
"is_available",
|
||||
"KERNEL_NAMES",
|
||||
"apply_rotary_emb",
|
||||
]
|
||||
@@ -0,0 +1,27 @@
|
||||
"""Backend selection, fallbacks, and execution policies."""
|
||||
|
||||
from astrai.extension.backend.attention import (
|
||||
ATTN_BACKEND,
|
||||
AttentionBackend,
|
||||
AttentionBackendFactory,
|
||||
CudaBackend,
|
||||
FlashAttnBackend,
|
||||
TorchNativeBackend,
|
||||
attention,
|
||||
attn_backend,
|
||||
get_backend,
|
||||
)
|
||||
from astrai.extension.backend.rotary import apply_rotary_emb
|
||||
|
||||
__all__ = [
|
||||
"ATTN_BACKEND",
|
||||
"AttentionBackend",
|
||||
"AttentionBackendFactory",
|
||||
"CudaBackend",
|
||||
"FlashAttnBackend",
|
||||
"TorchNativeBackend",
|
||||
"apply_rotary_emb",
|
||||
"attention",
|
||||
"attn_backend",
|
||||
"get_backend",
|
||||
]
|
||||
@@ -0,0 +1,702 @@
|
||||
"""Attention backend abstraction with context-manager switching.
|
||||
|
||||
The backend encapsulates KV cache I/O and attention computation. The
|
||||
attention module (GQA/MLA) keeps projections, rotary, QK-norm, gating,
|
||||
and output projection; the backend handles everything from "write K/V
|
||||
to cache" through "SDPA output".
|
||||
|
||||
Usage — mirroring ``torch.nn.attention.sdpa_kernel``:
|
||||
|
||||
from astrai.extension import attn_backend, ATTN_BACKEND
|
||||
|
||||
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
|
||||
engine.generate("hello")
|
||||
|
||||
# or with an instance:
|
||||
with attn_backend(TorchNativeBackend()):
|
||||
...
|
||||
|
||||
# or the shorthand (instance is itself a context manager):
|
||||
with TorchNativeBackend():
|
||||
...
|
||||
|
||||
Thread-safe via ``contextvars`` — each scheduler thread gets its own
|
||||
active backend. ``get_backend()`` returns the active one, falling back
|
||||
to a process-wide default (cuda > flash > torch, overridable via
|
||||
``ASTR_BACKEND``).
|
||||
|
||||
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
||||
(blhd). The backend returns ``[batch, seq_len, n_heads * head_dim]``.
|
||||
"""
|
||||
|
||||
import contextvars
|
||||
import enum
|
||||
import functools
|
||||
import os
|
||||
import threading
|
||||
from abc import ABC, abstractmethod
|
||||
from contextlib import contextmanager
|
||||
from typing import TYPE_CHECKING, Optional, Union
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.extension.loader import is_available
|
||||
from astrai.extension.ops.attention import (
|
||||
attn_paged_decode,
|
||||
attn_paged_prefill,
|
||||
)
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
try:
|
||||
import flash_attn as _flash_attn
|
||||
except Exception:
|
||||
_flash_attn = None
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from astrai.inference.cache import KVCache
|
||||
|
||||
|
||||
_default_backend: Optional["AttentionBackend"] = None
|
||||
_default_backend_lock = threading.Lock()
|
||||
_env_backend_name: Optional[str] = None
|
||||
_env_backend: Optional["AttentionBackend"] = None
|
||||
_current_backend: contextvars.ContextVar[Optional["AttentionBackend"]] = (
|
||||
contextvars.ContextVar("attn_backend", default=None)
|
||||
)
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def flash_attn_available() -> bool:
|
||||
if not torch.cuda.is_available():
|
||||
return False
|
||||
fa = _flash_attn
|
||||
if fa is None:
|
||||
return False
|
||||
|
||||
try:
|
||||
major = int(fa.__version__.split(".")[0])
|
||||
cc = torch.cuda.get_device_capability()
|
||||
cc_num = cc[0] * 10 + cc[1]
|
||||
except Exception:
|
||||
major, cc_num = 0, 0
|
||||
if (major >= 3 and cc_num < 90) or (major < 3 and 0 < cc_num < 70):
|
||||
return False
|
||||
|
||||
try:
|
||||
if not hasattr(fa, "flash_attn_func"):
|
||||
return False
|
||||
x = torch.zeros(1, 1, 1, 64, device="cuda", dtype=torch.bfloat16)
|
||||
out = fa.flash_attn_func(x, x, x, causal=True)
|
||||
return bool(torch.isfinite(out).all().item())
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
class ATTN_BACKEND(enum.Enum):
|
||||
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
|
||||
|
||||
TORCH_NATIVE = "torch_native"
|
||||
CUDA = "cuda"
|
||||
FLASH = "flash"
|
||||
|
||||
|
||||
def _priority_backends() -> list["AttentionBackend"]:
|
||||
"""Available backends in priority order: cuda -> flash -> torch."""
|
||||
backends: list[AttentionBackend] = []
|
||||
if is_available("attn_paged_decode") and is_available("attn_paged_prefill"):
|
||||
backends.append(CudaBackend())
|
||||
if flash_attn_available():
|
||||
backends.append(FlashAttnBackend())
|
||||
backends.append(TorchNativeBackend())
|
||||
return backends
|
||||
|
||||
|
||||
def _backend_supports(
|
||||
backend: "AttentionBackend",
|
||||
q: Tensor,
|
||||
kv_cache: Optional["KVCache"],
|
||||
attn_mask: Optional[Tensor],
|
||||
is_causal: bool,
|
||||
fwd: Optional[str],
|
||||
) -> bool:
|
||||
"""Whether ``backend`` can run this attention call.
|
||||
|
||||
The CUDA kernels are bf16-only, support head_dim in 32/64/128/256, and
|
||||
need a KV cache (decode/prefill); everything else falls back to torch.
|
||||
"""
|
||||
if isinstance(backend, CudaBackend):
|
||||
return (
|
||||
fwd in ("prefill", "decode")
|
||||
and kv_cache is not None
|
||||
and q.ndim == 3
|
||||
and q.dtype == torch.bfloat16
|
||||
and q.size(-1) in (32, 64, 128, 256)
|
||||
and is_available(f"attn_paged_{fwd}")
|
||||
)
|
||||
if isinstance(backend, FlashAttnBackend):
|
||||
if not flash_attn_available():
|
||||
return False
|
||||
if q.dtype not in (torch.float16, torch.bfloat16):
|
||||
return False
|
||||
if fwd is not None:
|
||||
return q.ndim == 3 and hasattr(_flash_attn, "flash_attn_varlen_func")
|
||||
if attn_mask is None or is_causal:
|
||||
return True
|
||||
return attn_mask.dim() == 4
|
||||
return True
|
||||
|
||||
|
||||
def _resolve_default_backend() -> "AttentionBackend":
|
||||
"""Pick the highest-priority available backend (cuda -> flash -> torch).
|
||||
|
||||
Resolved lazily on first ``get_backend()`` and cached. Per-call
|
||||
capability fallback happens in ``attention()``, so the default is
|
||||
safe for training and fp32 models.
|
||||
"""
|
||||
return _priority_backends()[0]
|
||||
|
||||
|
||||
def _environment_backend() -> Optional["AttentionBackend"]:
|
||||
"""Resolve the process-wide ``ASTR_BACKEND`` override, if configured."""
|
||||
global _env_backend, _env_backend_name
|
||||
name = os.environ.get("ASTR_BACKEND", "").strip().lower()
|
||||
if not name:
|
||||
return None
|
||||
if name != _env_backend_name:
|
||||
with _default_backend_lock:
|
||||
if name != _env_backend_name:
|
||||
try:
|
||||
_env_backend = AttentionBackendFactory.create(name)
|
||||
except (ValueError, RuntimeError):
|
||||
_env_backend = None
|
||||
_env_backend_name = name
|
||||
return _env_backend
|
||||
|
||||
|
||||
def _resolve_backend(
|
||||
backend: Optional[Union[str, ATTN_BACKEND, "AttentionBackend", type]] = None,
|
||||
) -> "AttentionBackend":
|
||||
"""Resolve a backend configuration, defaulting to the process policy."""
|
||||
if backend is not None:
|
||||
if isinstance(backend, ATTN_BACKEND):
|
||||
return AttentionBackendFactory.create(backend.value)
|
||||
if isinstance(backend, str):
|
||||
return AttentionBackendFactory.create(backend)
|
||||
if isinstance(backend, type) and issubclass(backend, AttentionBackend):
|
||||
return backend()
|
||||
if isinstance(backend, AttentionBackend):
|
||||
return backend
|
||||
raise TypeError(
|
||||
f"expected a registered name, ATTN_BACKEND, AttentionBackend type, "
|
||||
f"or instance, got {type(backend).__name__}"
|
||||
)
|
||||
|
||||
global _default_backend
|
||||
if _default_backend is None:
|
||||
with _default_backend_lock:
|
||||
if _default_backend is None:
|
||||
_default_backend = _resolve_default_backend()
|
||||
return _default_backend
|
||||
|
||||
|
||||
def get_backend(
|
||||
use_default: bool = True,
|
||||
) -> Optional["AttentionBackend"]:
|
||||
"""Return the context override, optionally falling back to the process default.
|
||||
|
||||
``ASTR_BACKEND`` is a process-wide override and takes precedence over the
|
||||
context value. Pass ``use_default=False`` at request submission to retain
|
||||
only an environment override or the caller's :func:`attn_backend` value.
|
||||
"""
|
||||
return (
|
||||
_environment_backend()
|
||||
or _current_backend.get()
|
||||
or (_resolve_backend() if use_default else None)
|
||||
)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def attn_backend(backend: Union[str, ATTN_BACKEND, "AttentionBackend", type]):
|
||||
"""Context manager to select an attention backend.
|
||||
|
||||
Mirrors ``torch.nn.attention.sdpa_kernel``. Accepts an
|
||||
registered name, ``ATTN_BACKEND`` enum value, backend class, or instance.
|
||||
|
||||
Examples::
|
||||
|
||||
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
|
||||
...
|
||||
with attn_backend(TorchNativeBackend):
|
||||
...
|
||||
with attn_backend(TorchNativeBackend()):
|
||||
...
|
||||
"""
|
||||
instance = _resolve_backend(backend)
|
||||
token = _current_backend.set(instance)
|
||||
try:
|
||||
yield instance
|
||||
finally:
|
||||
_current_backend.reset(token)
|
||||
|
||||
|
||||
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
||||
"""Expand KV heads to match Q heads for GQA."""
|
||||
if n_rep == 1:
|
||||
return x
|
||||
n_heads, head_dim = x.shape[-2:]
|
||||
return (
|
||||
x.unsqueeze(-2)
|
||||
.expand(*x.shape[:-2], n_heads, n_rep, head_dim)
|
||||
.reshape(*x.shape[:-2], n_heads * n_rep, head_dim)
|
||||
)
|
||||
|
||||
|
||||
def attention(
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional["KVCache"] = None,
|
||||
layer_id: int = 0,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
fwd: Optional[str] = None,
|
||||
) -> Tensor:
|
||||
"""Functional attention entry point — mirrors ``F.scaled_dot_product_attention``.
|
||||
|
||||
Delegates to the active backend (set via ``with attn_backend(...)``).
|
||||
Handles KV cache I/O, GQA head expansion, and causal masking so the
|
||||
caller only needs to provide projected q/k/v.
|
||||
|
||||
Args:
|
||||
q: [batch, q_len, n_heads, head_dim] (blhd)
|
||||
k: [batch, q_len, n_kv_heads, head_dim] (blhd)
|
||||
v: [batch, q_len, n_kv_heads, head_dim] (blhd)
|
||||
kv_cache: cache dataclass, or None for training (no cache).
|
||||
layer_id: transformer layer index for buffer access.
|
||||
attn_mask: pre-built attention mask (SDPA-compatible).
|
||||
is_causal: whether to apply causal masking.
|
||||
|
||||
Returns:
|
||||
[batch, q_len, n_heads * head_dim]
|
||||
"""
|
||||
explicit = get_backend(use_default=False)
|
||||
backend = get_backend()
|
||||
if fwd is None and explicit is None:
|
||||
backend = TorchNativeBackend()
|
||||
if not _backend_supports(backend, q, kv_cache, attn_mask, is_causal, fwd):
|
||||
if explicit is not None:
|
||||
raise RuntimeError(
|
||||
f"Explicitly-set backend {type(backend).__name__} cannot "
|
||||
f"handle this attention call (shape={q.shape}, "
|
||||
f"dtype={q.dtype}, kv_cache={'none' if kv_cache is None else 'present'}, "
|
||||
f"attn_mask={'none' if attn_mask is None else 'present'}). "
|
||||
f"Remove the attn_backend() context or switch to a compatible backend."
|
||||
)
|
||||
for candidate in _priority_backends():
|
||||
if isinstance(candidate, type(backend)):
|
||||
continue
|
||||
if _backend_supports(candidate, q, kv_cache, attn_mask, is_causal, fwd):
|
||||
backend = candidate
|
||||
break
|
||||
return backend.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal, fwd)
|
||||
|
||||
|
||||
class AttentionBackend(ABC):
|
||||
"""Abstract base for attention computation strategies.
|
||||
|
||||
Subclasses implement ``fwd_decode`` (q_len == 1, with cache) and
|
||||
``fwd_prefill`` (q_len > 1, with or without cache). The public
|
||||
``forward`` method dispatches based on q_len.
|
||||
|
||||
Three equivalent ways to activate a backend::
|
||||
|
||||
with attn_backend(ATTN_BACKEND.TORCH_NATIVE): # enum
|
||||
...
|
||||
with attn_backend(TorchNativeBackend): # class
|
||||
...
|
||||
with TorchNativeBackend(): # instance
|
||||
...
|
||||
"""
|
||||
|
||||
def __enter__(self) -> "AttentionBackend":
|
||||
self._token = _current_backend.set(self)
|
||||
return self
|
||||
|
||||
def __exit__(self, *exc) -> None:
|
||||
_current_backend.reset(self._token)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional["KVCache"],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
fwd: Optional[str] = None,
|
||||
) -> Tensor:
|
||||
"""Dispatch to decode or extend based on q_len.
|
||||
|
||||
Args:
|
||||
q: [batch, q_len, n_heads, head_dim]
|
||||
k: [batch, q_len, n_kv_heads, head_dim]
|
||||
v: [batch, q_len, n_kv_heads, head_dim]
|
||||
kv_cache: cache dataclass, or None for training (no cache).
|
||||
layer_id: transformer layer index for buffer access.
|
||||
attn_mask: pre-built attention mask compatible with SDPA.
|
||||
is_causal: whether to apply causal masking.
|
||||
|
||||
Returns:
|
||||
[batch, q_len, n_heads * head_dim]
|
||||
"""
|
||||
if fwd == "decode":
|
||||
return self.fwd_decode(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||
if fwd == "prefill" or fwd is None:
|
||||
return self.fwd_prefill(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||
raise ValueError(f"unsupported attention forward mode: {fwd}")
|
||||
|
||||
@abstractmethod
|
||||
def fwd_decode(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional["KVCache"],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
"""Single-token decode with KV cache."""
|
||||
|
||||
@abstractmethod
|
||||
def fwd_prefill(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional["KVCache"],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
"""Multi-token prefill or training forward."""
|
||||
|
||||
@staticmethod
|
||||
def supports_graph() -> bool:
|
||||
"""Return True if this backend supports CUDA-graph capture.
|
||||
|
||||
Override in subclasses that can run under ``torch.cuda.graph``.
|
||||
|
||||
Called on the *active* backend instance (or its class) — a cheap
|
||||
boolean check with no side-effects.
|
||||
"""
|
||||
return False
|
||||
|
||||
|
||||
class AttentionBackendFactory(BaseFactory[AttentionBackend]):
|
||||
"""Factory for registered attention backends."""
|
||||
|
||||
|
||||
@AttentionBackendFactory.register(ATTN_BACKEND.TORCH_NATIVE.value)
|
||||
class TorchNativeBackend(AttentionBackend):
|
||||
"""Reference backend using torch SDPA with indirect KV cache indexing.
|
||||
|
||||
Writes new K/V into the cache buffers, gathers the full sequence K/V
|
||||
via ``req_to_token`` indirect indexing, then calls
|
||||
``F.scaled_dot_product_attention``.
|
||||
|
||||
For training (``kv_cache is None``), skips cache I/O entirely and
|
||||
runs SDPA directly on the projected q/k/v.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def supports(**kwargs) -> bool:
|
||||
return True
|
||||
|
||||
def fwd_decode(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional["KVCache"],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||
|
||||
def fwd_prefill(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional["KVCache"],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||
|
||||
def _forward(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional["KVCache"],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
if q.ndim == 4:
|
||||
n_rep = q.size(2) // k.size(2)
|
||||
if n_rep > 1:
|
||||
k = repeat_kv(k, n_rep)
|
||||
v = repeat_kv(v, n_rep)
|
||||
return (
|
||||
F.scaled_dot_product_attention(
|
||||
q.permute(0, 2, 1, 3),
|
||||
k.permute(0, 2, 1, 3),
|
||||
v.permute(0, 2, 1, 3),
|
||||
attn_mask,
|
||||
is_causal=is_causal,
|
||||
)
|
||||
.permute(0, 2, 1, 3)
|
||||
.contiguous()
|
||||
)
|
||||
|
||||
if kv_cache is None or kv_cache.qo_indptr is None:
|
||||
raise ValueError("packed attention requires KV cache metadata")
|
||||
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||
outputs = []
|
||||
n_rep = q.size(1) // k.size(1)
|
||||
for i in range(kv_cache.req_pool_indices.numel()):
|
||||
q_start = int(kv_cache.qo_indptr[i])
|
||||
q_end = int(kv_cache.qo_indptr[i + 1])
|
||||
indices = kv_cache.req_to_token[
|
||||
kv_cache.req_pool_indices[i], : kv_cache.seq_lens[i]
|
||||
]
|
||||
k_i = kv_cache.k_buffer[layer_id, indices]
|
||||
v_i = kv_cache.v_buffer[layer_id, indices]
|
||||
if n_rep > 1:
|
||||
k_i = repeat_kv(k_i, n_rep)
|
||||
v_i = repeat_kv(v_i, n_rep)
|
||||
q_len = q_end - q_start
|
||||
kv_len = k_i.size(0)
|
||||
q_pos = torch.arange(kv_len - q_len, kv_len, device=q.device)
|
||||
causal_mask = q_pos[:, None] >= torch.arange(kv_len, device=q.device)
|
||||
out = F.scaled_dot_product_attention(
|
||||
q[q_start:q_end].transpose(0, 1).unsqueeze(0),
|
||||
k_i.transpose(0, 1).unsqueeze(0),
|
||||
v_i.transpose(0, 1).unsqueeze(0),
|
||||
attn_mask=causal_mask,
|
||||
)
|
||||
outputs.append(out.squeeze(0).transpose(0, 1))
|
||||
return torch.cat(outputs)
|
||||
|
||||
|
||||
@AttentionBackendFactory.register(ATTN_BACKEND.CUDA.value)
|
||||
class CudaBackend(AttentionBackend):
|
||||
"""CUDA kernel backend with direct KV cache access.
|
||||
|
||||
Decode path: writes K/V to the flat pool, then calls
|
||||
``attn_paged_decode`` with req_to_token + kv_indptr.
|
||||
|
||||
Prefill path: writes K/V to the flat pool, then calls
|
||||
``attn_paged_prefill`` with ragged-batch support via qo_indptr +
|
||||
kv_indptr.
|
||||
|
||||
``kv_cache is None`` (training) raises — the per-call fallback to
|
||||
torch SDPA for training / fp32 / unsupported head_dim happens in the
|
||||
``attention()`` entry point.
|
||||
|
||||
Raises ``RuntimeError`` if the required kernel is not available.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def supports(**kwargs) -> bool:
|
||||
head_dim = kwargs.get("head_dim", -1)
|
||||
return (
|
||||
torch.cuda.is_available()
|
||||
and head_dim in (32, 64, 128, 256)
|
||||
and is_available("attn_paged_decode")
|
||||
and is_available("attn_paged_prefill")
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def supports_graph() -> bool:
|
||||
return True
|
||||
|
||||
def fwd_decode(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional["KVCache"],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
if kv_cache is None:
|
||||
raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
|
||||
|
||||
kv_indptr = kv_cache.kv_indptr
|
||||
|
||||
out = attn_paged_decode(
|
||||
q,
|
||||
kv_cache.k_buffer[layer_id],
|
||||
kv_cache.v_buffer[layer_id],
|
||||
kv_cache.req_to_token,
|
||||
kv_cache.req_pool_indices,
|
||||
kv_indptr,
|
||||
new_k=k,
|
||||
new_v=v,
|
||||
is_causal=True,
|
||||
o_part_buf=kv_cache.decode_o_part,
|
||||
ml_part_buf=kv_cache.decode_ml_part,
|
||||
out_buf=kv_cache.decode_out,
|
||||
)
|
||||
return out
|
||||
|
||||
def fwd_prefill(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional["KVCache"],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
if kv_cache is None:
|
||||
raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
|
||||
|
||||
loc = kv_cache.out_cache_loc
|
||||
kv_cache.k_buffer[layer_id, loc] = k
|
||||
kv_cache.v_buffer[layer_id, loc] = v
|
||||
|
||||
out = attn_paged_prefill(
|
||||
q,
|
||||
kv_cache.k_buffer[layer_id],
|
||||
kv_cache.v_buffer[layer_id],
|
||||
kv_cache.req_to_token,
|
||||
kv_cache.req_pool_indices,
|
||||
kv_cache.kv_indptr,
|
||||
kv_cache.qo_indptr,
|
||||
kv_cache.q_tile_to_batch,
|
||||
kv_cache.q_tile_to_index,
|
||||
attn_mask,
|
||||
is_causal=is_causal,
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
@AttentionBackendFactory.register(ATTN_BACKEND.FLASH.value)
|
||||
class FlashAttnBackend(AttentionBackend):
|
||||
"""FlashAttention backend via the optional ``flash-attn`` package.
|
||||
|
||||
Decode (q_len=1, contiguous cache): uses ``flash_attn_with_kvcache``,
|
||||
which reads K/V directly from the flat pool via cache_batch_idx +
|
||||
cache_seqlens — no materialized KV gather.
|
||||
|
||||
Prefill / non-contiguous decode: falls back to KV gather +
|
||||
``flash_attn_func``.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def supports(**kwargs) -> bool:
|
||||
return flash_attn_available()
|
||||
|
||||
def fwd_decode(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional["KVCache"],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
return self._forward_packed(q, k, v, kv_cache, layer_id)
|
||||
|
||||
def fwd_prefill(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: Optional["KVCache"],
|
||||
layer_id: int,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
if q.ndim == 3:
|
||||
return self._forward_packed(q, k, v, kv_cache, layer_id)
|
||||
return self._forward_dense(q, k, v, attn_mask, is_causal)
|
||||
|
||||
def _forward_dense(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
attn_mask: Optional[Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
n_rep = q.size(2) // k.size(2)
|
||||
if n_rep > 1:
|
||||
k = repeat_kv(k, n_rep)
|
||||
v = repeat_kv(v, n_rep)
|
||||
|
||||
if attn_mask is not None and not is_causal and attn_mask.dim() != 4:
|
||||
raise ValueError(
|
||||
"FlashAttnBackend does not support a custom attention mask; "
|
||||
"use a causal mask or select TorchNativeBackend."
|
||||
)
|
||||
fa = _flash_attn
|
||||
if fa is None:
|
||||
raise RuntimeError(
|
||||
"FlashAttnBackend requires the optional 'flash-attn' package. "
|
||||
"Install with `pip install flash-attn`."
|
||||
)
|
||||
out = fa.flash_attn_func(
|
||||
q.contiguous(),
|
||||
k.contiguous(),
|
||||
v.contiguous(),
|
||||
causal=is_causal or (attn_mask is not None and attn_mask.dim() == 4),
|
||||
)
|
||||
return out.contiguous()
|
||||
|
||||
def _forward_packed(
|
||||
self,
|
||||
q: Tensor,
|
||||
k: Tensor,
|
||||
v: Tensor,
|
||||
kv_cache: "KVCache",
|
||||
layer_id: int,
|
||||
) -> Tensor:
|
||||
fa = _flash_attn
|
||||
if fa is None or not hasattr(fa, "flash_attn_varlen_func"):
|
||||
raise RuntimeError("packed inference requires flash_attn_varlen_func")
|
||||
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||
page_table = kv_cache.req_to_token[
|
||||
kv_cache.req_pool_indices, : kv_cache.max_len
|
||||
]
|
||||
positions = torch.arange(kv_cache.max_len, device=q.device)
|
||||
indices = page_table[positions.unsqueeze(0) < kv_cache.seq_lens.unsqueeze(1)]
|
||||
k_flat = kv_cache.k_buffer[layer_id, indices].contiguous()
|
||||
v_flat = kv_cache.v_buffer[layer_id, indices].contiguous()
|
||||
out = fa.flash_attn_varlen_func(
|
||||
q.contiguous(),
|
||||
k_flat,
|
||||
v_flat,
|
||||
kv_cache.qo_indptr,
|
||||
kv_cache.kv_indptr,
|
||||
int((kv_cache.qo_indptr[1:] - kv_cache.qo_indptr[:-1]).max()),
|
||||
int(kv_cache.seq_lens.max()),
|
||||
dropout_p=0.0,
|
||||
causal=True,
|
||||
)
|
||||
return out
|
||||
@@ -0,0 +1,53 @@
|
||||
"""Rotary embedding with auto-dispatch to CUDA kernel.
|
||||
|
||||
Single entry point ``apply_rotary_emb(x, freqs_cis)`` — uses the fused
|
||||
CUDA kernel when available, falls back to torch complex multiply otherwise.
|
||||
|
||||
Layout: x is [batch, seq_len, n_heads, head_dim] (bf16).
|
||||
freqs_cis is [batch, seq_len, dim/2, 2] (f32) — [cos, sin] pairs.
|
||||
"""
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.extension.loader import is_available
|
||||
from astrai.extension.ops.rotary import rotary_emb as _cuda_rotary
|
||||
|
||||
_cache = {"available": None}
|
||||
|
||||
|
||||
def _cuda_available() -> bool:
|
||||
if _cache["available"] is None:
|
||||
_cache["available"] = is_available("rotary_emb")
|
||||
return _cache["available"]
|
||||
|
||||
|
||||
def _torch_apply(x: Tensor, freqs_cis: Tensor) -> Tensor:
|
||||
cos, sin = freqs_cis[..., 0], freqs_cis[..., 1]
|
||||
dtype = x.dtype
|
||||
x_ = x.float().reshape(*x.shape[:-1], -1, 2)
|
||||
x_complex = torch.view_as_complex(x_)
|
||||
freqs_cis_complex = torch.complex(cos, sin).unsqueeze(-2)
|
||||
x_rotated = x_complex * freqs_cis_complex
|
||||
x_out = torch.view_as_real(x_rotated).flatten(-2)
|
||||
return x_out.to(dtype)
|
||||
|
||||
|
||||
def apply_rotary_emb(x: Tensor, freqs_cis: Tensor) -> Tensor:
|
||||
"""Apply rotary embedding to x.
|
||||
|
||||
Args:
|
||||
x: [batch, seq_len, n_heads, head_dim] (bf16)
|
||||
freqs_cis: [batch, seq_len, dim/2, 2] (f32) — [cos, sin] pairs
|
||||
|
||||
Returns:
|
||||
[batch, seq_len, n_heads, head_dim] (bf16)
|
||||
"""
|
||||
if (
|
||||
_cuda_available()
|
||||
and not torch.is_grad_enabled()
|
||||
and x.is_cuda
|
||||
and x.dtype == torch.bfloat16
|
||||
):
|
||||
return _cuda_rotary(x, freqs_cis)
|
||||
return _torch_apply(x, freqs_cis)
|
||||
@@ -0,0 +1,334 @@
|
||||
"""FP8 training: scaling state and aten::linear dispatch.
|
||||
|
||||
Layered (see also ``ops/fp8.py`` for the CUDA interface adapter):
|
||||
|
||||
1. Kernel interface: ``ops.fp8`` - the only module touching the pybind.
|
||||
2. Training state (this module): per-tensor scales, amax history, delayed
|
||||
scaling, and the ``fp8_autocast`` context (TE-style, like
|
||||
``torch.autocast``).
|
||||
3. aten::linear integration (this module): registers the CUDA impl and the
|
||||
M/N alignment guard.
|
||||
|
||||
Usage::
|
||||
|
||||
from astrai.extension.fp8 import fp8_autocast
|
||||
|
||||
with fp8_autocast(enabled=True):
|
||||
logits = model(input_ids)
|
||||
loss.backward()
|
||||
|
||||
Importing this module registers the aten::linear CUDA implementation.
|
||||
"""
|
||||
|
||||
from contextlib import contextmanager
|
||||
|
||||
import torch
|
||||
from torch.library import Library
|
||||
|
||||
from astrai.extension.ops.fp8 import (
|
||||
linear_backward_scaled,
|
||||
linear_forward_scaled,
|
||||
)
|
||||
|
||||
E4M3_MAX = 448.0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Layer 2: training state (scales, amax history, delayed scaling, autocast)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class FP8TensorMeta:
|
||||
"""Scales + amax state for one weight tensor and its paired activations.
|
||||
|
||||
- weight: delayed scale from a 16-step amax history window (TE style)
|
||||
- x/g: delayed one step, reuse the quantize kernel's free atomic amax
|
||||
"""
|
||||
|
||||
__slots__ = (
|
||||
"scale",
|
||||
"scale_inv",
|
||||
"amax_history",
|
||||
"idx",
|
||||
"x_scale",
|
||||
"x_scale_inv",
|
||||
"x_history",
|
||||
"x_idx",
|
||||
"g_scale",
|
||||
"g_scale_inv",
|
||||
"g_history",
|
||||
"g_idx",
|
||||
"w_init",
|
||||
"x_init",
|
||||
"g_init",
|
||||
)
|
||||
|
||||
def __init__(self, device: torch.device, update_interval: int):
|
||||
self.scale = torch.ones(1, device=device, dtype=torch.float32)
|
||||
self.scale_inv = torch.ones(1, device=device, dtype=torch.float32)
|
||||
self.amax_history = torch.ones(
|
||||
update_interval, device=device, dtype=torch.float32
|
||||
)
|
||||
self.idx = 0
|
||||
self.x_scale = torch.ones(1, device=device, dtype=torch.float32)
|
||||
self.x_scale_inv = torch.ones(1, device=device, dtype=torch.float32)
|
||||
self.x_history = torch.ones(update_interval, device=device, dtype=torch.float32)
|
||||
self.x_idx = 0
|
||||
self.g_scale = torch.ones(1, device=device, dtype=torch.float32)
|
||||
self.g_scale_inv = torch.ones(1, device=device, dtype=torch.float32)
|
||||
self.g_history = torch.ones(update_interval, device=device, dtype=torch.float32)
|
||||
self.g_idx = 0
|
||||
self.w_init = False
|
||||
self.x_init = False
|
||||
self.g_init = False
|
||||
|
||||
def init_scale(self, t: torch.Tensor) -> None:
|
||||
"""Immediate scale from the current amax; used on the first call.
|
||||
|
||||
A scale of 1 would underflow small activations/gradients (e4m3 min
|
||||
normal is 2^-6); initialize from the actual amax once, then delayed
|
||||
updates take over.
|
||||
"""
|
||||
amax = t.abs().amax().to(torch.float32).clamp_min(1e-12)
|
||||
self.scale.copy_(amax / E4M3_MAX)
|
||||
self.scale_inv.copy_(E4M3_MAX / amax)
|
||||
self.record(amax)
|
||||
|
||||
def push_x_scale(self, amax: torch.Tensor) -> None:
|
||||
"""Window update for the activation scale (delayed, TE style)."""
|
||||
self.x_history[self.x_idx] = amax.reshape(())
|
||||
self.x_idx = (self.x_idx + 1) % self.x_history.numel()
|
||||
m = self.x_history.max()
|
||||
self.x_scale.copy_(m / E4M3_MAX)
|
||||
self.x_scale_inv.copy_(E4M3_MAX / m)
|
||||
|
||||
def push_g_scale(self, amax: torch.Tensor) -> None:
|
||||
"""Window update for the gradient scale (delayed, TE style)."""
|
||||
self.g_history[self.g_idx] = amax.reshape(())
|
||||
self.g_idx = (self.g_idx + 1) % self.g_history.numel()
|
||||
m = self.g_history.max()
|
||||
self.g_scale.copy_(m / E4M3_MAX)
|
||||
self.g_scale_inv.copy_(E4M3_MAX / m)
|
||||
|
||||
def record(self, amax: torch.Tensor) -> None:
|
||||
"""Push the latest amax into the ring buffer (device-side copy, no sync)."""
|
||||
self.amax_history[self.idx] = amax.reshape(())
|
||||
self.idx = (self.idx + 1) % self.amax_history.numel()
|
||||
|
||||
def refresh(self) -> None:
|
||||
"""Recompute scale from the amax history window (delayed scaling)."""
|
||||
amax = self.amax_history.max()
|
||||
if amax > 0:
|
||||
self.scale.copy_(amax / E4M3_MAX)
|
||||
self.scale_inv.copy_(E4M3_MAX / amax)
|
||||
|
||||
|
||||
class FP8State:
|
||||
"""Global fp8 training state, TE-style."""
|
||||
|
||||
def __init__(self, update_interval: int = 16):
|
||||
self.enabled = False
|
||||
self.update_interval = update_interval
|
||||
self.step_count = 0
|
||||
self._metas: dict[tuple, FP8TensorMeta] = {}
|
||||
self._last_device: torch.device | None = None
|
||||
|
||||
def _get_device(self, t: torch.Tensor) -> torch.device:
|
||||
if self._last_device is None:
|
||||
self._last_device = t.device
|
||||
return t.device
|
||||
|
||||
def get_weight_meta(self, w: torch.Tensor) -> FP8TensorMeta:
|
||||
key = (w.data_ptr(), w.shape, w.dtype)
|
||||
meta = self._metas.get(key)
|
||||
if meta is None:
|
||||
meta = FP8TensorMeta(self._get_device(w), self.update_interval)
|
||||
self._metas[key] = meta
|
||||
return meta
|
||||
|
||||
def step(self) -> None:
|
||||
"""Advance the counter and refresh all weight scales every N steps."""
|
||||
self.step_count += 1
|
||||
if self.step_count % self.update_interval == 0:
|
||||
for meta in self._metas.values():
|
||||
meta.refresh()
|
||||
|
||||
def reset(self) -> None:
|
||||
self.enabled = False
|
||||
self.step_count = 0
|
||||
self._metas.clear()
|
||||
self._last_device = None
|
||||
|
||||
|
||||
# Global singleton: autograd backward runs on the engine worker threads, so
|
||||
# thread-local state would lose the fp8 flag during loss.backward(). The GIL
|
||||
# protects Python-side mutation; the CUDA kernels take their own mutex.
|
||||
_state = FP8State()
|
||||
|
||||
|
||||
def fp8_state() -> FP8State:
|
||||
return _state
|
||||
|
||||
|
||||
@contextmanager
|
||||
def fp8_autocast(enabled: bool = True, update_interval: int = 16):
|
||||
"""Autocast-style context: fp8 linear dispatch on this thread.
|
||||
|
||||
Usage::
|
||||
|
||||
with fp8_autocast(enabled=True):
|
||||
logits = model(input_ids) # aten::linear -> fp8 path
|
||||
loss.backward()
|
||||
|
||||
The scale-update counter advances once per ``enter`` (one training step),
|
||||
refreshing weight scales from their amax history every ``update_interval``.
|
||||
"""
|
||||
state = fp8_state()
|
||||
prev_enabled = state.enabled
|
||||
prev_interval = state.update_interval
|
||||
state.enabled = enabled
|
||||
state.update_interval = update_interval
|
||||
try:
|
||||
if enabled:
|
||||
state.step()
|
||||
yield
|
||||
finally:
|
||||
state.enabled = prev_enabled
|
||||
state.update_interval = prev_interval
|
||||
|
||||
|
||||
def fp8_linear_forward(x: torch.Tensor, w: torch.Tensor, bias=None):
|
||||
"""TE-style scaled fp8 linear forward (called from the aten::linear impl).
|
||||
|
||||
x uses the delayed scale of its paired weight meta (amax from the previous
|
||||
forward of this linear); the quantize kernel emits the current amax for the
|
||||
next step. No extra abs/max reduce.
|
||||
"""
|
||||
if bias is None:
|
||||
bias = torch.empty(0, device=x.device, dtype=x.dtype)
|
||||
state = fp8_state()
|
||||
meta = state.get_weight_meta(w)
|
||||
if not meta.w_init:
|
||||
meta.init_scale(w)
|
||||
meta.w_init = True
|
||||
if not meta.x_init:
|
||||
amax = x.abs().amax().to(torch.float32).clamp_min(1e-12)
|
||||
meta.x_history.fill_(amax)
|
||||
meta.x_scale.copy_(amax / E4M3_MAX)
|
||||
meta.x_scale_inv.copy_(E4M3_MAX / amax)
|
||||
meta.x_init = True
|
||||
amax_x = torch.empty(1, device=x.device, dtype=torch.float32)
|
||||
amax_w = torch.empty(1, device=x.device, dtype=torch.float32)
|
||||
out = linear_forward_scaled(
|
||||
x,
|
||||
w,
|
||||
bias,
|
||||
meta.x_scale,
|
||||
meta.scale,
|
||||
meta.x_scale_inv,
|
||||
meta.scale_inv,
|
||||
amax_x,
|
||||
amax_w,
|
||||
)
|
||||
meta.record(amax_w)
|
||||
meta.push_x_scale(amax_x)
|
||||
return out
|
||||
|
||||
|
||||
def fp8_linear_backward(g, x, w, masks):
|
||||
"""TE-style scaled fp8 linear backward (called from aten::linear_backward)."""
|
||||
state = fp8_state()
|
||||
meta = state.get_weight_meta(w)
|
||||
if not meta.g_init:
|
||||
amax = g.abs().amax().to(torch.float32).clamp_min(1e-12)
|
||||
meta.g_history.fill_(amax)
|
||||
meta.g_scale.copy_(amax / E4M3_MAX)
|
||||
meta.g_scale_inv.copy_(E4M3_MAX / amax)
|
||||
meta.g_init = True
|
||||
amax_g = torch.empty(1, device=g.device, dtype=torch.float32)
|
||||
out = linear_backward_scaled(
|
||||
g,
|
||||
x,
|
||||
w,
|
||||
masks,
|
||||
meta.g_scale,
|
||||
meta.scale,
|
||||
meta.x_scale,
|
||||
meta.g_scale_inv,
|
||||
meta.scale_inv,
|
||||
meta.x_scale_inv,
|
||||
amax_g,
|
||||
)
|
||||
meta.push_g_scale(amax_g)
|
||||
return out
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Layer 3: aten::linear integration
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def fp8_linear_enable(enabled: bool = True) -> None:
|
||||
"""Toggle fp8 dispatch for aten::linear (global; backward runs on engine
|
||||
worker threads, so a thread-local flag would be lost during backward)."""
|
||||
fp8_state().enabled = enabled
|
||||
|
||||
|
||||
def fp8_linear_enabled() -> bool:
|
||||
return fp8_state().enabled
|
||||
|
||||
|
||||
def _fp8_supported(x: torch.Tensor, w: torch.Tensor) -> bool:
|
||||
"""cuBLASLt fp8 requires M % 16 == 0 and N % 16 == 0 (K is padded)."""
|
||||
m = x.numel() // x.size(-1)
|
||||
return m % 16 == 0 and w.size(0) % 16 == 0
|
||||
|
||||
|
||||
def _linear_cuda_impl(x: torch.Tensor, w: torch.Tensor, bias=None):
|
||||
if (
|
||||
fp8_linear_enabled()
|
||||
and x.dtype == torch.bfloat16
|
||||
and w.dtype == torch.bfloat16
|
||||
and _fp8_supported(x, w)
|
||||
):
|
||||
return fp8_linear_forward(x, w, bias)
|
||||
return torch.ops.aten.linear.default.redispatch(
|
||||
torch._C.DispatchKeySet(torch._C.DispatchKey.CompositeImplicitAutograd),
|
||||
x,
|
||||
w,
|
||||
bias,
|
||||
)
|
||||
|
||||
|
||||
def _linear_backward_cuda_impl(input_tensor, grad_output, weight, output_mask):
|
||||
if (
|
||||
fp8_linear_enabled()
|
||||
and weight.dtype == torch.bfloat16
|
||||
and _fp8_supported(grad_output, weight)
|
||||
):
|
||||
return fp8_linear_backward(grad_output, input_tensor, weight, list(output_mask))
|
||||
compute_dtype = weight.dtype
|
||||
grad = grad_output.to(compute_dtype)
|
||||
grad_2d = grad.reshape(-1, weight.size(0))
|
||||
input_2d = input_tensor.reshape(-1, input_tensor.size(-1)).to(compute_dtype)
|
||||
grad_input = (
|
||||
torch.mm(grad_2d, weight)
|
||||
if output_mask[0]
|
||||
else torch.empty(0, device=input_tensor.device, dtype=input_tensor.dtype)
|
||||
)
|
||||
grad_weight = (
|
||||
torch.mm(grad_2d.t(), input_2d)
|
||||
if output_mask[1]
|
||||
else torch.empty(0, device=input_tensor.device, dtype=input_tensor.dtype)
|
||||
)
|
||||
grad_bias = (
|
||||
grad.sum(dim=0)
|
||||
if output_mask[2]
|
||||
else torch.empty(0, device=input_tensor.device, dtype=input_tensor.dtype)
|
||||
)
|
||||
return grad_input.reshape_as(input_tensor), grad_weight, grad_bias
|
||||
|
||||
|
||||
_lib = Library("aten", "IMPL", "CUDA")
|
||||
_lib.impl("linear", _linear_cuda_impl)
|
||||
_lib.impl("linear_backward", _linear_backward_cuda_impl)
|
||||
@@ -0,0 +1 @@
|
||||
"""Compiled CUDA kernel modules (``*.so``) live here, kept separate from Python source."""
|
||||
@@ -0,0 +1,43 @@
|
||||
"""Dynamic discovery and loading of compiled CUDA kernel modules.
|
||||
|
||||
Each kernel is registered in ``csrc/build.py`` and built into a ``.so`` placed
|
||||
in this package directory. On import we try to load each one; kernels that
|
||||
failed to build (or are running on a CPU-only machine) are marked unavailable
|
||||
so the wrapper functions can fall back to ``torch`` SDPA.
|
||||
"""
|
||||
|
||||
import importlib
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
KERNEL_NAMES = [
|
||||
"attn_decode",
|
||||
"attn_prefill",
|
||||
"attn_paged_decode",
|
||||
"attn_paged_prefill",
|
||||
"rotary_emb",
|
||||
"fp8_mm",
|
||||
]
|
||||
|
||||
_available: dict[str, bool] = {}
|
||||
_modules: dict[str, object] = {}
|
||||
|
||||
for _name in KERNEL_NAMES:
|
||||
try:
|
||||
_mod = importlib.import_module(f".lib.{_name}", package=__package__)
|
||||
_available[_name] = True
|
||||
_modules[_name] = _mod
|
||||
except ImportError:
|
||||
_available[_name] = False
|
||||
_modules[_name] = None
|
||||
|
||||
|
||||
def is_available(name: str) -> bool:
|
||||
"""Return ``True`` if the compiled kernel ``name`` was loaded."""
|
||||
return _available.get(name, False)
|
||||
|
||||
|
||||
def get_module(name: str) -> object:
|
||||
"""Return the loaded kernel module for ``name``, or ``None`` if unavailable."""
|
||||
return _modules.get(name)
|
||||
@@ -0,0 +1,19 @@
|
||||
"""Stateless wrappers around compiled extension kernels."""
|
||||
|
||||
from astrai.extension.ops.attention import (
|
||||
TensorLayout,
|
||||
attn_decode,
|
||||
attn_paged_decode,
|
||||
attn_paged_prefill,
|
||||
attn_prefill,
|
||||
)
|
||||
from astrai.extension.ops.rotary import rotary_emb
|
||||
|
||||
__all__ = [
|
||||
"TensorLayout",
|
||||
"attn_decode",
|
||||
"attn_paged_decode",
|
||||
"attn_paged_prefill",
|
||||
"attn_prefill",
|
||||
"rotary_emb",
|
||||
]
|
||||
@@ -0,0 +1,200 @@
|
||||
"""Attention kernel wrapper functions - one entry point per compiled kernel.
|
||||
|
||||
Each wrapper calls its CUDA kernel directly. If the kernel is not
|
||||
available, raises ``RuntimeError``. Fallback to torch SDPA is the
|
||||
responsibility of the attention backend, not this module.
|
||||
|
||||
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
||||
(blhd). Scale is always ``1/sqrt(head_dim)``.
|
||||
|
||||
Interface (all functions):
|
||||
is_causal: True = causal mask; False = non-causal
|
||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
|
||||
"""
|
||||
|
||||
import enum
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.extension.loader import _available, _modules
|
||||
|
||||
|
||||
class TensorLayout(enum.IntEnum):
|
||||
"""Q/K/V tensor layout, mirrors the C++ ``TensorLayout`` enum in ``attn_common.h``.
|
||||
|
||||
Kernels internally operate on BHLD; BLHD inputs are transposed at entry.
|
||||
"""
|
||||
|
||||
BHLD = 0 # [batch, n_heads, seq_len, head_dim]
|
||||
BLHD = 1 # [batch, seq_len, n_heads, head_dim]
|
||||
|
||||
|
||||
def _check_available(name: str):
|
||||
if not _available.get(name):
|
||||
raise RuntimeError(
|
||||
f"CUDA kernel '{name}' is not available. "
|
||||
f"Build with CSRC_KERNELS=true or use a torch-native backend."
|
||||
)
|
||||
|
||||
|
||||
def attn_decode(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: Optional[torch.Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""GQA decode attention (q_len == 1).
|
||||
|
||||
Args:
|
||||
q: [batch, 1, n_heads, head_dim] (blhd, bf16)
|
||||
k: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||
v: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||
mask: 2D [batch, kv_len] or 3D [batch, 1, kv_len] (bool, True=keep)
|
||||
is_causal: apply causal mask
|
||||
|
||||
Returns:
|
||||
[batch, 1, n_heads, head_dim] (blhd, bf16)
|
||||
"""
|
||||
_check_available("attn_decode")
|
||||
causal_offset = (k.size(1) - 1) if is_causal else -1
|
||||
return _modules["attn_decode"].attn_decode(
|
||||
q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD
|
||||
)
|
||||
|
||||
|
||||
def attn_prefill(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: Optional[torch.Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""GQA prefill attention (q_len > 1).
|
||||
|
||||
Args:
|
||||
q: [batch, q_len, n_heads, head_dim] (blhd, bf16)
|
||||
k: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||
v: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
|
||||
is_causal: apply causal mask
|
||||
|
||||
Returns:
|
||||
[batch, q_len, n_heads, head_dim] (blhd, bf16)
|
||||
"""
|
||||
_check_available("attn_prefill")
|
||||
causal_offset = (k.size(1) - q.size(1)) if is_causal else -1
|
||||
return _modules["attn_prefill"].attn_prefill(
|
||||
q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD
|
||||
)
|
||||
|
||||
|
||||
def attn_paged_decode(
|
||||
q: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
v_cache: torch.Tensor,
|
||||
req_to_token: torch.Tensor,
|
||||
req_pool_indices: torch.Tensor,
|
||||
kv_indptr: torch.Tensor,
|
||||
new_k: Optional[torch.Tensor] = None,
|
||||
new_v: Optional[torch.Tensor] = None,
|
||||
mask: Optional[torch.Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
o_part_buf: Optional[torch.Tensor] = None,
|
||||
ml_part_buf: Optional[torch.Tensor] = None,
|
||||
out_buf: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
"""SGLang-style paged decode (q_len == 1, flat KV pool).
|
||||
|
||||
Reads K/V directly from a flat pool [size, kv_head, head_dim] via
|
||||
req_to_token indirect indexing. Each request has its own seq_len
|
||||
(from kv_indptr), eliminating padding waste.
|
||||
|
||||
Args:
|
||||
q: [batch, n_heads, head_dim] (bf16, 3D — no seq dim)
|
||||
k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat)
|
||||
v_cache: same as k_cache
|
||||
req_to_token: [num_reqs, max_context_len] (int32) — token -> slot
|
||||
req_pool_indices: [batch] (int32) — rows into req_to_token
|
||||
kv_indptr: [batch+1] (int32) — prefix sum of per-request seq_lens
|
||||
new_k: current-token K to append, [batch, n_kv_heads, head_dim]
|
||||
new_v: current-token V to append, same shape as new_k
|
||||
mask: 2D [batch, max_context_len] (bool, True=keep) or None
|
||||
is_causal: apply causal mask
|
||||
o_part_buf: pre-allocated split-KV o partial buffer (workflow bypass)
|
||||
ml_part_buf: pre-allocated split-KV m/l buffer (workflow bypass)
|
||||
out_buf: pre-allocated output buffer [batch, n_heads, head_dim] (graph-safe)
|
||||
|
||||
Returns:
|
||||
[batch, n_heads, head_dim] (bf16, 3D)
|
||||
"""
|
||||
_check_available("attn_paged_decode")
|
||||
causal_offset = 0 if is_causal else -1
|
||||
return _modules["attn_paged_decode"].attn_paged_decode(
|
||||
q,
|
||||
k_cache,
|
||||
v_cache,
|
||||
req_to_token,
|
||||
req_pool_indices,
|
||||
kv_indptr,
|
||||
new_k=new_k,
|
||||
new_v=new_v,
|
||||
mask=mask,
|
||||
causal_offset=causal_offset,
|
||||
o_part_buf=o_part_buf,
|
||||
ml_part_buf=ml_part_buf,
|
||||
out_buf=out_buf,
|
||||
)
|
||||
|
||||
|
||||
def attn_paged_prefill(
|
||||
q: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
v_cache: torch.Tensor,
|
||||
req_to_token: torch.Tensor,
|
||||
req_pool_indices: torch.Tensor,
|
||||
kv_indptr: torch.Tensor,
|
||||
qo_indptr: torch.Tensor,
|
||||
q_tile_to_batch: torch.Tensor,
|
||||
q_tile_to_index: torch.Tensor,
|
||||
mask: Optional[torch.Tensor] = None,
|
||||
is_causal: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""SGLang-style paged prefill (ragged batch, flat KV pool).
|
||||
|
||||
Reads K/V directly from a flat pool [size, kv_head, head_dim] via
|
||||
req_to_token. Supports ragged batches: each request has its own
|
||||
q_len and kv_len, addressed via qo_indptr and kv_indptr.
|
||||
|
||||
Args:
|
||||
q: [total_q, n_heads, head_dim] (bf16, 3D — flattened across requests)
|
||||
k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat)
|
||||
v_cache: same as k_cache
|
||||
req_to_token: [num_reqs, max_context_len] (int32)
|
||||
req_pool_indices: [batch] (int32)
|
||||
kv_indptr: [batch+1] (int32) — prefix sum of per-request kv_lens
|
||||
qo_indptr: [batch+1] (int32) — prefix sum of per-request q_lens
|
||||
q_tile_to_batch: [num_q_tiles] (int32) — request index per Q tile
|
||||
q_tile_to_index: [num_q_tiles] (int32) — local Q tile index per request
|
||||
mask: 4D [batch, 1, q_len, kv_len] (bool, True=keep) or None
|
||||
is_causal: apply causal mask
|
||||
|
||||
Returns:
|
||||
[total_q, n_heads, head_dim] (bf16, 3D)
|
||||
"""
|
||||
_check_available("attn_paged_prefill")
|
||||
causal_offset = 0 if is_causal else -1
|
||||
return _modules["attn_paged_prefill"].attn_paged_prefill(
|
||||
q,
|
||||
k_cache,
|
||||
v_cache,
|
||||
req_to_token,
|
||||
req_pool_indices,
|
||||
kv_indptr,
|
||||
qo_indptr,
|
||||
q_tile_to_batch,
|
||||
q_tile_to_index,
|
||||
mask,
|
||||
causal_offset=causal_offset,
|
||||
)
|
||||
@@ -0,0 +1,127 @@
|
||||
"""FP8 CUDA kernel interface adapter (the only module touching the pybind.
|
||||
|
||||
Isolates the ``fp8_mm`` CUDA extension behind stable Python functions:
|
||||
- availability / dtype checks and clear errors
|
||||
- torch.library ``custom::fp8_mm`` registration (meta + CPU fallback)
|
||||
- quantize-in-GEMM primitives used by ``fp8.py`` training state
|
||||
|
||||
Policy (scales, amax history, delayed scaling, autocast) lives in ``fp8.py``;
|
||||
this module is stateless.
|
||||
"""
|
||||
|
||||
import torch
|
||||
from torch.library import custom_op
|
||||
|
||||
from astrai.extension.loader import get_module, is_available
|
||||
|
||||
|
||||
def _mod():
|
||||
if not is_available("fp8_mm"):
|
||||
raise RuntimeError(
|
||||
"CUDA kernel 'fp8_mm' is not available. Build with CSRC_KERNELS=true."
|
||||
)
|
||||
return get_module("fp8_mm")
|
||||
|
||||
|
||||
@custom_op("custom::fp8_mm", mutates_args=())
|
||||
def fp8_mm(
|
||||
a: torch.Tensor, b: torch.Tensor, sx: torch.Tensor, sw: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
"""BF16 inputs, fused FP8 GEMM with FP32 accumulation and BF16 output."""
|
||||
|
||||
|
||||
@fp8_mm.register_fake
|
||||
def _fp8_mm_fake(a, b, sx, sw):
|
||||
return torch.empty((a.size(0), b.size(0)), device=a.device, dtype=torch.bfloat16)
|
||||
|
||||
|
||||
@fp8_mm.register_kernel("cuda")
|
||||
def _fp8_mm_cuda(a, b, sx, sw):
|
||||
if not (a.dtype == torch.bfloat16 and b.dtype == torch.bfloat16):
|
||||
raise TypeError(f"bf16 GEMM requires bf16 inputs, got {a.dtype}/{b.dtype}")
|
||||
return _mod().fp8_mm(a, b, sx, sw)
|
||||
|
||||
|
||||
@fp8_mm.register_kernel("cpu")
|
||||
def _fp8_mm_cpu(a, b, sx, sw):
|
||||
return torch.mm(a.float(), b.float().t()).to(torch.bfloat16)
|
||||
|
||||
|
||||
@custom_op("custom::fp8_mm_prequant", mutates_args=())
|
||||
def fp8_mm_prequant(
|
||||
a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
"""Pre-quantized FP8 inputs, fused FP8 GEMM, FP32 accumulation, BF16 out."""
|
||||
|
||||
|
||||
@fp8_mm_prequant.register_fake
|
||||
def _fp8_mm_prequant_fake(a, b, scale):
|
||||
return torch.empty((a.size(0), b.size(0)), device=a.device, dtype=torch.bfloat16)
|
||||
|
||||
|
||||
@fp8_mm_prequant.register_kernel("cuda")
|
||||
def _fp8_mm_prequant_cuda(a, b, scale):
|
||||
if not (a.dtype == torch.float8_e4m3fn and b.dtype == torch.float8_e4m3fn):
|
||||
raise TypeError(
|
||||
f"pre-quantized FP8 GEMM requires fp8 inputs, got {a.dtype}/{b.dtype}"
|
||||
)
|
||||
return _mod().fp8_mm_prequant(a, b, scale)
|
||||
|
||||
|
||||
@fp8_mm_prequant.register_kernel("cpu")
|
||||
def _fp8_mm_prequant_cpu(a, b, scale):
|
||||
return (a.float() @ b.float().t() * scale).to(torch.bfloat16)
|
||||
|
||||
|
||||
@custom_op("custom::fp8_mm_prequant_fp8", mutates_args=())
|
||||
def fp8_mm_prequant_fp8(
|
||||
a: torch.Tensor, b: torch.Tensor, scale: torch.Tensor, out_scale: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
"""FP8 inputs and FP8 output: fused FP8 GEMM with FP32 accumulation."""
|
||||
|
||||
|
||||
@fp8_mm_prequant_fp8.register_fake
|
||||
def _fp8_mm_prequant_fp8_fake(a, b, scale, out_scale):
|
||||
return torch.empty((a.size(0), b.size(0)), device=a.device, dtype=a.dtype)
|
||||
|
||||
|
||||
@fp8_mm_prequant_fp8.register_kernel("cuda")
|
||||
def _fp8_mm_prequant_fp8_cuda(a, b, scale, out_scale):
|
||||
if not (a.dtype == torch.float8_e4m3fn and b.dtype == torch.float8_e4m3fn):
|
||||
raise TypeError(
|
||||
f"pre-quantized FP8 GEMM requires fp8 inputs, got {a.dtype}/{b.dtype}"
|
||||
)
|
||||
return _mod().fp8_mm_prequant_fp8(a, b, scale, out_scale)
|
||||
|
||||
|
||||
@fp8_mm_prequant_fp8.register_kernel("cpu")
|
||||
def _fp8_mm_prequant_fp8_cpu(a, b, scale, out_scale):
|
||||
return (a.float() @ b.float().t() * scale * out_scale).to(torch.float8_e4m3fn)
|
||||
|
||||
|
||||
def linear_forward_scaled(x, w, bias, sx, sw, sx_inv, sw_inv, amax_x, amax_w):
|
||||
"""Quantize BF16 inputs to FP8, accumulate in FP32, and return BF16.
|
||||
|
||||
x/w: [..., K] / [N, K] bf16; sx/sw and their inverses control the fused
|
||||
E4M3 conversion; amax_x/amax_w receive the input max-abs values.
|
||||
"""
|
||||
if not (x.dtype == torch.bfloat16 and w.dtype == torch.bfloat16):
|
||||
raise TypeError(f"fp8 forward requires bf16 inputs, got {x.dtype}/{w.dtype}")
|
||||
return _mod().fp8_linear_forward_scaled(
|
||||
x, w, bias, sx, sw, sx_inv, sw_inv, amax_x, amax_w
|
||||
)
|
||||
|
||||
|
||||
def linear_backward_scaled(g, x, w, masks, sg, sw, sx, sg_inv, sw_inv, sx_inv, amax_g):
|
||||
"""dX = g @ W, dW = g^T @ X, dB = sum(g) with per-tensor scales."""
|
||||
if not (
|
||||
g.dtype == torch.bfloat16
|
||||
and x.dtype == torch.bfloat16
|
||||
and w.dtype == torch.bfloat16
|
||||
):
|
||||
raise TypeError(
|
||||
f"fp8 backward requires bf16 inputs, got {g.dtype}/{x.dtype}/{w.dtype}"
|
||||
)
|
||||
return _mod().fp8_linear_backward_scaled(
|
||||
g, x, w, masks, sg, sw, sx, sg_inv, sw_inv, sx_inv, amax_g
|
||||
)
|
||||
@@ -0,0 +1,39 @@
|
||||
"""Rotary embedding CUDA kernel wrapper.
|
||||
|
||||
Calls the compiled CUDA kernel directly. If the kernel is not available,
|
||||
raises ``RuntimeError``. Fallback to torch complex multiply is the
|
||||
responsibility of ``astrai.extension.backend.rotary.apply_rotary_emb``.
|
||||
|
||||
Layout: x is packed [tokens, n_heads, head_dim] or dense
|
||||
[batch, seq_len, n_heads, head_dim]. ``freqs_cis`` has matching token axes.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.extension.loader import _available, _modules
|
||||
|
||||
|
||||
def _check_available():
|
||||
if not _available.get("rotary_emb"):
|
||||
raise RuntimeError(
|
||||
"CUDA kernel 'rotary_emb' is not available. "
|
||||
"Build with CSRC_KERNELS=true or use the torch fallback."
|
||||
)
|
||||
|
||||
|
||||
def rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
|
||||
"""Fused rotary embedding kernel.
|
||||
|
||||
Args:
|
||||
x: packed 3D or dense 4D bf16 tensor.
|
||||
freqs_cis: matching token axes followed by [head_dim/2, 2].
|
||||
|
||||
Returns:
|
||||
Tensor with the same shape as ``x``.
|
||||
"""
|
||||
_check_available()
|
||||
if not x.is_contiguous():
|
||||
x = x.contiguous()
|
||||
if not freqs_cis.is_contiguous():
|
||||
freqs_cis = freqs_cis.contiguous()
|
||||
return _modules["rotary_emb"].rotary_emb(x, freqs_cis)
|
||||
+102
-139
@@ -1,190 +1,153 @@
|
||||
"""Base factory class for extensible component registration."""
|
||||
"""Base factory with decorator-based registration and kwarg-filtered instantiation."""
|
||||
|
||||
import inspect
|
||||
import sys
|
||||
from abc import ABC
|
||||
from typing import Callable, Dict, Generic, List, Optional, Tuple, Type, TypeVar
|
||||
from typing import (
|
||||
Callable,
|
||||
Dict,
|
||||
ForwardRef,
|
||||
Generic,
|
||||
List,
|
||||
Optional,
|
||||
Type,
|
||||
TypeVar,
|
||||
Union,
|
||||
get_args,
|
||||
get_origin,
|
||||
)
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
class Registry:
|
||||
"""Flexible registry for component classes with category and priority support.
|
||||
def _resolve_base_type(
|
||||
arg: Union[Type, str, ForwardRef], factory_cls: type
|
||||
) -> Optional[Type]:
|
||||
"""Resolve the generic type-arg T to a concrete class.
|
||||
|
||||
This registry stores component classes with optional metadata (category, priority).
|
||||
It provides methods for registration, retrieval, and listing with filtering.
|
||||
- Concrete class (``BaseFactory[MyBase]``): returned directly.
|
||||
- Forward reference (``BaseFactory["MyBase"]``): ``Base["X"]``
|
||||
produces a ``ForwardRef("X")`` at class-creation time. We
|
||||
extract the name and evaluate it in the factory module's
|
||||
global namespace — the same mechanism ``typing.get_type_hints``
|
||||
uses internally.
|
||||
"""
|
||||
if isinstance(arg, type):
|
||||
return arg
|
||||
|
||||
def __init__(self):
|
||||
self._entries = {} # name -> (component_cls, category, priority)
|
||||
if isinstance(arg, str):
|
||||
name = arg
|
||||
elif isinstance(arg, ForwardRef):
|
||||
name = arg.__forward_arg__
|
||||
else:
|
||||
return None
|
||||
|
||||
def register(
|
||||
self,
|
||||
name: str,
|
||||
component_cls: Type,
|
||||
category: Optional[str] = None,
|
||||
priority: int = 0,
|
||||
) -> None:
|
||||
"""Register a component class with optional category and priority."""
|
||||
if name in self._entries:
|
||||
raise ValueError(f"Component '{name}' is already registered")
|
||||
self._entries[name] = (component_cls, category, priority)
|
||||
mod = sys.modules.get(factory_cls.__module__)
|
||||
if mod is None:
|
||||
return None
|
||||
try:
|
||||
return eval(name, vars(mod)) # noqa: S307
|
||||
except NameError:
|
||||
return None
|
||||
|
||||
def get(self, name: str) -> Type:
|
||||
"""Get component class by name."""
|
||||
if name not in self._entries:
|
||||
raise KeyError(f"Component '{name}' not found in registry")
|
||||
return self._entries[name][0]
|
||||
|
||||
def get_with_metadata(self, name: str) -> Tuple[Type, Optional[str], int]:
|
||||
"""Get component class with its metadata."""
|
||||
entry = self._entries.get(name)
|
||||
if entry is None:
|
||||
raise KeyError(f"Component '{name}' not found in registry")
|
||||
return entry
|
||||
def _validate_component(component_cls: Type, base: Optional[Type]) -> None:
|
||||
"""Validate that *component_cls* inherits from *base*.
|
||||
|
||||
def contains(self, name: str) -> bool:
|
||||
"""Check if a name is registered."""
|
||||
return name in self._entries
|
||||
|
||||
def list_names(self) -> List[str]:
|
||||
"""Return list of registered component names."""
|
||||
return sorted(self._entries.keys())
|
||||
|
||||
def list_by_category(self, category: str) -> List[str]:
|
||||
"""Return names of components belonging to a specific category."""
|
||||
return sorted(
|
||||
name for name, (_, cat, _) in self._entries.items() if cat == category
|
||||
)
|
||||
|
||||
def list_by_priority(self, reverse: bool = False) -> List[str]:
|
||||
"""Return names sorted by priority (default ascending)."""
|
||||
return sorted(
|
||||
self._entries.keys(),
|
||||
key=lambda name: self._entries[name][2],
|
||||
reverse=reverse,
|
||||
)
|
||||
|
||||
def entries(self) -> Dict[str, Tuple[Type, Optional[str], int]]:
|
||||
"""Return raw entries dictionary."""
|
||||
return self._entries.copy()
|
||||
No-op when *base* is ``None`` (e.g. forward-ref resolution failed).
|
||||
"""
|
||||
if base is not None and not issubclass(component_cls, base):
|
||||
raise TypeError(f"{component_cls.__name__} must inherit from {base.__name__}")
|
||||
|
||||
|
||||
class BaseFactory(ABC, Generic[T]):
|
||||
"""Generic factory class for component registration and creation.
|
||||
"""Generic factory with decorator-based registration.
|
||||
|
||||
This base class provides a decorator-based registration pattern
|
||||
for creating extensible component factories.
|
||||
Create a factory by subclassing with the desired base type::
|
||||
|
||||
Example usage:
|
||||
class MyFactory(BaseFactory[MyBaseClass]):
|
||||
class MyFactory(BaseFactory[MyBase]):
|
||||
pass
|
||||
|
||||
Register components with the ``register`` decorator::
|
||||
|
||||
@MyFactory.register("custom")
|
||||
class CustomComponent(MyBaseClass):
|
||||
class CustomComponent(MyBase):
|
||||
...
|
||||
|
||||
component = MyFactory.create("custom", *args, **kwargs)
|
||||
obj = MyFactory.create("custom", *args, **kwargs)
|
||||
|
||||
``create()`` filters kwargs to match the component's ``__init__``
|
||||
signature so components don't need ``**kwargs`` just to absorb
|
||||
unrelated parameters.
|
||||
"""
|
||||
|
||||
_registry: Registry
|
||||
_entries: Dict[str, Type[T]]
|
||||
|
||||
def __init_subclass__(cls, **kwargs):
|
||||
super().__init_subclass__(**kwargs)
|
||||
cls._registry = Registry()
|
||||
for orig_base in getattr(cls, "__orig_bases__", ()):
|
||||
if get_origin(orig_base) is BaseFactory:
|
||||
(arg,) = get_args(orig_base)
|
||||
cls._entries = {}
|
||||
cls._component_base = _resolve_base_type(arg, cls)
|
||||
return
|
||||
|
||||
@classmethod
|
||||
def register(
|
||||
cls, name: str, category: Optional[str] = None, priority: int = 0
|
||||
) -> Callable[[Type[T]], Type[T]]:
|
||||
"""Decorator to register a component class with optional category and priority.
|
||||
def register(cls, name: str) -> Callable[[Type[T]], Type[T]]:
|
||||
"""Decorator to register a component class.
|
||||
|
||||
Args:
|
||||
name: Registration name for the component
|
||||
category: Optional category for grouping components
|
||||
priority: Priority for ordering (default 0)
|
||||
|
||||
Returns:
|
||||
Decorator function that registers the component class
|
||||
|
||||
Raises:
|
||||
TypeError: If the decorated class doesn't inherit from the base type
|
||||
Validates that the decorated class inherits from the generic
|
||||
type parameter ``T`` declared on the factory.
|
||||
"""
|
||||
|
||||
def decorator(component_cls: Type[T]) -> Type[T]:
|
||||
cls._validate_component(component_cls)
|
||||
cls._registry.register(
|
||||
name, component_cls, category=category, priority=priority
|
||||
)
|
||||
_validate_component(component_cls, cls._component_base)
|
||||
if name in cls._entries:
|
||||
raise ValueError(f"Component '{name}' is already registered")
|
||||
cls._entries[name] = component_cls
|
||||
return component_cls
|
||||
|
||||
return decorator
|
||||
|
||||
@classmethod
|
||||
def create(cls, name: str, *args, **kwargs) -> T:
|
||||
"""Create a component instance by name.
|
||||
|
||||
Args:
|
||||
name: Registered name of the component
|
||||
*args: Positional arguments passed to component constructor
|
||||
**kwargs: Keyword arguments passed to component constructor
|
||||
|
||||
Returns:
|
||||
Component instance
|
||||
|
||||
Raises:
|
||||
ValueError: If the component name is not registered
|
||||
"""Create a component instance by name, filtering kwargs to match
|
||||
the component's ``__init__`` signature.
|
||||
"""
|
||||
if not cls._registry.contains(name):
|
||||
component_cls = cls._entries.get(name)
|
||||
if component_cls is None:
|
||||
raise ValueError(
|
||||
f"Unknown component: '{name}'. "
|
||||
f"Supported types: {sorted(cls._registry.list_names())}"
|
||||
f"Unknown component: '{name}'. Supported types: {sorted(cls._entries)}"
|
||||
)
|
||||
component_cls = cls._registry.get(name)
|
||||
sig = inspect.signature(component_cls.__init__)
|
||||
has_var_kwargs = any(
|
||||
p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()
|
||||
)
|
||||
if not has_var_kwargs:
|
||||
valid = {
|
||||
p.name
|
||||
for p in sig.parameters.values()
|
||||
if p.name != "self" and p.kind != inspect.Parameter.VAR_KEYWORD
|
||||
}
|
||||
kwargs = {k: v for k, v in kwargs.items() if k in valid}
|
||||
return component_cls(*args, **kwargs)
|
||||
|
||||
@classmethod
|
||||
def _validate_component(cls, component_cls: Type[T]) -> None:
|
||||
"""Validate that the component class is valid for this factory.
|
||||
|
||||
Override this method in subclasses to add custom validation.
|
||||
|
||||
Args:
|
||||
component_cls: Component class to validate
|
||||
|
||||
Raises:
|
||||
TypeError: If the component class is invalid
|
||||
"""
|
||||
pass
|
||||
def get_component_class(cls, name: str) -> Type[T]:
|
||||
"""Get the registered component class without instantiating it."""
|
||||
entry = cls._entries.get(name)
|
||||
if entry is None:
|
||||
raise ValueError(
|
||||
f"Unknown component: '{name}'. Supported types: {sorted(cls._entries)}"
|
||||
)
|
||||
return entry
|
||||
|
||||
@classmethod
|
||||
def list_registered(cls) -> list:
|
||||
"""List all registered component names.
|
||||
|
||||
Returns:
|
||||
List of registered component names
|
||||
"""
|
||||
return cls._registry.list_names()
|
||||
def list_registered(cls) -> List[str]:
|
||||
"""List all registered component names."""
|
||||
return sorted(cls._entries)
|
||||
|
||||
@classmethod
|
||||
def is_registered(cls, name: str) -> bool:
|
||||
"""Check if a component name is registered.
|
||||
|
||||
Args:
|
||||
name: Component name to check
|
||||
|
||||
Returns:
|
||||
True if registered, False otherwise
|
||||
"""
|
||||
return cls._registry.contains(name)
|
||||
|
||||
@classmethod
|
||||
def list_by_category(cls, category: str) -> List[str]:
|
||||
"""List registered component names in a category."""
|
||||
return cls._registry.list_by_category(category)
|
||||
|
||||
@classmethod
|
||||
def list_by_priority(cls, reverse: bool = False) -> List[str]:
|
||||
"""List registered component names sorted by priority."""
|
||||
return cls._registry.list_by_priority(reverse)
|
||||
|
||||
|
||||
__all__ = ["Registry", "BaseFactory"]
|
||||
"""Check if a component name is registered."""
|
||||
return name in cls._entries
|
||||
|
||||
@@ -1,46 +1,96 @@
|
||||
"""Inference module for continuous batching.
|
||||
|
||||
Layers:
|
||||
- engine.py: Facade (InferenceEngine), Value Object (GenerationParams, GenerationRequest)
|
||||
- scheduler.py: Continuous-batching loop, Task state machine, TaskStatus enum
|
||||
- cache.py: PagedCache (page-table-indirected KV cache with alloc/free)
|
||||
- sampling.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy)
|
||||
- server.py: FastAPI HTTP server (OpenAI-compatible endpoints)
|
||||
Subpackages:
|
||||
- cache/: KV cache (buffers, strategies, pool)
|
||||
- runtime/: Execution + sampling (executor, CUDA graph, sampling strategies)
|
||||
- task/: Request lifecycle + performance metrics
|
||||
- network/: HTTP protocol handling (server, protocol, OpenAI/Anthropic builders)
|
||||
|
||||
Modules:
|
||||
- scheduler.py: Continuous batching loop
|
||||
- workspace.py: Pre-allocated GPU buffers
|
||||
- engine.py: Facade (InferenceEngine)
|
||||
"""
|
||||
|
||||
from astrai.inference.engine import (
|
||||
GenerationParams,
|
||||
GenerationRequest,
|
||||
InferenceEngine,
|
||||
from astrai.inference.cache import (
|
||||
Allocator,
|
||||
KVCache,
|
||||
KVStorage,
|
||||
PagePool,
|
||||
RadixCache,
|
||||
ReqToTokenPool,
|
||||
TaskCacheManager,
|
||||
page_hash,
|
||||
)
|
||||
from astrai.inference.sampling import (
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
from astrai.inference.network import (
|
||||
AnthropicMessage,
|
||||
BaseToolParser,
|
||||
ChatCompletionRequest,
|
||||
ChatMessage,
|
||||
FunctionDef,
|
||||
GenContext,
|
||||
MessagesRequest,
|
||||
ProtocolHandler,
|
||||
SimpleJsonToolParser,
|
||||
StopChecker,
|
||||
ToolDef,
|
||||
ToolParserFactory,
|
||||
get_app,
|
||||
run_server,
|
||||
)
|
||||
from astrai.inference.network.anthropic import AnthropicResponseBuilder
|
||||
from astrai.inference.network.openai import OpenAIResponseBuilder
|
||||
from astrai.inference.runtime.executor import Executor
|
||||
from astrai.inference.runtime.sample import (
|
||||
BaseSamplingStrategy,
|
||||
FrequencyPenaltyStrategy,
|
||||
SamplingPipeline,
|
||||
TemperatureStrategy,
|
||||
TopKStrategy,
|
||||
TopPStrategy,
|
||||
sample,
|
||||
)
|
||||
from astrai.inference.scheduler import (
|
||||
InferenceScheduler,
|
||||
Task,
|
||||
TaskStatus,
|
||||
)
|
||||
from astrai.inference.scheduler import InferenceScheduler
|
||||
from astrai.inference.task import STOP, Task, TaskManager, TaskStatus
|
||||
|
||||
__all__ = [
|
||||
# Engine / Requests
|
||||
"InferenceEngine",
|
||||
"GenerationRequest",
|
||||
"GenerationParams",
|
||||
# Scheduler
|
||||
"InferenceScheduler",
|
||||
"Executor",
|
||||
"STOP",
|
||||
"Task",
|
||||
"TaskManager",
|
||||
"TaskStatus",
|
||||
# Sampling (Strategy pattern)
|
||||
"Allocator",
|
||||
"KVCache",
|
||||
"KVStorage",
|
||||
"PagePool",
|
||||
"RadixCache",
|
||||
"ReqToTokenPool",
|
||||
"TaskCacheManager",
|
||||
"page_hash",
|
||||
"sample",
|
||||
"BaseSamplingStrategy",
|
||||
"TemperatureStrategy",
|
||||
"TopKStrategy",
|
||||
"TopPStrategy",
|
||||
"FrequencyPenaltyStrategy",
|
||||
"SamplingPipeline",
|
||||
"ProtocolHandler",
|
||||
"StopChecker",
|
||||
"GenContext",
|
||||
"BaseToolParser",
|
||||
"SimpleJsonToolParser",
|
||||
"ToolParserFactory",
|
||||
"OpenAIResponseBuilder",
|
||||
"AnthropicResponseBuilder",
|
||||
"ChatMessage",
|
||||
"ChatCompletionRequest",
|
||||
"FunctionDef",
|
||||
"ToolDef",
|
||||
"AnthropicMessage",
|
||||
"MessagesRequest",
|
||||
"get_app",
|
||||
"run_server",
|
||||
]
|
||||
|
||||
@@ -1,174 +0,0 @@
|
||||
"""Page-based KV cache with page-table-indirected read/write.
|
||||
|
||||
Provides:
|
||||
- PagedCache: paged KV cache combining page pool and tensor storage.
|
||||
"""
|
||||
|
||||
from typing import Dict, List, Tuple
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
STOP = object()
|
||||
|
||||
|
||||
def page_hash(token_ids: List[int], page_idx: int, page_size: int) -> int:
|
||||
start = page_idx * page_size
|
||||
end = min(start + page_size, len(token_ids))
|
||||
h = 0
|
||||
for i in range(start, end):
|
||||
h = (h * 31 + token_ids[i]) & 0xFFFFFFFFFFFFFFFF
|
||||
return h
|
||||
|
||||
|
||||
class PagedCache:
|
||||
"""Paged KV cache with page-table-indirected read/write.
|
||||
|
||||
Combines:
|
||||
- Page pool (ref-counted alloc/free via bitmask)
|
||||
- KV tensor storage (k_cache, v_cache)
|
||||
- Prefix-cache hash lookup (page_content_hash -> physical_page_idx)
|
||||
|
||||
Call :meth:`bind` to obtain a batch view for the attention layers.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
n_layers: int,
|
||||
n_pages: int,
|
||||
page_size: int,
|
||||
n_kv_heads: int,
|
||||
head_dim: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
self.page_size = page_size
|
||||
self._free_mask = (1 << n_pages) - 1
|
||||
self._refs: List[int] = [0] * n_pages
|
||||
self.k_cache = torch.empty(
|
||||
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
self.v_cache = torch.empty(
|
||||
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
self._page_to_hash: Dict[int, int] = {}
|
||||
self._hash_to_page: Dict[int, int] = {}
|
||||
|
||||
def record_page(
|
||||
self, page_idx: int, token_ids: List[int], logical_page_idx: int
|
||||
) -> None:
|
||||
h = page_hash(token_ids, logical_page_idx, self.page_size)
|
||||
old_h = self._page_to_hash.pop(page_idx, None)
|
||||
if old_h is not None:
|
||||
self._hash_to_page.pop(old_h, None)
|
||||
self._page_to_hash[page_idx] = h
|
||||
self._hash_to_page[h] = page_idx
|
||||
|
||||
def lookup_prefix(self, token_ids: List[int]) -> List[int]:
|
||||
full_pages = len(token_ids) // self.page_size
|
||||
hits: List[int] = []
|
||||
for i in range(full_pages):
|
||||
h = page_hash(token_ids, i, self.page_size)
|
||||
p = self._hash_to_page.get(h)
|
||||
if p is None:
|
||||
break
|
||||
hits.append(p)
|
||||
return hits
|
||||
|
||||
def inc_ref(self, idx: int) -> None:
|
||||
self._refs[idx] += 1
|
||||
|
||||
def alloc(self) -> int:
|
||||
lsb = self._free_mask & -self._free_mask
|
||||
if lsb == 0:
|
||||
return -1
|
||||
idx = lsb.bit_length() - 1
|
||||
self._free_mask ^= lsb
|
||||
self._refs[idx] = 1
|
||||
return idx
|
||||
|
||||
def alloc_n(self, n: int) -> List[int]:
|
||||
pages = [self.alloc() for _ in range(n)]
|
||||
if any(p < 0 for p in pages):
|
||||
for p in pages:
|
||||
if p >= 0:
|
||||
self.free(p)
|
||||
return []
|
||||
return pages
|
||||
|
||||
def free(self, idx: int) -> None:
|
||||
self._refs[idx] -= 1
|
||||
if self._refs[idx] == 0:
|
||||
self._free_mask |= 1 << idx
|
||||
h = self._page_to_hash.pop(idx, None)
|
||||
if h is not None:
|
||||
self._hash_to_page.pop(h, None)
|
||||
|
||||
def bind(self, page_table: Tensor, total_len: int = 0) -> "CacheView":
|
||||
return CacheView(self, page_table, total_len)
|
||||
|
||||
def write(
|
||||
self, layer_id: int, page_table: Tensor, start_pos: int, k: Tensor, v: Tensor
|
||||
) -> None:
|
||||
seq_len = k.size(1)
|
||||
if seq_len == 0:
|
||||
return
|
||||
page_size = self.page_size
|
||||
written = 0
|
||||
first_page = start_pos // page_size
|
||||
last_page = (start_pos + seq_len - 1) // page_size
|
||||
for pi in range(first_page, last_page + 1):
|
||||
phys_pages = page_table[:, pi]
|
||||
page_start = pi * page_size
|
||||
write_start = max(page_start, start_pos)
|
||||
write_end = min(page_start + page_size, start_pos + seq_len)
|
||||
offset = write_start - page_start
|
||||
chunk = write_end - write_start
|
||||
self.k_cache[layer_id, phys_pages, offset : offset + chunk] = k[
|
||||
:, written : written + chunk
|
||||
]
|
||||
self.v_cache[layer_id, phys_pages, offset : offset + chunk] = v[
|
||||
:, written : written + chunk
|
||||
]
|
||||
written += chunk
|
||||
|
||||
def gather(self, layer_id: int, page_table: Tensor) -> Tuple[Tensor, Tensor]:
|
||||
k_parts, v_parts = [], []
|
||||
for pi in range(page_table.size(1)):
|
||||
phys_pages = page_table[:, pi]
|
||||
if not (phys_pages >= 0).any():
|
||||
break
|
||||
k_parts.append(self.k_cache[layer_id, phys_pages])
|
||||
v_parts.append(self.v_cache[layer_id, phys_pages])
|
||||
k = torch.cat(k_parts, dim=1)
|
||||
v = torch.cat(v_parts, dim=1)
|
||||
return k, v
|
||||
|
||||
|
||||
class CacheView:
|
||||
"""Per-batch view that bundles PagedCache + page_table + total_len.
|
||||
|
||||
Attention layers receive this as ``paged_cache`` and only see
|
||||
``write()`` / ``gather()``, never raw page tables or length params.
|
||||
"""
|
||||
|
||||
__slots__ = ("_cache", "_page_table", "_total_len")
|
||||
|
||||
def __init__(self, cache: PagedCache, page_table: Tensor, total_len: int = 0):
|
||||
self._cache = cache
|
||||
self._page_table = page_table
|
||||
self._total_len = total_len
|
||||
|
||||
def write(self, layer_id: int, start_pos: int, k: Tensor, v: Tensor) -> None:
|
||||
self._cache.write(layer_id, self._page_table, start_pos, k, v)
|
||||
|
||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
||||
k, v = self._cache.gather(layer_id, self._page_table)
|
||||
if self._total_len:
|
||||
k = k[:, : self._total_len]
|
||||
v = v[:, : self._total_len]
|
||||
return k, v
|
||||
Vendored
+27
@@ -0,0 +1,27 @@
|
||||
"""KV cache subsystem: buffers, strategies, pool management."""
|
||||
|
||||
from astrai.inference.cache.buffer import KVCache, KVStorage, ReqToTokenPool
|
||||
from astrai.inference.cache.pool import PagePool, TaskCacheManager, page_hash
|
||||
from astrai.inference.cache.strategy import (
|
||||
AllocationStrategy,
|
||||
Allocator,
|
||||
ContiguousStrategy,
|
||||
PagedStrategy,
|
||||
RadixCache,
|
||||
TaskCacheState,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"KVCache",
|
||||
"KVStorage",
|
||||
"ReqToTokenPool",
|
||||
"Allocator",
|
||||
"RadixCache",
|
||||
"TaskCacheState",
|
||||
"AllocationStrategy",
|
||||
"ContiguousStrategy",
|
||||
"PagedStrategy",
|
||||
"PagePool",
|
||||
"TaskCacheManager",
|
||||
"page_hash",
|
||||
]
|
||||
Vendored
+106
@@ -0,0 +1,106 @@
|
||||
"""Physical KV cache buffers.
|
||||
|
||||
Layer 1 — ``KVStorage``: flat token-level K/V GPU buffers [n_layers, size, n_kv_heads, head_dim]
|
||||
Layer 2 — ``ReqToTokenPool``: index table [req_idx, pos] → physical token slot
|
||||
Layer 3 — ``KVCache``: pure dataclass passed to the model for direct buffer access
|
||||
|
||||
These classes have no knowledge of tasks, allocation policies, or scheduling.
|
||||
They are the "dumb" physical storage layer.
|
||||
"""
|
||||
|
||||
import threading
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
class ReqToTokenPool:
|
||||
"""Maps [req_idx, pos] → physical token slot in KV storage.
|
||||
|
||||
Each row is one request; each column is a sequence position. The value
|
||||
at [req_idx, pos] is the flat index into the KV storage buffers.
|
||||
"""
|
||||
|
||||
def __init__(self, size: int, max_context_len: int, device: torch.device):
|
||||
self.size = size
|
||||
self.max_context_len = max_context_len
|
||||
self.req_to_token = torch.zeros(
|
||||
(size, max_context_len), dtype=torch.int32, device=device
|
||||
)
|
||||
self.free_slots = list(range(size))
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def alloc(self, num_reqs: int) -> Optional[List[int]]:
|
||||
with self._lock:
|
||||
if num_reqs > len(self.free_slots):
|
||||
return None
|
||||
slots = self.free_slots[:num_reqs]
|
||||
self.free_slots = self.free_slots[num_reqs:]
|
||||
return slots
|
||||
|
||||
def free(self, req_indices: List[int]):
|
||||
with self._lock:
|
||||
self.free_slots.extend(req_indices)
|
||||
|
||||
def write(self, indices, values):
|
||||
self.req_to_token[indices] = values
|
||||
|
||||
|
||||
class KVStorage:
|
||||
"""Token-level KV cache storage.
|
||||
|
||||
Buffers: ``[n_layers, size, n_kv_heads, head_dim]``. Each token occupies
|
||||
one slot indexed by ``ReqToTokenPool``.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size: int,
|
||||
n_layers: int,
|
||||
n_kv_heads: int,
|
||||
head_dim: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
self.size = size
|
||||
self.k_buffer = torch.empty(
|
||||
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
|
||||
)
|
||||
self.v_buffer = torch.empty(
|
||||
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
|
||||
)
|
||||
|
||||
def get_key_buffer(self, layer_id: int) -> Tensor:
|
||||
return self.k_buffer[layer_id]
|
||||
|
||||
def get_value_buffer(self, layer_id: int) -> Tensor:
|
||||
return self.v_buffer[layer_id]
|
||||
|
||||
def set_kv_buffer(self, layer_id: int, loc: Tensor, k: Tensor, v: Tensor) -> None:
|
||||
self.k_buffer[layer_id, loc] = k
|
||||
self.v_buffer[layer_id, loc] = v
|
||||
|
||||
|
||||
@dataclass
|
||||
class KVCache:
|
||||
"""Pure data struct passed to model for KV cache I/O.
|
||||
|
||||
The attention layer does raw buffer indexing — no methods, no abstraction.
|
||||
"""
|
||||
|
||||
k_buffer: Tensor
|
||||
v_buffer: Tensor
|
||||
req_to_token: Tensor
|
||||
req_pool_indices: Tensor
|
||||
seq_lens: Tensor
|
||||
out_cache_loc: Tensor
|
||||
max_len: int = 0
|
||||
kv_indptr: Optional[Tensor] = None
|
||||
qo_indptr: Optional[Tensor] = None
|
||||
q_tile_to_batch: Optional[Tensor] = None
|
||||
q_tile_to_index: Optional[Tensor] = None
|
||||
decode_o_part: Optional[Tensor] = None
|
||||
decode_ml_part: Optional[Tensor] = None
|
||||
decode_out: Optional[Tensor] = None
|
||||
Vendored
+382
@@ -0,0 +1,382 @@
|
||||
"""KV cache orchestration: PagePool + TaskCacheManager.
|
||||
|
||||
PagePool owns the physical buffers (``KVStorage`` + ``ReqToTokenPool``)
|
||||
and wires them to an allocation strategy. It assembles the ``KVCache``
|
||||
dataclass passed to the model forward.
|
||||
|
||||
TaskCacheManager owns the ``task_id`` → ``TaskCacheState`` mapping and
|
||||
delegates physical slot allocation to the strategy, and KV bind to the pool.
|
||||
|
||||
See ``cache_buffer.py`` for the raw buffer primitives and ``cache_strategy.py``
|
||||
for the allocation policies.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.inference.cache.buffer import KVCache, KVStorage, ReqToTokenPool
|
||||
from astrai.inference.cache.strategy import (
|
||||
AllocationStrategy,
|
||||
Allocator,
|
||||
ContiguousStrategy,
|
||||
PagedStrategy,
|
||||
RadixCache,
|
||||
TaskCacheState,
|
||||
)
|
||||
from astrai.inference.workspace import Q_TILE_ROWS, InferenceWorkspace
|
||||
|
||||
# Re-export everything so existing ``from astrai.inference.cache import ...``
|
||||
# continues to work unchanged after the file split.
|
||||
__all__ = [
|
||||
"KVCache",
|
||||
"KVStorage",
|
||||
"ReqToTokenPool",
|
||||
"Allocator",
|
||||
"RadixCache",
|
||||
"AllocationStrategy",
|
||||
"ContiguousStrategy",
|
||||
"PagedStrategy",
|
||||
"PagePool",
|
||||
"TaskCacheManager",
|
||||
"TaskCacheState",
|
||||
"page_hash",
|
||||
]
|
||||
|
||||
# ---- helpers ----
|
||||
|
||||
|
||||
def page_hash(
|
||||
token_ids: List[int], page_idx: int, page_size: int, parent_hash: int = 0
|
||||
) -> int:
|
||||
start = page_idx * page_size
|
||||
end = min(start + page_size, len(token_ids))
|
||||
h = parent_hash
|
||||
for i in range(start, end):
|
||||
h = (h * 31 + token_ids[i]) & 0xFFFFFFFFFFFFFFFF
|
||||
return h
|
||||
|
||||
|
||||
def _is_steady_increment(
|
||||
prev_sig: Optional[tuple],
|
||||
prev_vals: Optional[List[int]],
|
||||
cur_sig: tuple,
|
||||
cur_vals: List[int],
|
||||
) -> bool:
|
||||
return (
|
||||
prev_sig is not None
|
||||
and prev_vals is not None
|
||||
and prev_sig == cur_sig
|
||||
and len(prev_vals) == len(cur_vals)
|
||||
and all(c == p + 1 for c, p in zip(cur_vals, prev_vals))
|
||||
)
|
||||
|
||||
|
||||
# ---- task-scoped bind state ----
|
||||
@dataclass
|
||||
class _BindState:
|
||||
"""Cached bind metadata for steady-state decode increment detection."""
|
||||
|
||||
sig: tuple
|
||||
seq_lens: List[int]
|
||||
|
||||
|
||||
# ---- pool + manager ----
|
||||
|
||||
|
||||
class PagePool:
|
||||
"""Physical KV cache: buffers + req-to-token table + allocation strategy + bind.
|
||||
|
||||
Does not know about tasks — task lifecycle is managed by
|
||||
:class:`TaskCacheManager`, which holds a reference to this pool.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
n_layers: int,
|
||||
n_kv_heads: int,
|
||||
head_dim: int,
|
||||
max_batch_size: int,
|
||||
max_seq_len: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
page_size: int = 1,
|
||||
n_tokens: Optional[int] = None,
|
||||
):
|
||||
self.page_size = page_size
|
||||
self.max_batch_size = max_batch_size
|
||||
self.max_seq_len = max_seq_len
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
self.n_layers = n_layers
|
||||
self.n_kv_heads = n_kv_heads
|
||||
self.head_dim = head_dim
|
||||
|
||||
self.contiguous = n_tokens is None
|
||||
self.n_tokens = max_batch_size * max_seq_len if self.contiguous else n_tokens
|
||||
if self.n_tokens > torch.iinfo(torch.int32).max:
|
||||
raise ValueError("KV cache token count exceeds the int32 slot index limit")
|
||||
|
||||
self._storage = KVStorage(
|
||||
self.n_tokens, n_layers, n_kv_heads, head_dim, device, dtype
|
||||
)
|
||||
self._req_pool = ReqToTokenPool(max_batch_size, max_seq_len, device)
|
||||
|
||||
if self.contiguous:
|
||||
for i in range(max_batch_size):
|
||||
self._req_pool.req_to_token[i] = torch.arange(
|
||||
i * max_seq_len,
|
||||
(i + 1) * max_seq_len,
|
||||
dtype=torch.int32,
|
||||
device=device,
|
||||
)
|
||||
self._strategy: AllocationStrategy = ContiguousStrategy()
|
||||
else:
|
||||
n_pages = self.n_tokens // page_size
|
||||
alloc = Allocator(n_pages)
|
||||
prefix = RadixCache(page_size) if page_size > 1 else None
|
||||
if prefix is not None:
|
||||
alloc.on_evict = prefix.evict
|
||||
self._strategy = PagedStrategy(
|
||||
alloc, prefix, page_size, self._req_pool, device
|
||||
)
|
||||
|
||||
@property
|
||||
def strategy(self) -> AllocationStrategy:
|
||||
return self._strategy
|
||||
|
||||
@property
|
||||
def req_pool(self) -> ReqToTokenPool:
|
||||
return self._req_pool
|
||||
|
||||
def bind_tasks(
|
||||
self,
|
||||
req_indices: List[int],
|
||||
seq_lens: List[int],
|
||||
workspace: InferenceWorkspace,
|
||||
device: Optional[torch.device] = None,
|
||||
start_pos: Optional[int] = None,
|
||||
incremental: bool = False,
|
||||
) -> KVCache:
|
||||
"""Assemble the ``KVCache`` metadata for a batch of tasks.
|
||||
|
||||
Args:
|
||||
req_indices: request slot indices (from ``ReqToTokenPool``).
|
||||
seq_lens: current sequence length per task.
|
||||
workspace: pre-allocated fixed-shape buffers (CUDA-graph safe).
|
||||
start_pos: if set, produce **prefill** cache (full q_len range).
|
||||
If ``None``, produce **decode** cache (last position).
|
||||
incremental: if ``True``, reuse workspace state from previous step
|
||||
by incrementing counters in-place (decode hot path).
|
||||
|
||||
Returns:
|
||||
``KVCache`` dataclass with the correct output shapes for the
|
||||
attention backend (prefill: ``[B, q_len]``, decode: ``[B, 1]``).
|
||||
"""
|
||||
if device is None:
|
||||
device = workspace.device
|
||||
b = len(req_indices)
|
||||
|
||||
rpi_buf = workspace.req_pool_indices
|
||||
sl_buf = workspace.seq_lens
|
||||
kvp_buf = workspace.kv_indptr
|
||||
inc_buf = workspace.inc
|
||||
ocl_buf = workspace.out_cache_loc
|
||||
|
||||
if incremental:
|
||||
sl_buf[:b] += 1
|
||||
kvp_buf[: b + 1] += inc_buf[: b + 1]
|
||||
else:
|
||||
rpi_buf[:b].copy_(
|
||||
torch.tensor(req_indices, dtype=torch.int32, device=device)
|
||||
)
|
||||
sl_buf[:b].copy_(torch.tensor(seq_lens, dtype=torch.long, device=device))
|
||||
kvp_buf[: b + 1].zero_()
|
||||
kvp_buf[1 : b + 1] = sl_buf[:b].cumsum(0).to(torch.int32)
|
||||
|
||||
req_pool_indices = rpi_buf[:b]
|
||||
seq_lens_t = sl_buf[:b]
|
||||
kv_indptr = kvp_buf[: b + 1]
|
||||
|
||||
if start_pos is not None:
|
||||
# Packed prefill concatenates each request's query tokens.
|
||||
q_lens = [seq_len - start_pos for seq_len in seq_lens]
|
||||
if any(q_len <= 0 for q_len in q_lens):
|
||||
raise ValueError("prefill sequence lengths must exceed start_pos")
|
||||
out_cache_loc = torch.cat(
|
||||
[
|
||||
self._req_pool.req_to_token[
|
||||
req_pool_indices[i], start_pos : seq_lens[i]
|
||||
]
|
||||
for i in range(b)
|
||||
]
|
||||
)
|
||||
workspace.qo_indptr[: b + 1].zero_()
|
||||
workspace.qo_indptr[1 : b + 1].copy_(
|
||||
torch.tensor(q_lens, dtype=torch.int32, device=device).cumsum(0)
|
||||
)
|
||||
qo_indptr = workspace.qo_indptr[: b + 1]
|
||||
tile_batches = []
|
||||
tile_indices = []
|
||||
for batch, q_len in enumerate(q_lens):
|
||||
n_tiles = (q_len + Q_TILE_ROWS - 1) // Q_TILE_ROWS
|
||||
tile_batches.extend([batch] * n_tiles)
|
||||
tile_indices.extend(range(n_tiles))
|
||||
n_tiles = len(tile_batches)
|
||||
workspace.q_tile_to_batch[:n_tiles].copy_(
|
||||
torch.tensor(tile_batches, dtype=torch.int32, device=device)
|
||||
)
|
||||
workspace.q_tile_to_index[:n_tiles].copy_(
|
||||
torch.tensor(tile_indices, dtype=torch.int32, device=device)
|
||||
)
|
||||
q_tile_to_batch = workspace.q_tile_to_batch[:n_tiles]
|
||||
q_tile_to_index = workspace.q_tile_to_index[:n_tiles]
|
||||
decode_o_part = decode_ml_part = decode_out = None
|
||||
else:
|
||||
# ---- decode: out_cache_loc is a single column (last position) ----
|
||||
write_pos = seq_lens_t - 1
|
||||
loc = self._req_pool.req_to_token[req_pool_indices, write_pos].unsqueeze(-1)
|
||||
ocl_buf[:b].copy_(loc)
|
||||
out_cache_loc = ocl_buf[:b].reshape(-1)
|
||||
workspace.qo_indptr[: b + 1].copy_(inc_buf[: b + 1])
|
||||
qo_indptr = workspace.qo_indptr[: b + 1]
|
||||
q_tile_to_batch = q_tile_to_index = None
|
||||
decode_o_part = getattr(workspace, "decode_o_part", None)
|
||||
decode_ml_part = getattr(workspace, "decode_ml_part", None)
|
||||
decode_out = getattr(workspace, "decode_out", None)
|
||||
|
||||
return KVCache(
|
||||
k_buffer=self._storage.k_buffer,
|
||||
v_buffer=self._storage.v_buffer,
|
||||
req_to_token=self._req_pool.req_to_token,
|
||||
req_pool_indices=req_pool_indices,
|
||||
seq_lens=seq_lens_t,
|
||||
out_cache_loc=out_cache_loc,
|
||||
max_len=max(seq_lens),
|
||||
kv_indptr=kv_indptr,
|
||||
qo_indptr=qo_indptr,
|
||||
q_tile_to_batch=q_tile_to_batch,
|
||||
q_tile_to_index=q_tile_to_index,
|
||||
decode_o_part=decode_o_part,
|
||||
decode_ml_part=decode_ml_part,
|
||||
decode_out=decode_out,
|
||||
)
|
||||
|
||||
|
||||
class TaskCacheManager:
|
||||
"""Task ↔ KV slot lifecycle manager.
|
||||
|
||||
Sole owner of ``task_id → TaskCacheState``. Delegates physical slot
|
||||
allocation to the strategy (via ``pool.strategy``) and KV bind to
|
||||
``pool.bind_tasks()``.
|
||||
|
||||
Usage::
|
||||
|
||||
pool = PagePool(...)
|
||||
mgr = TaskCacheManager(pool)
|
||||
mgr.task_alloc("req_1", [101, 202, 303])
|
||||
...
|
||||
kv = mgr.bind(["req_1"], workspace)
|
||||
"""
|
||||
|
||||
def __init__(self, pool: PagePool):
|
||||
self._pool = pool
|
||||
self._strategy = pool.strategy
|
||||
self._req_pool = pool.req_pool
|
||||
self._max_seq_len = pool.max_seq_len
|
||||
self._states: Dict[str, TaskCacheState] = {}
|
||||
self._bind_state: Optional[_BindState] = None
|
||||
self._bind_was_steady = False
|
||||
|
||||
# -- public task lifecycle --
|
||||
|
||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
||||
self._bind_state = None
|
||||
req_slots = self._req_pool.alloc(1)
|
||||
if req_slots is None:
|
||||
return False
|
||||
state = TaskCacheState(req_idx=req_slots[0])
|
||||
self._states[task_id] = state
|
||||
if not self._strategy.alloc(state, prompt_ids):
|
||||
self._rollback(state, task_id)
|
||||
return False
|
||||
self._strategy.write_indices(state, prompt_ids)
|
||||
state.length = len(prompt_ids)
|
||||
return True
|
||||
|
||||
def task_free(self, task_id: str):
|
||||
self._bind_state = None
|
||||
state = self._states.pop(task_id, None)
|
||||
if state is None:
|
||||
return
|
||||
self._strategy.free(state)
|
||||
self._req_pool.free([state.req_idx])
|
||||
|
||||
def task_extend(self, task_id: str, pos: int) -> bool:
|
||||
state = self._states.get(task_id)
|
||||
if state is None or pos >= self._max_seq_len:
|
||||
return False
|
||||
if not self._strategy.extend(state, pos):
|
||||
return False
|
||||
state.length = pos + 1
|
||||
return True
|
||||
|
||||
def task_cached(self, task_id: str) -> int:
|
||||
state = self._states.get(task_id)
|
||||
return state.cached if state is not None else 0
|
||||
|
||||
def task_record_hashes(
|
||||
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
||||
):
|
||||
state = self._states.get(task_id)
|
||||
if state is not None:
|
||||
self._strategy.record_hashes(state, prompt_ids, start_logical_page)
|
||||
|
||||
@staticmethod
|
||||
def task_cacheable_ids(task_id: str, prompt_ids: List[int], output_ids: List[int]):
|
||||
return list(prompt_ids) + list(output_ids[:-1])
|
||||
|
||||
# -- bind (assemble KVCache for the model forward) --
|
||||
|
||||
def bind(
|
||||
self,
|
||||
task_ids: List[str],
|
||||
workspace: InferenceWorkspace,
|
||||
device: Optional[torch.device] = None,
|
||||
start_pos: Optional[int] = None,
|
||||
) -> KVCache:
|
||||
"""Build ``KVCache`` for an ordered list of task IDs."""
|
||||
states = [self._states[tid] for tid in task_ids]
|
||||
req_indices = [s.req_idx for s in states]
|
||||
seq_lens = [s.length for s in states]
|
||||
sig = tuple(req_indices)
|
||||
|
||||
prev = self._bind_state
|
||||
incremental = (
|
||||
start_pos is None
|
||||
and prev is not None
|
||||
and _is_steady_increment(prev.sig, prev.seq_lens, sig, seq_lens)
|
||||
)
|
||||
self._bind_state = _BindState(sig, list(seq_lens))
|
||||
self._bind_was_steady = incremental
|
||||
|
||||
return self._pool.bind_tasks(
|
||||
req_indices,
|
||||
seq_lens,
|
||||
workspace,
|
||||
device=device,
|
||||
start_pos=start_pos,
|
||||
incremental=incremental,
|
||||
)
|
||||
|
||||
@property
|
||||
def bind_was_steady(self) -> bool:
|
||||
return self._bind_was_steady
|
||||
|
||||
# -- internals --
|
||||
|
||||
def _rollback(self, state: TaskCacheState, task_id: str):
|
||||
self._strategy.free(state)
|
||||
self._req_pool.free([state.req_idx])
|
||||
self._states.pop(task_id, None)
|
||||
Vendored
+320
@@ -0,0 +1,320 @@
|
||||
"""KV cache allocation layer.
|
||||
|
||||
Encapsulates the physical slot allocation policy, isolated from GPU buffers
|
||||
and task lifecycle management.
|
||||
|
||||
- ``TaskCacheState``: data contract between strategy and manager (per-task slot state)
|
||||
- ``Allocator``: bitmask-based page allocator with LRU eviction
|
||||
- ``RadixCache``: page-granular prefix index (exact token match)
|
||||
- ``AllocationStrategy``: ABC for physical slot allocation
|
||||
- ``ContiguousStrategy``: statically partitioned, no dynamic allocation
|
||||
- ``PagedStrategy``: dynamic paged allocation from a shared pool
|
||||
"""
|
||||
|
||||
import threading
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, Dict, List, Optional, OrderedDict
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.inference.cache.buffer import ReqToTokenPool
|
||||
|
||||
# ---- data contract: per-task slot state ----
|
||||
|
||||
|
||||
@dataclass
|
||||
class TaskCacheState:
|
||||
"""Per-task cache allocation state.
|
||||
|
||||
Co-locates all task-owned cache metadata so the alloc/free/extend
|
||||
lifecycle is atomic. Owned by ``TaskCacheManager``, consumed by
|
||||
every ``AllocationStrategy`` method.
|
||||
"""
|
||||
|
||||
req_idx: int
|
||||
length: int = 0
|
||||
cached: int = 0
|
||||
pages: List[int] = field(default_factory=list)
|
||||
|
||||
|
||||
# ---- allocation primitives ----
|
||||
|
||||
|
||||
class Allocator:
|
||||
"""Bitmask-based page allocator with ref-counting and LRU eviction."""
|
||||
|
||||
def __init__(self, n_pages: int):
|
||||
self._free_mask = (1 << n_pages) - 1
|
||||
self._refs: List[int] = [0] * n_pages
|
||||
self._lru: OrderedDict[int, None] = OrderedDict()
|
||||
self.on_evict: Optional[Callable[[int], None]] = None
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def alloc(self) -> int:
|
||||
with self._lock:
|
||||
if self._free_mask:
|
||||
lsb = self._free_mask & -self._free_mask
|
||||
idx = lsb.bit_length() - 1
|
||||
self._free_mask ^= lsb
|
||||
self._refs[idx] = 1
|
||||
return idx
|
||||
if self._lru:
|
||||
idx, _ = self._lru.popitem(last=False)
|
||||
if self.on_evict:
|
||||
self.on_evict(idx)
|
||||
self._refs[idx] = 1
|
||||
self._free_mask &= ~(1 << idx)
|
||||
return idx
|
||||
return -1
|
||||
|
||||
def free(self, idx: int, keep_cached: bool = False):
|
||||
with self._lock:
|
||||
self._refs[idx] -= 1
|
||||
if self._refs[idx] == 0:
|
||||
if keep_cached:
|
||||
self._lru[idx] = None
|
||||
else:
|
||||
self._free_mask |= 1 << idx
|
||||
|
||||
def inc_ref(self, idx: int):
|
||||
with self._lock:
|
||||
self._refs[idx] += 1
|
||||
self._lru.pop(idx, None)
|
||||
|
||||
def ref_count(self, idx: int) -> int:
|
||||
with self._lock:
|
||||
return self._refs[idx]
|
||||
|
||||
def touch(self, idx: int):
|
||||
with self._lock:
|
||||
if idx in self._lru:
|
||||
self._lru.move_to_end(idx)
|
||||
|
||||
|
||||
class RadixNode:
|
||||
"""A page-aligned edge in the CPU-side prefix radix trie."""
|
||||
|
||||
__slots__ = ("parent", "children", "page_idx", "tokens", "lock_ref")
|
||||
|
||||
def __init__(self, parent=None, tokens=(), page_idx=None):
|
||||
self.parent = parent
|
||||
self.children: Dict[tuple, "RadixNode"] = {}
|
||||
self.page_idx = page_idx
|
||||
self.tokens = tuple(tokens)
|
||||
self.lock_ref = 0
|
||||
|
||||
|
||||
class RadixCache:
|
||||
"""Page-granular radix prefix index with exact token matching."""
|
||||
|
||||
def __init__(self, page_size: int):
|
||||
self._page_size = page_size
|
||||
self._root = RadixNode()
|
||||
self._page_to_node: Dict[int, RadixNode] = {}
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def evict(self, idx: int):
|
||||
with self._lock:
|
||||
node = self._page_to_node.pop(idx, None)
|
||||
if node is None:
|
||||
return
|
||||
node.page_idx = None
|
||||
parent = node.parent
|
||||
if parent is not None:
|
||||
parent.children.pop(node.tokens, None)
|
||||
|
||||
def has_page(self, idx: int) -> bool:
|
||||
with self._lock:
|
||||
return idx in self._page_to_node
|
||||
|
||||
def lookup(self, token_ids: List[int]) -> List[int]:
|
||||
with self._lock:
|
||||
full_pages = len(token_ids) // self._page_size
|
||||
hits: List[int] = []
|
||||
node = self._root
|
||||
for i in range(full_pages):
|
||||
start = i * self._page_size
|
||||
page_tokens = tuple(token_ids[start : start + self._page_size])
|
||||
child = node.children.get(page_tokens)
|
||||
if child is None or child.page_idx is None:
|
||||
break
|
||||
hits.append(child.page_idx)
|
||||
node = child
|
||||
return hits
|
||||
|
||||
def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int):
|
||||
with self._lock:
|
||||
full_pages = len(token_ids) // self._page_size
|
||||
if logical_page_idx >= full_pages:
|
||||
return
|
||||
old = self._page_to_node.pop(page_idx, None)
|
||||
if old is not None and old.parent is not None:
|
||||
old.parent.children.pop(old.tokens, None)
|
||||
|
||||
node = self._root
|
||||
for i in range(logical_page_idx + 1):
|
||||
start = i * self._page_size
|
||||
page_tokens = tuple(token_ids[start : start + self._page_size])
|
||||
child = node.children.get(page_tokens)
|
||||
if child is None:
|
||||
child = RadixNode(node, page_tokens)
|
||||
node.children[page_tokens] = child
|
||||
node = child
|
||||
if node.page_idx is not None and node.page_idx != page_idx:
|
||||
replaced = node.page_idx
|
||||
self._page_to_node.pop(replaced, None)
|
||||
node.page_idx = page_idx
|
||||
self._page_to_node[page_idx] = node
|
||||
|
||||
def release(self, pages: List[int]) -> None:
|
||||
with self._lock:
|
||||
for page_idx in pages:
|
||||
node = self._page_to_node.get(page_idx)
|
||||
if node is not None and node.lock_ref:
|
||||
node.lock_ref -= 1
|
||||
|
||||
|
||||
class AllocationStrategy(ABC):
|
||||
"""Physical slot allocation policy.
|
||||
|
||||
Subclasses implement the actual allocation semantics. This ABC declares
|
||||
the contract; there are no default implementations.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def alloc(self, state: TaskCacheState, prompt_ids: List[int]) -> bool: ...
|
||||
|
||||
@abstractmethod
|
||||
def free(self, state: TaskCacheState) -> None: ...
|
||||
|
||||
@abstractmethod
|
||||
def extend(self, state: TaskCacheState, pos: int) -> bool: ...
|
||||
|
||||
@abstractmethod
|
||||
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None: ...
|
||||
|
||||
@abstractmethod
|
||||
def record_hashes(
|
||||
self,
|
||||
state: TaskCacheState,
|
||||
prompt_ids: List[int],
|
||||
start: int,
|
||||
) -> None: ...
|
||||
|
||||
|
||||
class ContiguousStrategy(AllocationStrategy):
|
||||
"""Static contiguous allocation: slots are pre-assigned at pool init.
|
||||
|
||||
No dynamic allocation or prefix caching. All operations are no-ops
|
||||
because ``ReqToTokenPool`` is pre-filled with contiguous ranges.
|
||||
"""
|
||||
|
||||
def alloc(self, state: TaskCacheState, prompt_ids: List[int]) -> bool:
|
||||
return True
|
||||
|
||||
def free(self, state: TaskCacheState) -> None:
|
||||
pass
|
||||
|
||||
def extend(self, state: TaskCacheState, pos: int) -> bool:
|
||||
return True
|
||||
|
||||
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None:
|
||||
pass
|
||||
|
||||
def record_hashes(
|
||||
self,
|
||||
state: TaskCacheState,
|
||||
prompt_ids: List[int],
|
||||
start: int,
|
||||
) -> None:
|
||||
pass
|
||||
|
||||
|
||||
class PagedStrategy(AllocationStrategy):
|
||||
"""Dynamic paged allocation from a shared bitmask pool.
|
||||
|
||||
``page_size`` is a parameter, not a separate strategy: at ``page_size=1``
|
||||
each allocated page *is* one token slot (``page * 1 + 0``), and prefix
|
||||
caching is simply disabled (``prefix=None``). The unified page formula
|
||||
``pages[page_idx] * page_size + offset`` holds for both.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
alloc: Allocator,
|
||||
prefix: Optional[RadixCache],
|
||||
page_size: int,
|
||||
req_pool: ReqToTokenPool,
|
||||
device,
|
||||
):
|
||||
self._alloc = alloc
|
||||
self._prefix = prefix
|
||||
self._page_size = page_size
|
||||
self._req_pool = req_pool
|
||||
self._device = device
|
||||
|
||||
def alloc(self, state: TaskCacheState, prompt_ids: List[int]) -> bool:
|
||||
if self._prefix is not None:
|
||||
hits = self._prefix.lookup(prompt_ids)
|
||||
state.cached = len(hits) * self._page_size
|
||||
for p in hits:
|
||||
self._alloc.inc_ref(p)
|
||||
state.pages = list(hits)
|
||||
|
||||
remaining = len(prompt_ids) - state.cached
|
||||
if remaining <= 0:
|
||||
return True
|
||||
n_new = (remaining + self._page_size - 1) // self._page_size
|
||||
for _ in range(n_new):
|
||||
p = self._alloc.alloc()
|
||||
if p < 0:
|
||||
return False
|
||||
state.pages.append(p)
|
||||
return True
|
||||
|
||||
def free(self, state: TaskCacheState) -> None:
|
||||
if self._prefix is not None:
|
||||
for p in state.pages:
|
||||
keep = self._prefix.has_page(p)
|
||||
self._alloc.free(p, keep_cached=keep)
|
||||
if not keep:
|
||||
self._prefix.evict(p)
|
||||
else:
|
||||
for p in state.pages:
|
||||
self._alloc.free(p)
|
||||
|
||||
def extend(self, state: TaskCacheState, pos: int) -> bool:
|
||||
page_idx = pos // self._page_size
|
||||
if page_idx >= len(state.pages):
|
||||
p = self._alloc.alloc()
|
||||
if p < 0:
|
||||
return False
|
||||
state.pages.append(p)
|
||||
offset = pos % self._page_size
|
||||
self._req_pool.req_to_token[state.req_idx, pos] = (
|
||||
state.pages[page_idx] * self._page_size + offset
|
||||
)
|
||||
return True
|
||||
|
||||
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None:
|
||||
total = len(prompt_ids)
|
||||
for pos in range(total):
|
||||
page_idx = pos // self._page_size
|
||||
offset = pos % self._page_size
|
||||
if page_idx < len(state.pages):
|
||||
self._req_pool.req_to_token[state.req_idx, pos] = (
|
||||
state.pages[page_idx] * self._page_size + offset
|
||||
)
|
||||
|
||||
def record_hashes(
|
||||
self,
|
||||
state: TaskCacheState,
|
||||
prompt_ids: List[int],
|
||||
start: int,
|
||||
) -> None:
|
||||
if self._prefix is None:
|
||||
return
|
||||
full = len(prompt_ids) // self._page_size
|
||||
for i in range(start, min(full, len(state.pages))):
|
||||
self._prefix.record(state.pages[i], prompt_ids, i)
|
||||
+116
-333
@@ -1,132 +1,35 @@
|
||||
"""Unified inference engine for continuous batching.
|
||||
|
||||
Layers:
|
||||
- GenerationParams: Immutable value object for sampling parameters.
|
||||
- GenerationRequest: User-facing request DTO with validation.
|
||||
- _Result: Thread-safe token accumulator (Observer pattern).
|
||||
- InferenceEngine: Facade over InferenceScheduler + async wrapper.
|
||||
"""
|
||||
"""Unified inference engine for continuous batching."""
|
||||
|
||||
import asyncio
|
||||
import gc
|
||||
import threading
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Union
|
||||
from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from astrai.inference.cache import STOP
|
||||
from astrai.extension import ATTN_BACKEND, AttentionBackend, get_backend
|
||||
from astrai.inference.cache import PagePool
|
||||
from astrai.inference.scheduler import InferenceScheduler
|
||||
from astrai.inference.task import STOP
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GenerationParams:
|
||||
"""Immutable value object for sampling hyperparameters."""
|
||||
|
||||
top_k: int = 50
|
||||
top_p: float = 1.0
|
||||
temperature: float = 1.0
|
||||
max_tokens: int = 1024
|
||||
|
||||
|
||||
class GenerationRequest:
|
||||
"""Request parameters for text generation.
|
||||
|
||||
Encapsulates messages, sampling parameters (via GenerationParams),
|
||||
and streaming preference for a single generation request.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
messages: List[Dict[str, str]],
|
||||
top_k: int = 50,
|
||||
top_p: float = 1.0,
|
||||
temperature: float = 1.0,
|
||||
max_len: int = 1024,
|
||||
stream: bool = False,
|
||||
):
|
||||
"""Initializes a generation request.
|
||||
|
||||
Args:
|
||||
messages: Conversation history as list of {"role": ..., "content": ...}.
|
||||
top_k: Top-k sampling count (0 disables).
|
||||
top_p: Nucleus sampling probability threshold.
|
||||
temperature: Sampling temperature.
|
||||
max_len: Maximum tokens to generate.
|
||||
stream: Whether to return output as a token stream.
|
||||
"""
|
||||
self.messages = messages
|
||||
self.params = GenerationParams(
|
||||
top_k=top_k,
|
||||
top_p=top_p,
|
||||
temperature=temperature,
|
||||
max_tokens=max_len,
|
||||
)
|
||||
self.stream = stream
|
||||
self._validate()
|
||||
|
||||
@property
|
||||
def top_k(self) -> int:
|
||||
return self.params.top_k
|
||||
|
||||
@property
|
||||
def top_p(self) -> float:
|
||||
return self.params.top_p
|
||||
|
||||
@property
|
||||
def temperature(self) -> float:
|
||||
return self.params.temperature
|
||||
|
||||
@property
|
||||
def max_len(self) -> int:
|
||||
return self.params.max_tokens
|
||||
|
||||
def _validate(self):
|
||||
"""Validates sampling parameter ranges."""
|
||||
if not (isinstance(self.top_k, int) and self.top_k >= 0):
|
||||
raise ValueError("top_k must be a non-negative integer")
|
||||
if not (0.0 <= self.top_p <= 1.0):
|
||||
raise ValueError("top_p must be a float between 0.0 and 1.0")
|
||||
if not (isinstance(self.temperature, (int, float)) and self.temperature >= 0):
|
||||
raise ValueError("temperature must be a non-negative number")
|
||||
|
||||
|
||||
class _Result:
|
||||
"""Thread-safe token accumulator for streaming and non-streaming modes.
|
||||
|
||||
Supports multiple concurrent generation tasks with per-index result tracking.
|
||||
Uses a threading.Condition for efficient completion notification
|
||||
and a threading.Event for streaming wakeup.
|
||||
"""
|
||||
class GenerateResult:
|
||||
"""Thread-safe token accumulator for streaming and non-streaming modes."""
|
||||
|
||||
def __init__(self, count: int = 1):
|
||||
"""Initializes the accumulator.
|
||||
|
||||
Args:
|
||||
count: Number of concurrent generation tasks to track.
|
||||
"""
|
||||
self._cond = threading.Condition()
|
||||
self._event = threading.Event()
|
||||
self.tokens: List[str] = []
|
||||
self.tokens: List[Tuple[int, str]] = []
|
||||
self.results: List[str] = [""] * count
|
||||
self._done: List[bool] = [False] * count
|
||||
self._completed = 0
|
||||
self._total = count
|
||||
|
||||
def append(self, token: str, idx: int = 0):
|
||||
"""Appends a token to the result buffer.
|
||||
|
||||
In non-streaming mode, tokens are concatenated into results[idx].
|
||||
The sentinel STOP marks a task as complete.
|
||||
|
||||
Args:
|
||||
token: The decoded token string, or STOP sentinel.
|
||||
idx: Index of the generation task this token belongs to.
|
||||
"""
|
||||
with self._cond:
|
||||
self.tokens.append(token)
|
||||
self.tokens.append((idx, token))
|
||||
if token is not STOP:
|
||||
self.results[idx] += token
|
||||
else:
|
||||
@@ -136,12 +39,7 @@ class _Result:
|
||||
self._cond.notify_all()
|
||||
self._event.set()
|
||||
|
||||
def pop_all(self) -> List[str]:
|
||||
"""Returns and clears all accumulated tokens.
|
||||
|
||||
Returns:
|
||||
List of token strings since the last call.
|
||||
"""
|
||||
def pop_all(self) -> List[Tuple[int, str]]:
|
||||
with self._cond:
|
||||
out = self.tokens.copy()
|
||||
self.tokens.clear()
|
||||
@@ -150,45 +48,25 @@ class _Result:
|
||||
return out
|
||||
|
||||
def wait(self, timeout: Optional[float] = None) -> bool:
|
||||
"""Blocks until new tokens arrive or the timeout expires.
|
||||
|
||||
Args:
|
||||
timeout: Maximum wait time in seconds (None = infinite).
|
||||
|
||||
Returns:
|
||||
True if the event was set (new data available), False on timeout.
|
||||
"""
|
||||
return self._event.wait(timeout=timeout)
|
||||
|
||||
def wait_completion(self) -> None:
|
||||
"""Blocks until all tasks complete (non-streaming).
|
||||
|
||||
Uses a Condition to sleep efficiently instead of busy-waiting.
|
||||
The calling thread is parked until a STOP signal arrives.
|
||||
"""
|
||||
def wait_completion(self, timeout: float = 300.0):
|
||||
with self._cond:
|
||||
self._cond.wait_for(lambda: self._completed >= self._total)
|
||||
if not self._cond.wait_for(
|
||||
lambda: self._completed >= self._total, timeout=timeout
|
||||
):
|
||||
raise TimeoutError(
|
||||
f"Generation timeout after {timeout}s "
|
||||
f"({self._completed}/{self._total} completed)"
|
||||
)
|
||||
|
||||
def get_results(self) -> List[str]:
|
||||
"""Returns all accumulated results for non-streaming mode.
|
||||
|
||||
Returns:
|
||||
List of complete generated strings, one per task index.
|
||||
"""
|
||||
with self._cond:
|
||||
return self.results.copy()
|
||||
|
||||
|
||||
class InferenceEngine:
|
||||
"""Unified inference engine backed by continuous-batching scheduler.
|
||||
|
||||
Usage:
|
||||
with InferenceEngine(model, tokenizer) as engine:
|
||||
for token in engine.generate("hello", stream=True):
|
||||
print(token, end="")
|
||||
|
||||
text = engine.generate("hello")
|
||||
"""
|
||||
"""Unified inference engine backed by continuous-batching scheduler."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -196,20 +74,10 @@ class InferenceEngine:
|
||||
tokenizer: AutoTokenizer,
|
||||
max_batch_size: int = 1,
|
||||
max_seq_len: Optional[int] = None,
|
||||
max_prompt_len: int = 2048,
|
||||
page_size: int = 128,
|
||||
cache: Optional[PagePool] = None,
|
||||
enable_cuda_graph: bool = True,
|
||||
backend: Optional[Union[str, ATTN_BACKEND, AttentionBackend, type]] = None,
|
||||
):
|
||||
"""Initializes the inference engine.
|
||||
|
||||
Args:
|
||||
model: The model instance.
|
||||
tokenizer: The tokenizer instance.
|
||||
max_batch_size: Maximum number of concurrent tasks.
|
||||
max_seq_len: Maximum sequence length.
|
||||
max_prompt_len: Maximum prompt tokens.
|
||||
compile: Whether to compile the model with torch.compile.
|
||||
page_size: Number of tokens per KV cache page.
|
||||
"""
|
||||
self.model = model
|
||||
self.tokenizer = tokenizer
|
||||
self.scheduler = InferenceScheduler(
|
||||
@@ -217,8 +85,9 @@ class InferenceEngine:
|
||||
tokenizer=self.tokenizer,
|
||||
max_batch_size=max_batch_size,
|
||||
max_seq_len=max_seq_len,
|
||||
max_prompt_len=max_prompt_len,
|
||||
page_size=page_size,
|
||||
cache=cache,
|
||||
enable_cuda_graph=enable_cuda_graph,
|
||||
backend=backend,
|
||||
)
|
||||
|
||||
self.scheduler.start()
|
||||
@@ -234,226 +103,140 @@ class InferenceEngine:
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
stream: bool = False,
|
||||
max_tokens: int = 1024,
|
||||
max_tokens: Optional[int] = None,
|
||||
temperature: float = 1.0,
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
) -> Union[Generator[str, None, None], str, List[str]]:
|
||||
"""Generates text from a prompt.
|
||||
|
||||
Args:
|
||||
prompt: Single string or list of strings for batch generation.
|
||||
stream: If True, returns a generator yielding tokens one by one.
|
||||
max_tokens: Maximum number of tokens to generate.
|
||||
temperature: Sampling temperature.
|
||||
top_p: Nucleus sampling probability threshold.
|
||||
top_k: Top-k sampling count (0 disables).
|
||||
|
||||
Returns:
|
||||
Generator (stream=True), single string (non-stream, single prompt),
|
||||
or list of strings (non-stream, batch prompts).
|
||||
"""
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
) -> Union[Generator, str, List[str]]:
|
||||
is_batch = isinstance(prompt, list)
|
||||
prompts = prompt if is_batch else [prompt]
|
||||
|
||||
if stream:
|
||||
return self._generate_streaming(
|
||||
prompts, is_batch, max_tokens, temperature, top_p, top_k
|
||||
)
|
||||
else:
|
||||
return self._generate_non_streaming(
|
||||
prompts, is_batch, max_tokens, temperature, top_p, top_k
|
||||
)
|
||||
if max_tokens is not None and max_tokens <= 0:
|
||||
if stream:
|
||||
return iter(())
|
||||
results = [""] * len(prompts)
|
||||
return results if is_batch else results[0]
|
||||
|
||||
return self._generate(
|
||||
prompts,
|
||||
is_batch,
|
||||
stream,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
|
||||
def generate_async(
|
||||
self,
|
||||
prompt: str,
|
||||
max_tokens: int = 1024,
|
||||
max_tokens: Optional[int] = None,
|
||||
temperature: float = 1.0,
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""Async streaming generator that does not block the event loop.
|
||||
|
||||
Runs the synchronous generator in a background thread pool executor,
|
||||
yielding tokens to the async consumer as they arrive.
|
||||
|
||||
Args:
|
||||
prompt: Input text to generate from.
|
||||
max_tokens: Maximum tokens to generate.
|
||||
temperature: Sampling temperature.
|
||||
top_p: Nucleus sampling threshold.
|
||||
top_k: Top-k sampling count.
|
||||
|
||||
Yields:
|
||||
Decoded token strings as they are generated.
|
||||
"""
|
||||
sync_gen = self._generate_streaming(
|
||||
[prompt], False, max_tokens, temperature, top_p, top_k
|
||||
sync_gen = self._generate(
|
||||
[prompt],
|
||||
False,
|
||||
True,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
|
||||
async def _agen():
|
||||
loop = asyncio.get_event_loop()
|
||||
while True:
|
||||
token = await loop.run_in_executor(None, self._next_token, sync_gen)
|
||||
if token is None:
|
||||
try:
|
||||
token = await loop.run_in_executor(None, next, sync_gen)
|
||||
except StopIteration:
|
||||
break
|
||||
yield token
|
||||
|
||||
return _agen()
|
||||
|
||||
@staticmethod
|
||||
def _next_token(gen: Generator) -> Optional[str]:
|
||||
"""Retrieves the next token from a synchronous generator.
|
||||
|
||||
Args:
|
||||
gen: A synchronous generator yielding token strings.
|
||||
|
||||
Returns:
|
||||
The next token, or None if the generator is exhausted.
|
||||
"""
|
||||
try:
|
||||
return next(gen)
|
||||
except StopIteration:
|
||||
return None
|
||||
|
||||
def generate_with_request(
|
||||
self, request: GenerationRequest
|
||||
) -> Union[Generator[str, None, None], str, List[str]]:
|
||||
"""Generates text from a structured GenerationRequest.
|
||||
|
||||
Applies the chat template to the request's messages before generation.
|
||||
|
||||
Args:
|
||||
request: A GenerationRequest with messages and parameters.
|
||||
|
||||
Returns:
|
||||
Generator, string, or list of strings (see generate()).
|
||||
"""
|
||||
prompt = self.tokenizer.apply_chat_template(request.messages, tokenize=False)
|
||||
return self.generate(
|
||||
prompt=prompt,
|
||||
stream=request.stream,
|
||||
max_tokens=request.params.max_tokens,
|
||||
temperature=request.params.temperature,
|
||||
top_p=request.params.top_p,
|
||||
top_k=request.params.top_k,
|
||||
)
|
||||
|
||||
def _generate_streaming(
|
||||
def _generate(
|
||||
self,
|
||||
prompts: List[str],
|
||||
is_batch: bool,
|
||||
max_tokens: int,
|
||||
stream: bool,
|
||||
max_tokens: Optional[int],
|
||||
temperature: float,
|
||||
top_p: float,
|
||||
top_k: int,
|
||||
) -> Generator[str, None, None]:
|
||||
"""Internal streaming generator.
|
||||
|
||||
Polls the _Result accumulator in a loop, yielding tokens as they arrive.
|
||||
Cleans up the scheduler task on GeneratorExit.
|
||||
|
||||
Args:
|
||||
prompts: List of prompts (only first is used; batch not yet supported).
|
||||
is_batch: If True, raises NotImplementedError.
|
||||
max_tokens: Maximum tokens to generate.
|
||||
temperature: Sampling temperature.
|
||||
top_p: Nucleus sampling threshold.
|
||||
top_k: Top-k sampling count.
|
||||
|
||||
Yields:
|
||||
Decoded token strings.
|
||||
"""
|
||||
if is_batch:
|
||||
raise NotImplementedError("Batch streaming not yet supported")
|
||||
|
||||
result = _Result()
|
||||
|
||||
task_id = self.scheduler.add_task(
|
||||
prompt=prompts[0],
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
stream_callback=lambda tok: result.append(tok, 0),
|
||||
)
|
||||
|
||||
def gen():
|
||||
try:
|
||||
while True:
|
||||
tokens = result.pop_all()
|
||||
for token in tokens:
|
||||
if token is STOP:
|
||||
return
|
||||
yield token
|
||||
if not result.wait(timeout=0.05):
|
||||
pass
|
||||
finally:
|
||||
self.scheduler.remove_task(task_id)
|
||||
|
||||
return gen()
|
||||
|
||||
def _generate_non_streaming(
|
||||
self,
|
||||
prompts: List[str],
|
||||
is_batch: bool,
|
||||
max_tokens: int,
|
||||
temperature: float,
|
||||
top_p: float,
|
||||
top_k: int,
|
||||
) -> Union[str, List[str]]:
|
||||
"""Internal non-streaming generator.
|
||||
|
||||
Submits all prompts to the scheduler and waits for all to complete.
|
||||
|
||||
Args:
|
||||
prompts: List of prompt strings.
|
||||
is_batch: Whether multiple prompts were provided.
|
||||
max_tokens: Maximum tokens to generate.
|
||||
temperature: Sampling temperature.
|
||||
top_p: Nucleus sampling threshold.
|
||||
top_k: Top-k sampling count.
|
||||
|
||||
Returns:
|
||||
Single string for one prompt, list of strings for batch.
|
||||
"""
|
||||
result = _Result(count=len(prompts))
|
||||
task_ids = []
|
||||
|
||||
for i, p in enumerate(prompts):
|
||||
|
||||
def make_cb(idx):
|
||||
return lambda tok: result.append(tok, idx)
|
||||
|
||||
task_id = self.scheduler.add_task(
|
||||
frequency_penalty: float,
|
||||
rep_window: int,
|
||||
) -> Union[Generator, str, List[str]]:
|
||||
n = len(prompts)
|
||||
request_backend = get_backend(use_default=False)
|
||||
result = GenerateResult(count=n)
|
||||
task_ids = [
|
||||
self.scheduler.add_task(
|
||||
prompt=p,
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
stream_callback=make_cb(i),
|
||||
frequency_penalty=frequency_penalty,
|
||||
rep_window=rep_window,
|
||||
backend=request_backend,
|
||||
stream_callback=lambda token, idx=i: result.append(token, idx),
|
||||
)
|
||||
task_ids.append(task_id)
|
||||
for i, p in enumerate(prompts)
|
||||
]
|
||||
|
||||
result.wait_completion()
|
||||
if not stream:
|
||||
try:
|
||||
result.wait_completion()
|
||||
except TimeoutError:
|
||||
for tid in task_ids:
|
||||
self.scheduler.remove_task(tid)
|
||||
raise
|
||||
for tid in task_ids:
|
||||
self.scheduler.remove_task(tid)
|
||||
res = result.get_results()
|
||||
return res if is_batch else res[0]
|
||||
|
||||
for task_id in task_ids:
|
||||
self.scheduler.remove_task(task_id)
|
||||
remaining = n
|
||||
finished = [False] * n
|
||||
|
||||
res = result.get_results()
|
||||
return res if is_batch else res[0]
|
||||
def gen():
|
||||
nonlocal remaining
|
||||
while remaining > 0:
|
||||
items = result.pop_all()
|
||||
for idx, token in items:
|
||||
if token is STOP:
|
||||
if not finished[idx]:
|
||||
finished[idx] = True
|
||||
remaining -= 1
|
||||
else:
|
||||
yield (idx, token) if is_batch else token
|
||||
if remaining > 0:
|
||||
result.wait(timeout=0.05)
|
||||
|
||||
return gen()
|
||||
|
||||
def get_stats(self) -> Dict[str, Any]:
|
||||
"""Returns current engine statistics.
|
||||
|
||||
Returns:
|
||||
Dict with total_tasks, total_tokens, active_tasks, waiting_queue.
|
||||
"""
|
||||
return self.scheduler.get_stats()
|
||||
|
||||
def shutdown(self) -> None:
|
||||
"""Shuts down the engine, stops the scheduler, and frees GPU memory."""
|
||||
@property
|
||||
def backend_name(self) -> str:
|
||||
return self.scheduler.backend_name
|
||||
|
||||
@property
|
||||
def cuda_graph_enabled(self) -> bool:
|
||||
return self.scheduler.cuda_graph_enabled
|
||||
|
||||
def shutdown(self):
|
||||
self.scheduler.stop()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
@@ -0,0 +1,223 @@
|
||||
"""Unified per-task perf/stats: timing records, context-manager scopes, aggregate reporting."""
|
||||
|
||||
import time
|
||||
from collections import deque
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Deque, Dict, Generator, List, Literal, Optional
|
||||
|
||||
|
||||
@dataclass
|
||||
class TaskTiming:
|
||||
"""Timestamp snapshots and computed metrics for one generation task.
|
||||
|
||||
Created by :class:`MetricsCollector` at task-registration time;
|
||||
updated via ``record`` / ``mark_finished``.
|
||||
"""
|
||||
|
||||
task_id: str
|
||||
arrival_time: float
|
||||
prefill_start_time: Optional[float] = None
|
||||
first_token_time: Optional[float] = None
|
||||
finish_time: Optional[float] = None
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
_decode_steps: int = 0
|
||||
_decode_total_s: float = 0.0
|
||||
|
||||
# derived metrics
|
||||
|
||||
@property
|
||||
def queue_wait_ms(self) -> Optional[float]:
|
||||
if self.prefill_start_time is not None:
|
||||
return (self.prefill_start_time - self.arrival_time) * 1000
|
||||
return None
|
||||
|
||||
@property
|
||||
def ttft_ms(self) -> Optional[float]:
|
||||
if self.first_token_time is not None:
|
||||
return (self.first_token_time - self.arrival_time) * 1000
|
||||
return None
|
||||
|
||||
@property
|
||||
def prefill_tps(self) -> Optional[float]:
|
||||
if self.prefill_start_time is not None and self.first_token_time is not None:
|
||||
d = self.first_token_time - self.prefill_start_time
|
||||
if d > 0 and self.input_tokens > 0:
|
||||
return self.input_tokens / d
|
||||
return None
|
||||
|
||||
@property
|
||||
def decode_tps(self) -> Optional[float]:
|
||||
if self.first_token_time is not None and self.finish_time is not None:
|
||||
d = self.finish_time - self.first_token_time
|
||||
dt = self.output_tokens - 1
|
||||
if dt > 0 and d > 0:
|
||||
return dt / d
|
||||
return None
|
||||
|
||||
@property
|
||||
def decode_avg_ms(self) -> Optional[float]:
|
||||
if self._decode_steps > 0 and self._decode_total_s > 0:
|
||||
return (self._decode_total_s / self._decode_steps) * 1000
|
||||
return None
|
||||
|
||||
@property
|
||||
def e2e_latency_ms(self) -> Optional[float]:
|
||||
if self.finish_time is not None:
|
||||
return (self.finish_time - self.arrival_time) * 1000
|
||||
return None
|
||||
|
||||
@property
|
||||
def total_tps(self) -> Optional[float]:
|
||||
if self.finish_time is not None:
|
||||
total = self.input_tokens + self.output_tokens
|
||||
d = self.finish_time - self.arrival_time
|
||||
if total > 0 and d > 0:
|
||||
return total / d
|
||||
return None
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"task_id": self.task_id,
|
||||
"input_tokens": self.input_tokens,
|
||||
"output_tokens": self.output_tokens,
|
||||
"queue_wait_ms": (
|
||||
round(self.queue_wait_ms, 2) if self.queue_wait_ms is not None else None
|
||||
),
|
||||
"ttft_ms": (round(self.ttft_ms, 2) if self.ttft_ms is not None else None),
|
||||
"prefill_tps": (
|
||||
round(self.prefill_tps, 2) if self.prefill_tps is not None else None
|
||||
),
|
||||
"decode_tps": (
|
||||
round(self.decode_tps, 2) if self.decode_tps is not None else None
|
||||
),
|
||||
"decode_avg_ms": (
|
||||
round(self.decode_avg_ms, 2) if self.decode_avg_ms is not None else None
|
||||
),
|
||||
"total_tps": (
|
||||
round(self.total_tps, 2) if self.total_tps is not None else None
|
||||
),
|
||||
"e2e_latency_ms": (
|
||||
round(self.e2e_latency_ms, 2)
|
||||
if self.e2e_latency_ms is not None
|
||||
else None
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
class MetricsCollector:
|
||||
"""Single-owner perf/stats hub for all generation tasks.
|
||||
|
||||
Usage::
|
||||
|
||||
metrics = MetricsCollector()
|
||||
metrics.register(task_id, arrival_time)
|
||||
|
||||
with metrics.record(task_ids, "prefill"):
|
||||
run_prefill(...)
|
||||
|
||||
metrics.mark_finished(task_id, input_tokens, output_tokens)
|
||||
|
||||
stats = metrics.get_stats()
|
||||
"""
|
||||
|
||||
def __init__(self, max_recent: int = 128):
|
||||
self._timings: Dict[str, TaskTiming] = {}
|
||||
self._completed: Deque[TaskTiming] = deque(maxlen=max_recent)
|
||||
|
||||
self._ttft_ms_sum = 0.0
|
||||
self._ttft_ms_count = 0
|
||||
self._decode_tps_sum = 0.0
|
||||
self._decode_tps_count = 0
|
||||
self._e2e_ms_sum = 0.0
|
||||
self._e2e_ms_count = 0
|
||||
|
||||
def register(self, task_id: str):
|
||||
"""Create a timing record for a newly-created task."""
|
||||
self._timings[task_id] = TaskTiming(task_id=task_id, arrival_time=time.time())
|
||||
|
||||
def mark_finished(self, task_id: str, input_tokens: int, output_tokens: int):
|
||||
"""Close timing for a finished/aborted task and move it to completed."""
|
||||
timing = self._timings.pop(task_id, None)
|
||||
if timing is None:
|
||||
return
|
||||
timing.finish_time = time.time()
|
||||
timing.input_tokens = input_tokens
|
||||
timing.output_tokens = output_tokens
|
||||
self._completed.append(timing)
|
||||
self._accumulate(timing)
|
||||
|
||||
def clear(self):
|
||||
"""Reset all state (e.g. on engine shutdown)."""
|
||||
self._timings.clear()
|
||||
self._completed.clear()
|
||||
self._ttft_ms_sum = 0.0
|
||||
self._ttft_ms_count = 0
|
||||
self._decode_tps_sum = 0.0
|
||||
self._decode_tps_count = 0
|
||||
self._e2e_ms_sum = 0.0
|
||||
self._e2e_ms_count = 0
|
||||
|
||||
# timing scopes
|
||||
|
||||
@contextmanager
|
||||
def record(
|
||||
self, task_ids: List[str], phase: Literal["prefill", "decode"]
|
||||
) -> Generator[None, None, None]:
|
||||
tic = time.time()
|
||||
yield
|
||||
toc = time.time()
|
||||
dt = toc - tic
|
||||
for tid in task_ids:
|
||||
t = self._timings.get(tid)
|
||||
if t is None:
|
||||
continue
|
||||
if phase == "prefill":
|
||||
t.prefill_start_time = tic
|
||||
t.first_token_time = toc
|
||||
elif phase == "decode":
|
||||
t._decode_steps += 1
|
||||
t._decode_total_s += dt
|
||||
|
||||
# access
|
||||
|
||||
def get_timing(self, task_id: str) -> Optional[TaskTiming]:
|
||||
"""Return the timing record for *task_id* (active or completed)."""
|
||||
if task_id in self._timings:
|
||||
return self._timings[task_id]
|
||||
for t in self._completed:
|
||||
if t.task_id == task_id:
|
||||
return t
|
||||
return None
|
||||
|
||||
# aggregate stats
|
||||
|
||||
def get_stats(self) -> Dict[str, Any]:
|
||||
stats: Dict[str, Any] = {}
|
||||
if self._ttft_ms_count > 0:
|
||||
stats["avg_ttft_ms"] = round(self._ttft_ms_sum / self._ttft_ms_count, 2)
|
||||
if self._decode_tps_count > 0:
|
||||
stats["avg_decode_tps"] = round(
|
||||
self._decode_tps_sum / self._decode_tps_count, 2
|
||||
)
|
||||
if self._e2e_ms_count > 0:
|
||||
stats["avg_e2e_latency_ms"] = round(
|
||||
self._e2e_ms_sum / self._e2e_ms_count, 2
|
||||
)
|
||||
if self._completed:
|
||||
stats["recent_tasks"] = [t.to_dict() for t in self._completed]
|
||||
return stats
|
||||
|
||||
# internal
|
||||
|
||||
def _accumulate(self, t: TaskTiming):
|
||||
if t.ttft_ms is not None:
|
||||
self._ttft_ms_sum += t.ttft_ms
|
||||
self._ttft_ms_count += 1
|
||||
if t.decode_tps is not None:
|
||||
self._decode_tps_sum += t.decode_tps
|
||||
self._decode_tps_count += 1
|
||||
if t.e2e_latency_ms is not None:
|
||||
self._e2e_ms_sum += t.e2e_latency_ms
|
||||
self._e2e_ms_count += 1
|
||||
@@ -0,0 +1,39 @@
|
||||
"""Inference API: protocol handler, stop checker, tool parsers, and FastAPI server.
|
||||
|
||||
``app`` is no longer a module-level global. Use :func:`get_app` to access the
|
||||
lazy singleton FastAPI instance.
|
||||
"""
|
||||
|
||||
from astrai.inference.network.app import (
|
||||
AnthropicMessage,
|
||||
ChatCompletionRequest,
|
||||
ChatMessage,
|
||||
FunctionDef,
|
||||
MessagesRequest,
|
||||
ToolDef,
|
||||
get_app,
|
||||
run_server,
|
||||
)
|
||||
from astrai.inference.network.protocol import GenContext, ProtocolHandler, StopChecker
|
||||
from astrai.inference.network.tool_parser import (
|
||||
BaseToolParser,
|
||||
SimpleJsonToolParser,
|
||||
ToolParserFactory,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ProtocolHandler",
|
||||
"StopChecker",
|
||||
"GenContext",
|
||||
"BaseToolParser",
|
||||
"SimpleJsonToolParser",
|
||||
"ToolParserFactory",
|
||||
"AnthropicMessage",
|
||||
"ChatCompletionRequest",
|
||||
"ChatMessage",
|
||||
"FunctionDef",
|
||||
"ToolDef",
|
||||
"MessagesRequest",
|
||||
"get_app",
|
||||
"run_server",
|
||||
]
|
||||
@@ -0,0 +1,142 @@
|
||||
"""Anthropic message completion response builder."""
|
||||
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Tuple, Union
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
from astrai.inference.network.protocol import (
|
||||
GenContext,
|
||||
ResponseBuilder,
|
||||
StopInfo,
|
||||
sse_event,
|
||||
)
|
||||
|
||||
|
||||
def _extract_text(content: Union[str, List[Dict[str, Any]]]) -> str:
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
for block in content:
|
||||
if isinstance(block, dict) and block.get("type") == "text":
|
||||
return block.get("text", "")
|
||||
return ""
|
||||
|
||||
|
||||
class AnthropicResponseBuilder(ResponseBuilder):
|
||||
def prepare(
|
||||
self, request: BaseModel, engine: InferenceEngine
|
||||
) -> Tuple[str, GenContext, List[str]]:
|
||||
messages: List[Dict[str, str]] = []
|
||||
system = getattr(request, "system", None)
|
||||
if system:
|
||||
messages.append({"role": "system", "content": system})
|
||||
for m in request.messages:
|
||||
text = _extract_text(m.content)
|
||||
if text:
|
||||
messages.append({"role": m.role, "content": text})
|
||||
prompt = engine.tokenizer.apply_chat_template(messages, tokenize=False)
|
||||
ctx = GenContext(
|
||||
resp_id=f"msg_{uuid.uuid4().hex[:24]}",
|
||||
created=int(time.time()),
|
||||
model=request.model,
|
||||
)
|
||||
stop_sequences = getattr(request, "stop_sequences", None) or []
|
||||
return prompt, ctx, stop_sequences
|
||||
|
||||
def format_stream_start(self, ctx: GenContext) -> List[str]:
|
||||
return [
|
||||
sse_event(
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": ctx.resp_id,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": ctx.model,
|
||||
"content": [],
|
||||
"usage": {"input_tokens": ctx.prompt_tokens},
|
||||
},
|
||||
},
|
||||
event="message_start",
|
||||
),
|
||||
sse_event(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 0,
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
},
|
||||
event="content_block_start",
|
||||
),
|
||||
]
|
||||
|
||||
def format_chunk(self, token: str, **kwargs) -> List[str]:
|
||||
return [
|
||||
sse_event(
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": token},
|
||||
},
|
||||
event="content_block_delta",
|
||||
)
|
||||
]
|
||||
|
||||
def format_stream_end(self, ctx: GenContext, stop: StopInfo) -> List[str]:
|
||||
events: List[str] = []
|
||||
if stop.matched:
|
||||
trimmed = stop.body[: stop.body.rfind(stop.matched)]
|
||||
unyielded = trimmed[len(stop.yielded) :]
|
||||
if unyielded:
|
||||
events.append(
|
||||
sse_event(
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": unyielded},
|
||||
},
|
||||
event="content_block_delta",
|
||||
)
|
||||
)
|
||||
events.append(
|
||||
sse_event(
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
event="content_block_stop",
|
||||
)
|
||||
)
|
||||
events.append(
|
||||
sse_event(
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {
|
||||
"stop_reason": "stop_sequence" if stop.matched else "end_turn",
|
||||
"stop_sequence": stop.matched,
|
||||
},
|
||||
"usage": {"output_tokens": ctx.completion_tokens},
|
||||
},
|
||||
event="message_delta",
|
||||
)
|
||||
)
|
||||
events.append(sse_event({"type": "message_stop"}, event="message_stop"))
|
||||
return events
|
||||
|
||||
def format_response(
|
||||
self, ctx: GenContext, content: str, stop: StopInfo
|
||||
) -> Dict[str, Any]:
|
||||
if stop.matched:
|
||||
content = content[: content.rfind(stop.matched)]
|
||||
return {
|
||||
"id": ctx.resp_id,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": ctx.model,
|
||||
"content": [{"type": "text", "text": content}],
|
||||
"stop_reason": "stop_sequence" if stop.matched else "end_turn",
|
||||
"stop_sequence": stop.matched,
|
||||
"usage": {
|
||||
"input_tokens": ctx.prompt_tokens,
|
||||
"output_tokens": ctx.completion_tokens,
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,206 @@
|
||||
"""
|
||||
OpenAI / Anthropic-compatible chat completion server backed by continuous-batching inference.
|
||||
|
||||
Protocol-specific formatting is delegated to ``astrai.inference.protocol``.
|
||||
This module owns the FastAPI app, request/response schemas, and dependency wiring.
|
||||
|
||||
``app`` is lazily constructed — importing this module does NOT create a FastAPI instance.
|
||||
Use :func:`get_app` to access the singleton.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import torch
|
||||
import uvicorn
|
||||
from fastapi import APIRouter, FastAPI, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
from astrai.inference.network.anthropic import AnthropicResponseBuilder
|
||||
from astrai.inference.network.openai import OpenAIResponseBuilder
|
||||
from astrai.inference.network.protocol import ProtocolHandler
|
||||
from astrai.model import AutoModel
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_app_instance: Optional[FastAPI] = None
|
||||
|
||||
|
||||
class ChatMessage(BaseModel):
|
||||
role: str
|
||||
content: Optional[str] = None
|
||||
tool_calls: Optional[List[Dict[str, Any]]] = None
|
||||
tool_call_id: Optional[str] = None
|
||||
|
||||
|
||||
class FunctionDef(BaseModel):
|
||||
name: str
|
||||
description: Optional[str] = None
|
||||
parameters: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
class ToolDef(BaseModel):
|
||||
type: str = "function"
|
||||
function: FunctionDef
|
||||
|
||||
|
||||
class ChatCompletionRequest(BaseModel):
|
||||
"""OpenAI Chat Completion API request body."""
|
||||
|
||||
model: str = "astrai"
|
||||
messages: List[ChatMessage]
|
||||
temperature: Optional[float] = Field(default=1.0, ge=0.0, le=2.0)
|
||||
top_p: Optional[float] = Field(default=1.0, ge=0.0, le=1.0)
|
||||
top_k: Optional[int] = Field(default=50, ge=1)
|
||||
stream: Optional[bool] = False
|
||||
stop: Optional[Union[str, List[str]]] = None
|
||||
max_tokens: Optional[int] = Field(default=2048, ge=1)
|
||||
n: Optional[int] = Field(default=1, ge=1)
|
||||
presence_penalty: Optional[float] = Field(default=0.0, ge=-2.0, le=2.0)
|
||||
frequency_penalty: Optional[float] = Field(default=0.0, ge=-2.0, le=2.0)
|
||||
logit_bias: Optional[Dict[int, float]] = None
|
||||
user: Optional[str] = None
|
||||
tools: Optional[List[ToolDef]] = None
|
||||
tool_choice: Optional[Union[str, Dict[str, Any]]] = "auto"
|
||||
|
||||
|
||||
class AnthropicMessage(BaseModel):
|
||||
role: str
|
||||
content: Union[str, List[Dict[str, Any]]]
|
||||
|
||||
|
||||
class MessagesRequest(BaseModel):
|
||||
"""Anthropic Messages API request body."""
|
||||
|
||||
model: str = "astrai"
|
||||
max_tokens: int = Field(default=1024, ge=1)
|
||||
messages: List[AnthropicMessage]
|
||||
system: Optional[str] = None
|
||||
temperature: Optional[float] = Field(default=1.0, ge=0.0, le=2.0)
|
||||
top_p: Optional[float] = Field(default=1.0, ge=0.0, le=1.0)
|
||||
top_k: Optional[int] = Field(default=50, ge=1)
|
||||
stream: Optional[bool] = False
|
||||
stop_sequences: Optional[List[str]] = None
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
config = app.state.server_config
|
||||
if not config.get("_test", False):
|
||||
try:
|
||||
app.state.engine = _create_engine(**config)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to load model: {e}")
|
||||
raise
|
||||
yield
|
||||
if app.state.engine:
|
||||
app.state.engine.shutdown()
|
||||
logger.info("Inference engine shutdown complete")
|
||||
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _create_engine(
|
||||
param_path: Path,
|
||||
device: str = "cuda",
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
max_batch_size: int = 16,
|
||||
max_seq_len: Optional[int] = None,
|
||||
) -> InferenceEngine:
|
||||
if not param_path.exists():
|
||||
raise FileNotFoundError(f"Parameter directory not found: {param_path}")
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||
model = AutoModel.from_pretrained(param_path)
|
||||
model.to(device=device, dtype=dtype)
|
||||
logger.info(f"Model loaded on {device} with dtype {dtype}")
|
||||
|
||||
engine = InferenceEngine(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=max_batch_size,
|
||||
max_seq_len=max_seq_len,
|
||||
)
|
||||
logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}")
|
||||
return engine
|
||||
|
||||
|
||||
def get_app() -> FastAPI:
|
||||
"""Return the singleton FastAPI instance (lazily created on first call)."""
|
||||
global _app_instance
|
||||
if _app_instance is None:
|
||||
_app_instance = FastAPI(
|
||||
title="AstrAI Inference Server",
|
||||
version="0.2.0",
|
||||
lifespan=lifespan,
|
||||
)
|
||||
_app_instance.include_router(router)
|
||||
_app_instance.state.server_config = {}
|
||||
_app_instance.state.engine = None
|
||||
return _app_instance
|
||||
|
||||
|
||||
def _get_engine() -> InferenceEngine:
|
||||
engine = get_app().state.engine
|
||||
if engine is None:
|
||||
raise HTTPException(status_code=503, detail="Engine not initialized")
|
||||
return engine
|
||||
|
||||
|
||||
@router.get("/health")
|
||||
async def health():
|
||||
app = get_app()
|
||||
return {
|
||||
"status": "ok",
|
||||
"model_loaded": app.state.engine is not None,
|
||||
}
|
||||
|
||||
|
||||
@router.get("/stats")
|
||||
async def get_stats():
|
||||
return _get_engine().get_stats()
|
||||
|
||||
|
||||
@router.post("/v1/chat/completions")
|
||||
async def chat_completion(request: ChatCompletionRequest):
|
||||
engine = _get_engine()
|
||||
handler = ProtocolHandler(request, engine, OpenAIResponseBuilder())
|
||||
return await handler.handle()
|
||||
|
||||
|
||||
@router.post("/v1/messages")
|
||||
async def create_message(request: MessagesRequest):
|
||||
engine = _get_engine()
|
||||
handler = ProtocolHandler(request, engine, AnthropicResponseBuilder())
|
||||
return await handler.handle()
|
||||
|
||||
|
||||
def run_server(
|
||||
param_path: Path,
|
||||
host: str = "0.0.0.0",
|
||||
port: int = 8000,
|
||||
reload: bool = False,
|
||||
device: str = "cuda",
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
max_batch_size: int = 16,
|
||||
max_seq_len: Optional[int] = None,
|
||||
):
|
||||
app = get_app()
|
||||
app.state.server_config = {
|
||||
"device": device,
|
||||
"dtype": dtype,
|
||||
"param_path": param_path,
|
||||
"max_batch_size": max_batch_size,
|
||||
"max_seq_len": max_seq_len,
|
||||
}
|
||||
uvicorn.run(
|
||||
app,
|
||||
host=host,
|
||||
port=port,
|
||||
reload=reload,
|
||||
)
|
||||
@@ -0,0 +1,277 @@
|
||||
"""OpenAI chat completion response builder."""
|
||||
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
from astrai.inference.network.protocol import (
|
||||
GenContext,
|
||||
ResponseBuilder,
|
||||
StopInfo,
|
||||
sse_event,
|
||||
)
|
||||
from astrai.inference.network.tool_parser import BaseToolParser, ToolParserFactory
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_UNSUPPORTED_PARAMS = (
|
||||
"n",
|
||||
"presence_penalty",
|
||||
"logit_bias",
|
||||
"user",
|
||||
)
|
||||
|
||||
|
||||
def _resolve_tool_choice(
|
||||
request: BaseModel,
|
||||
) -> Union[str, Dict[str, Any]]:
|
||||
tc = getattr(request, "tool_choice", None)
|
||||
if tc is None:
|
||||
return "auto"
|
||||
if isinstance(tc, str):
|
||||
return tc
|
||||
if isinstance(tc, dict):
|
||||
return tc
|
||||
return "auto"
|
||||
|
||||
|
||||
def _resolve_tools(request: BaseModel) -> Optional[List[Dict[str, Any]]]:
|
||||
raw = getattr(request, "tools", None)
|
||||
if not raw:
|
||||
return None
|
||||
if isinstance(raw, list):
|
||||
return [t.model_dump() if hasattr(t, "model_dump") else t for t in raw]
|
||||
return None
|
||||
|
||||
|
||||
class OpenAIResponseBuilder(ResponseBuilder):
|
||||
def prepare(
|
||||
self, request: BaseModel, engine: InferenceEngine
|
||||
) -> Tuple[str, GenContext, List[str]]:
|
||||
messages = [{"role": m.role, "content": m.content} for m in request.messages]
|
||||
tools = _resolve_tools(request)
|
||||
prompt = engine.tokenizer.apply_chat_template(
|
||||
messages, tokenize=False, tools=tools or []
|
||||
)
|
||||
|
||||
self._resp_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
||||
self._model = request.model
|
||||
|
||||
for param in _UNSUPPORTED_PARAMS:
|
||||
value = getattr(request, param, None)
|
||||
fields = getattr(type(request), "model_fields", {})
|
||||
default = fields[param].default if param in fields else None
|
||||
if value is not None and value != default:
|
||||
logger.warning(
|
||||
"ChatCompletionRequest param '%s'=%r is not supported"
|
||||
" and will be ignored",
|
||||
param,
|
||||
value,
|
||||
)
|
||||
|
||||
self._parser: Optional[BaseToolParser] = None
|
||||
if tools:
|
||||
tool_choice = _resolve_tool_choice(request)
|
||||
self._parser = ToolParserFactory.create(
|
||||
"simple_json", tools=tools, tool_choice=tool_choice
|
||||
)
|
||||
self._content_started = False
|
||||
|
||||
ctx = GenContext(
|
||||
resp_id=self._resp_id,
|
||||
created=int(time.time()),
|
||||
model=self._model,
|
||||
)
|
||||
stop = request.stop
|
||||
stop_sequences = (
|
||||
[] if stop is None else [stop] if isinstance(stop, str) else stop
|
||||
)
|
||||
return prompt, ctx, stop_sequences
|
||||
|
||||
def format_stream_start(self, ctx: GenContext) -> List[str]:
|
||||
return [
|
||||
sse_event(
|
||||
{
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": ctx.created,
|
||||
"model": self._model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"role": "assistant"},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
]
|
||||
|
||||
def format_chunk(self, token: str, **kwargs) -> List[str]:
|
||||
body = kwargs.get("body", "")
|
||||
if self._parser is not None:
|
||||
return self._format_tool_chunk(body, **kwargs)
|
||||
|
||||
return [
|
||||
sse_event(
|
||||
{
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 0,
|
||||
"model": self._model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"content": token},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
]
|
||||
|
||||
def _format_tool_chunk(self, body: str, **kwargs) -> List[str]:
|
||||
deltas = self._parser.feed(
|
||||
body,
|
||||
current_token_ids=kwargs.get("current_token_ids"),
|
||||
delta_token_ids=kwargs.get("delta_token_ids"),
|
||||
)
|
||||
events: List[str] = []
|
||||
for d in deltas:
|
||||
if "content" in d:
|
||||
if not self._content_started:
|
||||
events.append(self._role_chunk())
|
||||
self._content_started = True
|
||||
events.append(
|
||||
sse_event(
|
||||
{
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 0,
|
||||
"model": self._model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"content": d["content"]},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
)
|
||||
elif "tool_calls" in d:
|
||||
if not self._content_started:
|
||||
events.append(self._role_chunk())
|
||||
self._content_started = True
|
||||
events.append(
|
||||
sse_event(
|
||||
{
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 0,
|
||||
"model": self._model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"tool_calls": d["tool_calls"]},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
)
|
||||
return events
|
||||
|
||||
def _role_chunk(self) -> str:
|
||||
return sse_event(
|
||||
{
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 0,
|
||||
"model": self._model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"role": "assistant"},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
def format_stream_end(self, ctx: GenContext, stop: StopInfo) -> List[str]:
|
||||
finish_reason = "stop"
|
||||
if self._parser is not None and self._parser.has_tool_calls:
|
||||
finish_reason = "tool_calls"
|
||||
return [
|
||||
sse_event(
|
||||
{
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": ctx.created,
|
||||
"model": self._model,
|
||||
"choices": [
|
||||
{"index": 0, "delta": {}, "finish_reason": finish_reason}
|
||||
],
|
||||
}
|
||||
),
|
||||
sse_event(
|
||||
{
|
||||
"prompt_tokens": ctx.prompt_tokens,
|
||||
"completion_tokens": ctx.completion_tokens,
|
||||
"total_tokens": ctx.prompt_tokens + ctx.completion_tokens,
|
||||
}
|
||||
),
|
||||
]
|
||||
|
||||
def format_response(
|
||||
self, ctx: GenContext, content: str, stop: StopInfo
|
||||
) -> Dict[str, Any]:
|
||||
if self._parser is not None:
|
||||
parsed = self._parser.parse_complete(content)
|
||||
if parsed and parsed.get("tool_calls"):
|
||||
return {
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion",
|
||||
"created": ctx.created,
|
||||
"model": self._model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": parsed.get("content"),
|
||||
"tool_calls": parsed["tool_calls"],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": ctx.prompt_tokens,
|
||||
"completion_tokens": ctx.completion_tokens,
|
||||
"total_tokens": ctx.prompt_tokens + ctx.completion_tokens,
|
||||
},
|
||||
}
|
||||
|
||||
return {
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion",
|
||||
"created": ctx.created,
|
||||
"model": self._model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": content},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": ctx.prompt_tokens,
|
||||
"completion_tokens": ctx.completion_tokens,
|
||||
"total_tokens": ctx.prompt_tokens + ctx.completion_tokens,
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,197 @@
|
||||
"""Orchestration layer: ProtocolHandler, StopChecker, GenContext, StopInfo, ResponseBuilder, SSE utils.
|
||||
|
||||
ProtocolHandler orchestrates the async generation loop and delegates
|
||||
protocol-specific formatting to a ResponseBuilder.
|
||||
"""
|
||||
|
||||
import json
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, AsyncGenerator, Dict, List, Optional, Tuple, Union
|
||||
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import BaseModel
|
||||
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
|
||||
|
||||
def sse_event(data: Dict[str, Any], event: Optional[str] = None) -> str:
|
||||
lines: List[str] = []
|
||||
if event:
|
||||
lines.append(f"event: {event}")
|
||||
lines.append(f"data: {json.dumps(data, ensure_ascii=False)}")
|
||||
lines.append("")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def sse_done() -> str:
|
||||
return "data: [DONE]\n\n"
|
||||
|
||||
|
||||
@dataclass
|
||||
class GenContext:
|
||||
"""Per-generation metadata passed to builder format methods."""
|
||||
|
||||
resp_id: str
|
||||
created: int
|
||||
model: str
|
||||
prompt_tokens: int = 0
|
||||
completion_tokens: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class StopInfo:
|
||||
"""Stop-check result passed to format_stream_end / format_response."""
|
||||
|
||||
matched: Optional[str] = None
|
||||
body: str = ""
|
||||
yielded: str = ""
|
||||
|
||||
|
||||
class StopChecker:
|
||||
"""Scans accumulated text for stop sequence matches."""
|
||||
|
||||
def __init__(self, sequences: List[str]):
|
||||
self._sequences = [s for s in sequences if s]
|
||||
|
||||
def check(self, text: str) -> Optional[str]:
|
||||
for seq in self._sequences:
|
||||
if seq in text:
|
||||
return seq
|
||||
return None
|
||||
|
||||
|
||||
class ResponseBuilder(ABC):
|
||||
"""Interface for protocol-specific response formatting.
|
||||
|
||||
A new protocol requires one concrete builder implementing 5 methods.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def prepare(
|
||||
self, request: BaseModel, engine: InferenceEngine
|
||||
) -> Tuple[str, GenContext, List[str]]:
|
||||
"""Return (prompt, ctx, stop_sequences) for a generation request."""
|
||||
|
||||
@abstractmethod
|
||||
def format_stream_start(self, ctx: GenContext) -> List[str]:
|
||||
"""SSE events that open the stream."""
|
||||
|
||||
@abstractmethod
|
||||
def format_chunk(self, token: str, **kwargs) -> List[str]:
|
||||
"""SSE events for a single generated token.
|
||||
|
||||
``body`` (the full accumulated text so far) is always provided
|
||||
as a keyword argument. Additional keyword arguments such as
|
||||
``current_token_ids`` and ``delta_token_ids`` may be included
|
||||
for tool parsers that need token-level information.
|
||||
Returns a list of SSE event strings (may be empty).
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def format_stream_end(self, ctx: GenContext, stop: StopInfo) -> List[str]:
|
||||
"""SSE events that close the stream."""
|
||||
|
||||
@abstractmethod
|
||||
def format_response(
|
||||
self, ctx: GenContext, content: str, stop: StopInfo
|
||||
) -> Dict[str, Any]:
|
||||
"""JSON response body for non-streaming mode."""
|
||||
|
||||
|
||||
class ProtocolHandler:
|
||||
"""Orchestrates the generation loop, delegates formatting to a builder.
|
||||
|
||||
Usage::
|
||||
|
||||
handler = ProtocolHandler(request, engine, OpenAIResponseBuilder())
|
||||
response = await handler.handle()
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, request: BaseModel, engine: InferenceEngine, builder: ResponseBuilder
|
||||
):
|
||||
self.request = request
|
||||
self.engine = engine
|
||||
self.builder = builder
|
||||
|
||||
async def handle(self) -> Union[StreamingResponse, Dict[str, Any]]:
|
||||
prompt, ctx, stop_sequences = self.builder.prepare(self.request, self.engine)
|
||||
ctx.prompt_tokens = len(self.engine.tokenizer.encode(prompt))
|
||||
|
||||
agen = self.engine.generate_async(
|
||||
prompt=prompt,
|
||||
max_tokens=self.request.max_tokens,
|
||||
temperature=self.request.temperature,
|
||||
top_p=self.request.top_p,
|
||||
top_k=self.request.top_k,
|
||||
frequency_penalty=getattr(self.request, "frequency_penalty", 0.0),
|
||||
)
|
||||
|
||||
if self.request.stream:
|
||||
return self._handle_stream(agen, ctx, stop_sequences)
|
||||
else:
|
||||
return await self._handle_non_stream(agen, ctx, stop_sequences)
|
||||
|
||||
def _handle_stream(
|
||||
self, agen: AsyncGenerator, ctx: GenContext, stop_sequences: List[str]
|
||||
) -> StreamingResponse:
|
||||
checker = StopChecker(stop_sequences)
|
||||
|
||||
async def event_stream():
|
||||
for event in self.builder.format_stream_start(ctx):
|
||||
yield event
|
||||
|
||||
body = ""
|
||||
yielded = ""
|
||||
matched = None
|
||||
token_ids: List[int] = []
|
||||
async for token in agen:
|
||||
body += token
|
||||
|
||||
new_ids = self.engine.tokenizer.encode(token)
|
||||
token_ids.extend(new_ids)
|
||||
|
||||
matched = checker.check(body)
|
||||
if matched:
|
||||
break
|
||||
|
||||
ctx.completion_tokens += 1
|
||||
for event in self.builder.format_chunk(
|
||||
token,
|
||||
body=body,
|
||||
current_token_ids=token_ids,
|
||||
delta_token_ids=new_ids,
|
||||
):
|
||||
yield event
|
||||
yielded += token
|
||||
|
||||
stop = StopInfo(matched=matched, body=body, yielded=yielded)
|
||||
for event in self.builder.format_stream_end(ctx, stop):
|
||||
yield event
|
||||
yield sse_done()
|
||||
|
||||
return StreamingResponse(
|
||||
event_stream(),
|
||||
media_type="text/event-stream",
|
||||
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
||||
)
|
||||
|
||||
async def _handle_non_stream(
|
||||
self, agen: AsyncGenerator, ctx: GenContext, stop_sequences: List[str]
|
||||
) -> Dict[str, Any]:
|
||||
checker = StopChecker(stop_sequences)
|
||||
body = ""
|
||||
matched = None
|
||||
|
||||
async for token in agen:
|
||||
body += token
|
||||
|
||||
matched = checker.check(body)
|
||||
if matched:
|
||||
break
|
||||
|
||||
ctx.completion_tokens += 1
|
||||
|
||||
stop = StopInfo(matched=matched, body=body)
|
||||
return self.builder.format_response(ctx, body, stop)
|
||||
@@ -0,0 +1,339 @@
|
||||
"""Tool call parsers for extracting structured tool calls from model output.
|
||||
|
||||
Patterned after vLLM's ToolParser abstraction. Each parser knows how to
|
||||
detect and incrementally extract tool calls from raw generated text.
|
||||
|
||||
Subclasses may optionally consume ``token_ids`` for token-level parsing
|
||||
(e.g. Harmony / VLM-style parsers).
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
import uuid
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
|
||||
class BaseToolParser(ABC):
|
||||
"""Abstract tool call parser — one instance per request.
|
||||
|
||||
Maintains streaming state internally so that each call to :meth:`feed`
|
||||
can diff against previously emitted content.
|
||||
|
||||
Args:
|
||||
tools (list of dict, optional): Tool definitions from the request.
|
||||
tool_choice (str): ``"auto"`` / ``"required"`` / ``"none"`` or a named
|
||||
tool choice dict.
|
||||
"""
|
||||
|
||||
def __init__(self, tools: Optional[List[Dict]] = None, tool_choice: str = "auto"):
|
||||
self.tools = tools or []
|
||||
self.tool_choice = tool_choice
|
||||
|
||||
@abstractmethod
|
||||
def feed(
|
||||
self,
|
||||
body: str,
|
||||
current_token_ids: Optional[List[int]] = None,
|
||||
delta_token_ids: Optional[List[int]] = None,
|
||||
) -> List[Dict]:
|
||||
"""Feed the *full* accumulated text each step.
|
||||
|
||||
Returns a list of delta dicts to emit. Each delta is one of:
|
||||
|
||||
- ``{"content": "text"}`` — plain text delta
|
||||
- ``{"tool_calls": [...]}`` — tool-call delta (OpenAI format)
|
||||
|
||||
Returns an empty list when nothing new should be emitted.
|
||||
|
||||
Args:
|
||||
body (str): The complete accumulated generated text so far.
|
||||
current_token_ids (list of int, optional): All token IDs decoded
|
||||
into *body* (cumulative).
|
||||
delta_token_ids (list of int, optional): Only the token IDs for
|
||||
this chunk.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def parse_complete(self, body: str) -> Optional[Dict]:
|
||||
"""Parse the *complete* generated text after generation ends.
|
||||
|
||||
Returns ``None`` when no tool calls were found, otherwise a dict
|
||||
with ``content`` (str or None) and ``tool_calls`` (list of dicts).
|
||||
"""
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def has_tool_calls(self) -> bool:
|
||||
"""True if the parser detected at least one tool call in the stream."""
|
||||
|
||||
|
||||
class ToolParserFactory(BaseFactory["BaseToolParser"]):
|
||||
pass
|
||||
|
||||
|
||||
_TOOL_CALL_HEAD_RE = re.compile(r'\{\s*"name"\s*:')
|
||||
|
||||
|
||||
def _scan_json(text: str, start: int = 0):
|
||||
"""Scan for a complete JSON object starting at *start*.
|
||||
|
||||
Returns ``(end, complete)`` where *end* is one-past the closing
|
||||
brace (or ``len(text)`` if unclosed), and *complete* is a bool.
|
||||
"""
|
||||
depth = 0
|
||||
in_string = False
|
||||
escape = False
|
||||
for i in range(start, len(text)):
|
||||
c = text[i]
|
||||
if escape:
|
||||
escape = False
|
||||
continue
|
||||
if c == "\\":
|
||||
escape = True
|
||||
continue
|
||||
if c == '"':
|
||||
in_string = not in_string
|
||||
continue
|
||||
if in_string:
|
||||
continue
|
||||
if c == "{":
|
||||
depth += 1
|
||||
elif c == "}":
|
||||
depth -= 1
|
||||
if depth == 0:
|
||||
return i + 1, True
|
||||
return len(text), False
|
||||
|
||||
|
||||
def _parse_tool_call_json(json_str: str, complete: bool):
|
||||
"""Extract *name* and *arguments* from a tool-call JSON string.
|
||||
|
||||
Returns ``(name, args, valid)``.
|
||||
"""
|
||||
if complete:
|
||||
try:
|
||||
obj = json.loads(json_str)
|
||||
except json.JSONDecodeError:
|
||||
return None, "", False
|
||||
name = obj.get("name")
|
||||
if not isinstance(name, str) or not name:
|
||||
return None, "", False
|
||||
args = obj.get("arguments")
|
||||
if isinstance(args, dict):
|
||||
if not args:
|
||||
args = ""
|
||||
else:
|
||||
args = json.dumps(args, ensure_ascii=False)
|
||||
args = args[1:-1].rstrip()
|
||||
elif isinstance(args, list):
|
||||
args = json.dumps(args, ensure_ascii=False) if args else ""
|
||||
elif isinstance(args, str):
|
||||
pass
|
||||
else:
|
||||
args = str(args) if args is not None else ""
|
||||
return name, args, True
|
||||
|
||||
name_match = re.search(r'"name"\s*:\s*"([^"]*)"', json_str)
|
||||
if not name_match:
|
||||
return None, "", False
|
||||
name = name_match.group(1)
|
||||
|
||||
args_match = re.search(r'"arguments"\s*:\s*(.*)', json_str, re.DOTALL)
|
||||
if not args_match:
|
||||
return name, "", True
|
||||
|
||||
raw = args_match.group(1).rstrip()
|
||||
if raw.startswith("{"):
|
||||
inner = raw[1:].rstrip()
|
||||
if inner.endswith("}"):
|
||||
inner = inner[:-1].rstrip()
|
||||
raw = inner
|
||||
return name, raw, True
|
||||
|
||||
|
||||
def _find_tool_calls(text: str, start_pos: int = 0):
|
||||
"""Find all complete ``{...}`` tool-call objects in *text*.
|
||||
|
||||
Returns a list of dicts with keys *start*, *end*, *name*, *args*,
|
||||
*complete*.
|
||||
"""
|
||||
results = []
|
||||
pos = start_pos
|
||||
|
||||
while True:
|
||||
brace = text.find("{", pos)
|
||||
if brace == -1:
|
||||
break
|
||||
|
||||
end, complete = _scan_json(text, brace)
|
||||
if not complete:
|
||||
break
|
||||
|
||||
json_str = text[brace:end]
|
||||
|
||||
name, args, valid = _parse_tool_call_json(json_str, complete=True)
|
||||
if not valid or name is None:
|
||||
pos = end
|
||||
continue
|
||||
|
||||
results.append(
|
||||
{
|
||||
"start": brace,
|
||||
"end": end,
|
||||
"name": name,
|
||||
"args": args,
|
||||
"complete": True,
|
||||
}
|
||||
)
|
||||
pos = end
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def _find_partial_tool_call(text: str, start_pos: int = 0):
|
||||
"""Find one incomplete (still-generating) tool-call JSON object."""
|
||||
brace = text.find("{", start_pos)
|
||||
if brace == -1:
|
||||
return None
|
||||
|
||||
json_str = text[brace:]
|
||||
if '"name"' not in json_str:
|
||||
return None
|
||||
|
||||
name, args, valid = _parse_tool_call_json(json_str, complete=False)
|
||||
if not valid or name is None:
|
||||
return None
|
||||
|
||||
return {
|
||||
"start": brace,
|
||||
"name": name,
|
||||
"args": args,
|
||||
"complete": False,
|
||||
}
|
||||
|
||||
|
||||
@ToolParserFactory.register("simple_json")
|
||||
class SimpleJsonToolParser(BaseToolParser):
|
||||
"""Parser for models that output tool calls as plain JSON objects.
|
||||
|
||||
Detects ``{"name": "<func>", "arguments": {...}}`` anywhere in the
|
||||
generated text. Handles single and (non-overlapping) multiple tool
|
||||
calls. Text preceding the first tool call is emitted as plain
|
||||
``content`` deltas.
|
||||
"""
|
||||
|
||||
def __init__(self, tools=None, tool_choice="auto"):
|
||||
super().__init__(tools, tool_choice)
|
||||
self._emitted_content_len = 0
|
||||
self._tc_state: List[Dict] = []
|
||||
self._has_tool_calls = False
|
||||
|
||||
# -------------------------------------------------------------- feed
|
||||
|
||||
def feed(
|
||||
self,
|
||||
body: str,
|
||||
current_token_ids: Optional[List[int]] = None,
|
||||
delta_token_ids: Optional[List[int]] = None,
|
||||
) -> List[Dict]:
|
||||
deltas: List[Dict] = []
|
||||
|
||||
completed = _find_tool_calls(body)
|
||||
|
||||
if not completed:
|
||||
partial = _find_partial_tool_call(body)
|
||||
if not partial:
|
||||
return self._emit_plain_content(body, deltas)
|
||||
all_tcs = [partial]
|
||||
else:
|
||||
all_tcs = completed
|
||||
partial = _find_partial_tool_call(body, completed[-1]["end"])
|
||||
if partial:
|
||||
all_tcs = completed + [partial]
|
||||
|
||||
first_start = all_tcs[0]["start"]
|
||||
if first_start > self._emitted_content_len:
|
||||
content = body[self._emitted_content_len : first_start]
|
||||
self._emitted_content_len = first_start
|
||||
if content:
|
||||
deltas.append({"content": content})
|
||||
|
||||
for i, tc in enumerate(all_tcs):
|
||||
if i >= len(self._tc_state):
|
||||
self._tc_state.append(
|
||||
{
|
||||
"id": f"call_{uuid.uuid4().hex[:12]}",
|
||||
"name_emitted": False,
|
||||
"args_emitted_len": 0,
|
||||
}
|
||||
)
|
||||
self._has_tool_calls = True
|
||||
st = self._tc_state[i]
|
||||
|
||||
if not st["name_emitted"]:
|
||||
st["name_emitted"] = True
|
||||
deltas.append(
|
||||
{
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": i,
|
||||
"id": st["id"],
|
||||
"type": "function",
|
||||
"function": {"name": tc["name"], "arguments": ""},
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
|
||||
new_args = tc["args"]
|
||||
if len(new_args) > st["args_emitted_len"]:
|
||||
diff = new_args[st["args_emitted_len"] :]
|
||||
st["args_emitted_len"] = len(new_args)
|
||||
deltas.append(
|
||||
{
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": i,
|
||||
"function": {"arguments": diff},
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
|
||||
return deltas
|
||||
|
||||
def _emit_plain_content(self, body: str, deltas: List[Dict]) -> List[Dict]:
|
||||
new_content = body[self._emitted_content_len :]
|
||||
if new_content:
|
||||
self._emitted_content_len = len(body)
|
||||
deltas.append({"content": new_content})
|
||||
return deltas
|
||||
|
||||
# -------------------------------------------------------- complete
|
||||
|
||||
def parse_complete(self, body: str) -> Optional[Dict]:
|
||||
completed = _find_tool_calls(body)
|
||||
if not completed:
|
||||
return None
|
||||
|
||||
content = body[: completed[0]["start"]].strip() or None
|
||||
tool_calls = []
|
||||
for i, tc in enumerate(completed):
|
||||
tool_calls.append(
|
||||
{
|
||||
"id": f"call_{uuid.uuid4().hex[:12]}",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tc["name"],
|
||||
"arguments": tc["args"],
|
||||
},
|
||||
}
|
||||
)
|
||||
return {"content": content, "tool_calls": tool_calls}
|
||||
|
||||
@property
|
||||
def has_tool_calls(self) -> bool:
|
||||
return self._has_tool_calls
|
||||
@@ -0,0 +1,25 @@
|
||||
"""Execution primitives: forward passes, CUDA graphs, and sampling."""
|
||||
|
||||
from astrai.inference.runtime.executor import Executor
|
||||
from astrai.inference.runtime.graph import CudaGraphContext
|
||||
from astrai.inference.runtime.sample import (
|
||||
BaseSamplingStrategy,
|
||||
FrequencyPenaltyStrategy,
|
||||
SamplingPipeline,
|
||||
TemperatureStrategy,
|
||||
TopKStrategy,
|
||||
TopPStrategy,
|
||||
sample,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Executor",
|
||||
"CudaGraphContext",
|
||||
"BaseSamplingStrategy",
|
||||
"FrequencyPenaltyStrategy",
|
||||
"SamplingPipeline",
|
||||
"TemperatureStrategy",
|
||||
"TopKStrategy",
|
||||
"TopPStrategy",
|
||||
"sample",
|
||||
]
|
||||
@@ -0,0 +1,421 @@
|
||||
import logging
|
||||
import time
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.extension.backend.attention import (
|
||||
CudaBackend,
|
||||
get_backend,
|
||||
)
|
||||
from astrai.inference.cache import PagePool, TaskCacheManager
|
||||
from astrai.inference.runtime.graph import CudaGraphContext
|
||||
from astrai.inference.runtime.sample import sample
|
||||
from astrai.inference.task import Task
|
||||
from astrai.inference.workspace import InferenceWorkspace
|
||||
from astrai.model.automodel import AutoModel
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def timed(label: str, log: Optional[logging.Logger] = None):
|
||||
"""GPU-precise timer via CUDA events; falls back to perf_counter on CPU."""
|
||||
log = log or logger
|
||||
if not log.isEnabledFor(logging.DEBUG):
|
||||
yield
|
||||
return
|
||||
use_cuda = torch.cuda.is_available()
|
||||
if use_cuda:
|
||||
start = torch.cuda.Event(enable_timing=True)
|
||||
end = torch.cuda.Event(enable_timing=True)
|
||||
start.record()
|
||||
else:
|
||||
tic = time.perf_counter()
|
||||
yield
|
||||
if use_cuda:
|
||||
end.record()
|
||||
torch.cuda.synchronize()
|
||||
elapsed_ms = start.elapsed_time(end)
|
||||
else:
|
||||
elapsed_ms = (time.perf_counter() - tic) * 1000
|
||||
log.debug("%s %.2fms", label, elapsed_ms)
|
||||
|
||||
|
||||
@dataclass
|
||||
class SamplingBatchInfo:
|
||||
"""Per-batch sampling parameters, cached across decode steps.
|
||||
|
||||
Sampling params are constant for a given ordered task set, so they are
|
||||
built once (pinned-memory async H2D) and reused until the task set
|
||||
changes. ``top_ks`` is int32 to match the native consumers.
|
||||
"""
|
||||
|
||||
temperatures: Tensor # float32 [B]
|
||||
top_ks: Tensor # int32 [B]
|
||||
top_ps: Tensor # float32 [B]
|
||||
freq_penalties: Tensor # float32 [B]
|
||||
has_freq: bool # any frequency_penalty != 0 (avoids per-step GPU .any())
|
||||
|
||||
|
||||
@dataclass
|
||||
class DecodeSteadyState:
|
||||
"""Cached decode metadata for the steady-state case.
|
||||
|
||||
When the same ordered task set decodes one token per step, sampling
|
||||
params and task signature are reused; only positions advance by 1.
|
||||
"""
|
||||
|
||||
task_sig: tuple
|
||||
positions: list[int]
|
||||
sampling_info: SamplingBatchInfo
|
||||
|
||||
|
||||
def _build_sampling_batch_info(tasks: List[Task], device) -> SamplingBatchInfo:
|
||||
pin = str(device).startswith("cuda")
|
||||
freq_penalties = torch.tensor(
|
||||
[t.frequency_penalty for t in tasks], dtype=torch.float32, pin_memory=pin
|
||||
).to(device, non_blocking=True)
|
||||
return SamplingBatchInfo(
|
||||
temperatures=torch.tensor(
|
||||
[t.temperature for t in tasks], dtype=torch.float32, pin_memory=pin
|
||||
).to(device, non_blocking=True),
|
||||
top_ks=torch.tensor(
|
||||
[t.top_k for t in tasks], dtype=torch.int32, pin_memory=pin
|
||||
).to(device, non_blocking=True),
|
||||
top_ps=torch.tensor(
|
||||
[t.top_p for t in tasks], dtype=torch.float32, pin_memory=pin
|
||||
).to(device, non_blocking=True),
|
||||
freq_penalties=freq_penalties,
|
||||
has_freq=bool((freq_penalties != 0).any()),
|
||||
)
|
||||
|
||||
|
||||
def _warmup_cuda_graphs(
|
||||
model: AutoModel,
|
||||
pool: PagePool,
|
||||
task_cache: TaskCacheManager,
|
||||
ws: InferenceWorkspace,
|
||||
gctx: CudaGraphContext,
|
||||
max_batch_size: int,
|
||||
prompt_len: int = 1,
|
||||
device: Optional[str] = None,
|
||||
):
|
||||
dev = device or next(model.parameters()).device
|
||||
|
||||
# Prefill warmup: cuBLAS auto-tunes for the actual prompt-length tensor
|
||||
# shapes on first call (F.linear is the dominant cost). This also warms
|
||||
# up the CUDA context (driver init) and compiles the graph-capture trace
|
||||
# that follows. Custom .so kernels do NOT need this — they are pre-built.
|
||||
warmup_len = 64
|
||||
tid = "_warmup_prefill"
|
||||
if task_cache.task_alloc(tid, list(range(warmup_len))):
|
||||
with (
|
||||
torch.inference_mode(),
|
||||
timed("warmup prefill", logger),
|
||||
):
|
||||
kv = task_cache.bind([tid], ws, start_pos=0)
|
||||
ids_in = torch.arange(warmup_len, device=dev)
|
||||
pos_in = ids_in
|
||||
model(
|
||||
ids_in,
|
||||
kv_cache=kv,
|
||||
position_ids=pos_in,
|
||||
fwd="prefill",
|
||||
)
|
||||
task_cache.task_free(tid)
|
||||
|
||||
batch_sizes = [1]
|
||||
n = 2
|
||||
while n <= max_batch_size:
|
||||
batch_sizes.append(n)
|
||||
n *= 2
|
||||
if max_batch_size not in batch_sizes:
|
||||
batch_sizes.append(max_batch_size)
|
||||
|
||||
for b in batch_sizes:
|
||||
task_ids = [f"_warmup_decode_{b}_{i}" for i in range(b)]
|
||||
prompt_tokens = [list(range(prompt_len)) for _ in range(b)]
|
||||
alloc_ok = True
|
||||
for tid, pt in zip(task_ids, prompt_tokens):
|
||||
if not task_cache.task_alloc(tid, pt):
|
||||
alloc_ok = False
|
||||
break
|
||||
if not alloc_ok:
|
||||
for tid in task_ids:
|
||||
task_cache.task_free(tid)
|
||||
continue
|
||||
|
||||
with (
|
||||
torch.inference_mode(),
|
||||
timed(f"warmup decode b={b}", logger),
|
||||
):
|
||||
for step in range(2):
|
||||
seq_pos = step
|
||||
ws.position_ids[:b] = seq_pos
|
||||
for tid in task_ids:
|
||||
task_cache.task_extend(tid, seq_pos)
|
||||
kv = task_cache.bind(task_ids, ws)
|
||||
ids_buf = ws.fill_input_ids([step] * b)
|
||||
gctx.forward(
|
||||
model,
|
||||
key=(b,),
|
||||
input_ids=ids_buf,
|
||||
kv_cache=kv,
|
||||
position_ids=ws.position_ids[:b],
|
||||
fwd="decode",
|
||||
)
|
||||
|
||||
for tid in task_ids:
|
||||
task_cache.task_free(tid)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
|
||||
class Executor:
|
||||
"""Model forward passes for prefill and decode phases."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: AutoModel,
|
||||
kv_cache: PagePool,
|
||||
task_cache: TaskCacheManager,
|
||||
device: Optional[str] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
enable_cuda_graph: bool = True,
|
||||
):
|
||||
self.model = model
|
||||
self.kv_cache = kv_cache
|
||||
self.task_cache = task_cache
|
||||
self.device = device or next(model.parameters()).device
|
||||
self.dtype = dtype or next(model.parameters()).dtype
|
||||
|
||||
# Per-step decode cache for the steady-state case (same ordered
|
||||
# task set decodes one token per step). Sampling params stay
|
||||
# constant; only positions advance.
|
||||
self._decode_cache: Optional[DecodeSteadyState] = None
|
||||
|
||||
# Pre-allocated fixed-shape buffers for the decode hot path
|
||||
# (input_ids, decode mask, KV bind metadata). Eagerly sized at init
|
||||
# so the workspace is CUDA-graph-capture friendly — no allocation
|
||||
# during capture.
|
||||
config = model.config
|
||||
max_q_heads = config.num_attention_heads
|
||||
head_dim = config.hidden_size // config.num_attention_heads
|
||||
backend = get_backend()
|
||||
self._graph_supported = backend.supports_graph() and CudaBackend.supports(
|
||||
head_dim=head_dim
|
||||
)
|
||||
self._workspace = InferenceWorkspace(
|
||||
max_batch_size=kv_cache.max_batch_size,
|
||||
max_seq_len=kv_cache.max_seq_len,
|
||||
max_q_heads=max_q_heads,
|
||||
head_dim=head_dim,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
|
||||
# CUDA-graph capture: one graph per (batch_size,) key.
|
||||
# Enabled at init-time via _warmup_cuda_graphs for CudaBackend
|
||||
# on supported head_dims; left disabled otherwise.
|
||||
self._graph_ctx = CudaGraphContext()
|
||||
if enable_cuda_graph:
|
||||
self._try_enable_cuda_graph()
|
||||
|
||||
def _try_enable_cuda_graph(self):
|
||||
if not self._graph_supported:
|
||||
return
|
||||
|
||||
self._graph_ctx.set_enabled(True)
|
||||
_warmup_cuda_graphs(
|
||||
self.model,
|
||||
self.kv_cache,
|
||||
self.task_cache,
|
||||
self._workspace,
|
||||
self._graph_ctx,
|
||||
max_batch_size=self.kv_cache.max_batch_size,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
@property
|
||||
def cuda_graph_enabled(self) -> bool:
|
||||
return self._graph_ctx.enabled and self._graph_supported
|
||||
|
||||
def _sample_logits(
|
||||
self,
|
||||
logits: Tensor,
|
||||
tasks: List[Task],
|
||||
return_logprobs: bool = False,
|
||||
info: Optional[SamplingBatchInfo] = None,
|
||||
):
|
||||
info = info or _build_sampling_batch_info(tasks, self.device)
|
||||
if info.has_freq:
|
||||
history_lists = [
|
||||
t.prompt_ids[-t.rep_window :] + t.output_ids for t in tasks
|
||||
]
|
||||
history_lens = [len(ids) for ids in history_lists]
|
||||
max_len = max(history_lens, default=0)
|
||||
padded_ids = torch.zeros(
|
||||
len(tasks), max_len, dtype=torch.long, device=self.device
|
||||
)
|
||||
padded_mask = torch.zeros(
|
||||
len(tasks), max_len, dtype=torch.bool, device=self.device
|
||||
)
|
||||
for i, ids in enumerate(history_lists):
|
||||
length = len(ids)
|
||||
padded_ids[i, :length] = torch.as_tensor(
|
||||
ids, dtype=torch.long, device=self.device
|
||||
)
|
||||
padded_mask[i, :length] = True
|
||||
else:
|
||||
padded_ids = None
|
||||
padded_mask = None
|
||||
|
||||
result = sample(
|
||||
logits,
|
||||
temperature=info.temperatures,
|
||||
top_k=info.top_ks,
|
||||
top_p=info.top_ps,
|
||||
frequency_penalty=info.freq_penalties,
|
||||
input_ids=padded_ids,
|
||||
input_mask=padded_mask,
|
||||
return_logprobs=return_logprobs,
|
||||
)
|
||||
if not return_logprobs:
|
||||
return result.tolist()
|
||||
|
||||
tokens, logprobs = result
|
||||
tokens_list = tokens.tolist()
|
||||
logprobs_list = logprobs.tolist()
|
||||
for task, logprob in zip(tasks, logprobs_list):
|
||||
task.output_logprobs.append(float(logprob))
|
||||
return list(zip(tokens_list, logprobs_list))
|
||||
|
||||
def execute_prefill(
|
||||
self,
|
||||
tasks: List[Task],
|
||||
prompt_len: int,
|
||||
start_pos: int = 0,
|
||||
return_logprobs: bool = False,
|
||||
):
|
||||
if start_pos >= prompt_len:
|
||||
return []
|
||||
|
||||
tasks = sorted(tasks, key=lambda t: t.task_id)
|
||||
batch_sz = len(tasks)
|
||||
|
||||
input_ids = torch.tensor(
|
||||
[token for t in tasks for token in t.prompt_ids[start_pos:prompt_len]],
|
||||
dtype=torch.long,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
task_ids = [t.task_id for t in tasks]
|
||||
position_ids = torch.arange(
|
||||
start_pos, prompt_len, dtype=torch.long, device=self.device
|
||||
).repeat(batch_sz)
|
||||
|
||||
with (
|
||||
torch.inference_mode(),
|
||||
timed(f"execute_prefill b={batch_sz} prompt_len={prompt_len}", logger),
|
||||
):
|
||||
outputs = self.model(
|
||||
input_ids,
|
||||
position_ids=position_ids,
|
||||
kv_cache=self.task_cache.bind(
|
||||
task_ids,
|
||||
self._workspace,
|
||||
start_pos=start_pos,
|
||||
),
|
||||
fwd="prefill",
|
||||
)
|
||||
q_len = prompt_len - start_pos
|
||||
logits = outputs["logits"][
|
||||
torch.arange(1, batch_sz + 1, device=self.device) * q_len - 1
|
||||
]
|
||||
|
||||
return tasks, self._sample_logits(logits, tasks, return_logprobs)
|
||||
|
||||
def execute_decode(
|
||||
self, tasks: List[Task], return_logprobs: bool = False
|
||||
) -> List[int]:
|
||||
"""Decode next token for each task.
|
||||
|
||||
Args:
|
||||
return_logprobs: When ``True``, also record (and return)
|
||||
the log-probability of each sampled token under the
|
||||
post-strategy sampling distribution. The logprob is
|
||||
appended to ``task.output_logprobs`` and the return
|
||||
list becomes ``List[Tuple[int, float]]``.
|
||||
|
||||
Returns:
|
||||
``List[int]`` of sampled token IDs, or
|
||||
``List[Tuple[int, float]]`` of ``(token_id, logprob)`` when
|
||||
``return_logprobs`` is ``True``.
|
||||
"""
|
||||
if not tasks:
|
||||
return []
|
||||
|
||||
b = len(tasks)
|
||||
ws = self._workspace
|
||||
|
||||
# ---- pre-replay: update input buffers in-place ----
|
||||
|
||||
input_ids = ws.fill_input_ids(
|
||||
[t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks]
|
||||
)
|
||||
|
||||
task_ids = [t.task_id for t in tasks]
|
||||
cur_positions = [t.next_pos for t in tasks]
|
||||
|
||||
kv_cache = self.task_cache.bind(task_ids, ws)
|
||||
|
||||
task_sig = tuple(task_ids)
|
||||
reuse_decode_state = (
|
||||
self.task_cache.bind_was_steady
|
||||
and self._decode_cache is not None
|
||||
and self._decode_cache.task_sig == task_sig
|
||||
)
|
||||
if reuse_decode_state:
|
||||
info = self._decode_cache.sampling_info
|
||||
ws.position_ids[:b] += 1
|
||||
else:
|
||||
info = _build_sampling_batch_info(tasks, self.device)
|
||||
ws.position_ids[:b].copy_(
|
||||
torch.tensor(cur_positions, dtype=torch.long, device=self.device)
|
||||
)
|
||||
self._decode_cache = DecodeSteadyState(task_sig, cur_positions, info)
|
||||
|
||||
# ---- forward (graph replay or live run + capture) ----
|
||||
|
||||
use_graph = (
|
||||
self._graph_ctx.enabled
|
||||
and self._graph_supported
|
||||
and get_backend().supports_graph()
|
||||
)
|
||||
key = (b,)
|
||||
with (
|
||||
torch.inference_mode(),
|
||||
timed(f"execute_decode forward b={b}", logger),
|
||||
):
|
||||
if use_graph:
|
||||
outputs = self._graph_ctx.forward(
|
||||
self.model,
|
||||
key=key,
|
||||
input_ids=input_ids,
|
||||
kv_cache=kv_cache,
|
||||
position_ids=ws.position_ids[:b],
|
||||
fwd="decode",
|
||||
)
|
||||
else:
|
||||
outputs = self.model(
|
||||
input_ids,
|
||||
kv_cache=kv_cache,
|
||||
position_ids=ws.position_ids[:b],
|
||||
fwd="decode",
|
||||
)
|
||||
logits = outputs["logits"]
|
||||
|
||||
return self._sample_logits(logits, tasks, return_logprobs, info=info)
|
||||
@@ -0,0 +1,103 @@
|
||||
"""CUDA-graph capture for the decode model-forward step.
|
||||
|
||||
Mirrors SGLang's cuda-graph manager: one graph per batch size. The graph
|
||||
pair. The graph captures ``model.forward()`` with workspace-backed inputs
|
||||
(all at fixed addresses). Before each replay the caller updates the input
|
||||
buffer content in-place so the graph sees fresh data at the same tensor
|
||||
addresses.
|
||||
|
||||
Only the model forward is captured — sampling runs outside the graph
|
||||
(via ``torch.multinomial`` which consumes a mutable RNG state).
|
||||
"""
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
class CudaGraphContext:
|
||||
"""CUDA-graph capture/replay for decode steps.
|
||||
|
||||
Parameters:
|
||||
enabled: When ``False``, ``forward()`` always runs the live model
|
||||
forward without capture/replay (graphs are cleared). Toggle at
|
||||
runtime via the ``set_enabled()`` method.
|
||||
|
||||
Usage::
|
||||
|
||||
gctx = CudaGraphContext()
|
||||
with torch.inference_mode():
|
||||
outputs = gctx.forward(
|
||||
model,
|
||||
key=(batch_size,),
|
||||
input_ids=workspace.input_ids[:b].unsqueeze(1),
|
||||
input_mask=input_mask,
|
||||
kv_cache=kv_cache,
|
||||
position_ids=workspace.position_ids[:b].unsqueeze(1),
|
||||
)
|
||||
|
||||
The first call at a given key runs *without* capture (warmup). The
|
||||
second call captures the graph. Subsequent calls replay the captured
|
||||
graph. A ``torch.cuda.synchronize()`` before capture drains in-flight
|
||||
work so the graph trace is clean.
|
||||
"""
|
||||
|
||||
def __init__(self, enabled: bool = False):
|
||||
self._enabled = enabled
|
||||
self._graphs: dict[tuple, torch.cuda.CUDAGraph] = {}
|
||||
self._outputs: dict[tuple, dict[str, Tensor]] = {}
|
||||
self._warmed: set[tuple] = set()
|
||||
|
||||
@property
|
||||
def enabled(self) -> bool:
|
||||
return self._enabled
|
||||
|
||||
def set_enabled(self, flag: bool):
|
||||
"""Enable or disable CUDA-graph capture at runtime.
|
||||
|
||||
Disabling clears all captured graphs (frees GPU memory) and warmup
|
||||
state. Re-enabling after disable starts fresh — graphs are
|
||||
re-captured on the next warmup cycle.
|
||||
"""
|
||||
if flag == self._enabled:
|
||||
return
|
||||
self._enabled = flag
|
||||
if not flag:
|
||||
self._graphs.clear()
|
||||
self._outputs.clear()
|
||||
self._warmed.clear()
|
||||
|
||||
def forward(self, model, *, key, **kwargs) -> dict[str, Tensor]:
|
||||
"""Run ``model(**kwargs)`` via graph replay or live forward.
|
||||
|
||||
Args:
|
||||
model: callable, e.g. ``self.model.forward``.
|
||||
key: ``(batch_size,)`` — the dispatch key (one graph per batch size).
|
||||
**kwargs: arguments forwarded to ``model``. All tensor arguments
|
||||
must reside at stable addresses (workspace buffers).
|
||||
|
||||
Returns:
|
||||
The dict produced by ``model(**kwargs)``, e.g.
|
||||
``{"logits": ..., "h0": ...}``.
|
||||
"""
|
||||
if not self._enabled:
|
||||
self._outputs[key] = model(**kwargs)
|
||||
return self._outputs[key]
|
||||
|
||||
if key in self._graphs:
|
||||
self._graphs[key].replay()
|
||||
elif key in self._warmed:
|
||||
cap_output = model(**kwargs)
|
||||
torch.cuda.synchronize()
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
self._outputs[key] = model(**kwargs)
|
||||
self._graphs[key] = graph
|
||||
self._warmed.discard(key)
|
||||
return cap_output
|
||||
else:
|
||||
self._warmed.add(key)
|
||||
self._outputs[key] = model(**kwargs)
|
||||
return self._outputs[key]
|
||||
|
||||
def has_graph(self, key: tuple) -> bool:
|
||||
return key in self._graphs
|
||||
@@ -0,0 +1,386 @@
|
||||
"""Composable sampling strategies for logit transformation.
|
||||
|
||||
Implements the Strategy pattern: each sampling technique
|
||||
(temperature, top-k, top-p, frequency penalty) is a pluggable
|
||||
strategy that can be composed into a pipeline.
|
||||
|
||||
All strategies accept both scalar and per-sample tensor
|
||||
parameters, so a single pipeline works for any batch size.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
class BaseSamplingStrategy(ABC):
|
||||
"""Abstract base for a logit transformation strategy."""
|
||||
|
||||
@abstractmethod
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
"""Applies the strategy to logits.
|
||||
|
||||
Args:
|
||||
logits: Raw logits tensor (batch, vocab_size).
|
||||
filter_value: Value assigned to filtered-out positions.
|
||||
input_ids: Previously generated token IDs ``[batch, seq_len]``,
|
||||
padded with 0. Used by frequency penalty.
|
||||
input_mask: Boolean mask ``[batch, seq_len]``, True for real
|
||||
tokens, False for padding. Used to exclude padding from
|
||||
penalty computation.
|
||||
|
||||
Returns:
|
||||
Transformed logits tensor.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class TemperatureStrategy(BaseSamplingStrategy):
|
||||
"""Divides logits by temperature to control randomness.
|
||||
|
||||
Args:
|
||||
temperature: Scalar or ``[batch]`` tensor.
|
||||
"""
|
||||
|
||||
def __init__(self, temperature: Union[float, Tensor] = 1.0):
|
||||
self.temperature = temperature
|
||||
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
t = self.temperature
|
||||
if isinstance(t, Tensor):
|
||||
t = t.to(logits.device, non_blocking=True).view(-1, 1)
|
||||
t = torch.clamp(t, min=1e-8)
|
||||
if (t != 1.0).any():
|
||||
logits = logits / t
|
||||
elif t != 1.0:
|
||||
logits = logits / max(t, 1e-8)
|
||||
return logits
|
||||
|
||||
|
||||
class TopKStrategy(BaseSamplingStrategy):
|
||||
"""Keeps only the top-k logits, setting the rest to filter_value.
|
||||
|
||||
Args:
|
||||
top_k: Scalar or ``[batch]`` tensor (0 disables).
|
||||
"""
|
||||
|
||||
def __init__(self, top_k: Union[int, Tensor] = 0):
|
||||
self.top_k = top_k
|
||||
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
tk = self.top_k
|
||||
if isinstance(tk, Tensor):
|
||||
tk = tk.to(logits.device, non_blocking=True).long().clamp(min=0)
|
||||
max_k = int(tk.max().item())
|
||||
if max_k <= 0:
|
||||
return logits
|
||||
max_k = min(max_k, logits.size(-1))
|
||||
values, _ = torch.topk(logits, max_k, dim=-1)
|
||||
per_row_k = tk.clamp(max=max_k)
|
||||
thresholds = torch.full_like(logits[..., -1:], -float("inf"))
|
||||
positive = per_row_k > 0
|
||||
if positive.any():
|
||||
row_idx = torch.arange(logits.size(0), device=logits.device)[positive]
|
||||
thresholds[positive] = values[
|
||||
row_idx, per_row_k[positive] - 1
|
||||
].unsqueeze(-1)
|
||||
logits[logits < thresholds] = filter_value
|
||||
return logits
|
||||
if tk > 0:
|
||||
k = min(tk, logits.size(-1))
|
||||
thresholds = torch.topk(logits, k, dim=-1)[0][..., -1:]
|
||||
logits[logits < thresholds] = filter_value
|
||||
return logits
|
||||
|
||||
|
||||
class TopPStrategy(BaseSamplingStrategy):
|
||||
"""Nucleus (top-p) filtering: keeps the smallest set of tokens whose
|
||||
cumulative probability exceeds top_p.
|
||||
|
||||
Args:
|
||||
top_p: Scalar or ``[batch]`` tensor (1.0 disables).
|
||||
"""
|
||||
|
||||
def __init__(self, top_p: Union[float, Tensor] = 1.0):
|
||||
self.top_p = top_p
|
||||
|
||||
def _apply(
|
||||
self, logits: Tensor, top_p: Union[float, Tensor], filter_value: float
|
||||
) -> Tensor:
|
||||
sorted_logits, sorted_indices = torch.sort(logits, descending=True, dim=-1)
|
||||
cum_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
|
||||
remove = cum_probs > top_p
|
||||
remove[..., 1:] = remove[..., :-1].clone()
|
||||
remove[..., 0] = False
|
||||
mask = torch.zeros_like(logits, dtype=torch.bool)
|
||||
mask.scatter_(1, sorted_indices, remove)
|
||||
logits[mask] = filter_value
|
||||
return logits
|
||||
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
tp = self.top_p
|
||||
if isinstance(tp, Tensor):
|
||||
tp = tp.to(logits.device, non_blocking=True)
|
||||
if (tp < 1.0).any():
|
||||
logits = self._apply(logits, tp.view(-1, 1), filter_value)
|
||||
elif tp < 1.0:
|
||||
logits = self._apply(logits, tp, filter_value)
|
||||
return logits
|
||||
|
||||
|
||||
class FrequencyPenaltyStrategy(BaseSamplingStrategy):
|
||||
"""Penalizes tokens based on how many times they appeared in history.
|
||||
|
||||
Subtracts ``penalty * count(token)`` from each token's logit, where
|
||||
``count(token)`` is the number of occurrences in the generation history
|
||||
(prompt + output). A penalty of ``0.0`` disables the strategy.
|
||||
|
||||
Unlike repetition penalty (which only checks *presence*), frequency
|
||||
penalty scales linearly with occurrence count: the first use is
|
||||
penalized once, the third use three times. This allows natural
|
||||
repetition of common words while suppressing degenerate loops.
|
||||
|
||||
Reference: OpenAI API ``frequency_penalty`` parameter.
|
||||
|
||||
Args:
|
||||
penalty: Scalar or ``[batch]`` tensor (0.0 disables, range -2.0~2.0).
|
||||
"""
|
||||
|
||||
def __init__(self, penalty: Union[float, Tensor] = 0.0):
|
||||
self.penalty = penalty
|
||||
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
if input_ids is None:
|
||||
return logits
|
||||
|
||||
p = self.penalty
|
||||
if isinstance(p, Tensor):
|
||||
p = p.to(logits.device, non_blocking=True).view(-1, 1)
|
||||
if (p == 0.0).all():
|
||||
return logits
|
||||
elif p == 0.0:
|
||||
return logits
|
||||
|
||||
input_ids = input_ids.to(logits.device, non_blocking=True)
|
||||
|
||||
if input_mask is not None:
|
||||
input_mask = input_mask.to(logits.device, non_blocking=True)
|
||||
masked_ids = input_ids.clone()
|
||||
masked_ids[~input_mask] = -1
|
||||
else:
|
||||
masked_ids = input_ids
|
||||
|
||||
batch_sz, seq_len = masked_ids.shape
|
||||
vocab_size = logits.size(-1)
|
||||
|
||||
if isinstance(p, Tensor):
|
||||
penalty_per_row = p.expand(batch_sz, 1)
|
||||
else:
|
||||
penalty_per_row = torch.full(
|
||||
(batch_sz, 1), float(p), device=logits.device, dtype=logits.dtype
|
||||
)
|
||||
|
||||
counts = torch.zeros(
|
||||
batch_sz, vocab_size, device=logits.device, dtype=logits.dtype
|
||||
)
|
||||
valid_mask = masked_ids >= 0
|
||||
if valid_mask.any():
|
||||
valid_ids = masked_ids[valid_mask]
|
||||
row_indices = (
|
||||
torch.arange(batch_sz, device=logits.device)
|
||||
.unsqueeze(1)
|
||||
.expand_as(masked_ids)[valid_mask]
|
||||
)
|
||||
counts.index_put_(
|
||||
(row_indices, valid_ids),
|
||||
torch.ones_like(valid_ids, dtype=logits.dtype),
|
||||
accumulate=True,
|
||||
)
|
||||
|
||||
return logits - penalty_per_row * counts
|
||||
|
||||
|
||||
class SamplingPipeline(BaseSamplingStrategy):
|
||||
"""Composes multiple sampling strategies into a single transformation.
|
||||
|
||||
Strategies are applied sequentially in the order they are provided,
|
||||
matching the original temperature -> top-k -> top-p ordering.
|
||||
|
||||
Usage::
|
||||
|
||||
pipeline = SamplingPipeline([
|
||||
TemperatureStrategy(0.8),
|
||||
TopKStrategy(50),
|
||||
TopPStrategy(0.95),
|
||||
])
|
||||
logits = pipeline.apply(logits)
|
||||
token = pipeline.sample(logits) # softmax + multinomial
|
||||
"""
|
||||
|
||||
def __init__(self, strategies: List[BaseSamplingStrategy]):
|
||||
self.strategies = strategies
|
||||
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
for strategy in self.strategies:
|
||||
logits = strategy.apply(logits, filter_value, input_ids, input_mask)
|
||||
return logits
|
||||
|
||||
@staticmethod
|
||||
def _is_greedy(temperature: Union[float, Tensor]) -> bool:
|
||||
if isinstance(temperature, Tensor):
|
||||
return bool((temperature == 0).all())
|
||||
return temperature == 0
|
||||
|
||||
@torch.inference_mode()
|
||||
def sample(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
return_logprobs: bool = False,
|
||||
):
|
||||
"""Apply strategies then sample (softmax + multinomial).
|
||||
|
||||
Short-circuits to ``argmax`` when temperature is exactly 0
|
||||
(deterministic / greedy decode).
|
||||
|
||||
Args:
|
||||
logits: Raw logits ``[batch, vocab_size]``.
|
||||
input_ids: Previously generated token IDs ``[batch, seq_len]``.
|
||||
input_mask: Boolean mask for ``input_ids`` padding.
|
||||
return_logprobs: If ``True``, return ``(tokens, logprobs)``
|
||||
where ``logprobs[i]`` is the log-probability of
|
||||
``tokens[i]`` under the (post-strategy) sampling
|
||||
distribution.
|
||||
|
||||
Returns:
|
||||
Sampled token IDs ``[batch]``, or — when ``return_logprobs``
|
||||
is ``True`` — a ``(token_ids, chosen_logprobs)`` tuple.
|
||||
"""
|
||||
if self._is_greedy_pipeline():
|
||||
tokens = logits.argmax(dim=-1)
|
||||
if not return_logprobs:
|
||||
return tokens
|
||||
log_probs = torch.log_softmax(logits.float(), dim=-1)
|
||||
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
|
||||
return tokens, chosen
|
||||
|
||||
transformed = self.apply(logits, filter_value, input_ids, input_mask)
|
||||
tokens = torch.multinomial(
|
||||
torch.softmax(transformed, dim=-1), num_samples=1
|
||||
).squeeze(-1)
|
||||
if not return_logprobs:
|
||||
return tokens
|
||||
log_probs = torch.log_softmax(transformed.float(), dim=-1)
|
||||
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
|
||||
return tokens, chosen
|
||||
|
||||
def _is_greedy_pipeline(self) -> bool:
|
||||
"""True if the first strategy is greedy temperature (temp=0)."""
|
||||
if not self.strategies:
|
||||
return False
|
||||
first = self.strategies[0]
|
||||
return isinstance(first, TemperatureStrategy) and self._is_greedy(
|
||||
first.temperature
|
||||
)
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def sample(
|
||||
logits: Tensor,
|
||||
temperature: Union[float, Tensor] = 1.0,
|
||||
top_k: Union[int, Tensor] = 0,
|
||||
top_p: Union[float, Tensor] = 1.0,
|
||||
frequency_penalty: Union[float, Tensor] = 0.0,
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
filter_value: float = -float("inf"),
|
||||
return_logprobs: bool = False,
|
||||
):
|
||||
"""Apply sampling strategies then sample (softmax + multinomial).
|
||||
|
||||
Shortcut for ``SamplingPipeline(...).sample(logits, return_logprobs=)``.
|
||||
|
||||
When **temperature** is exactly 0 (scalar or single-element tensor)
|
||||
the function short-circuits to ``argmax`` for deterministic decode.
|
||||
|
||||
When **frequency_penalty** is 0 (the common decode case), the entire
|
||||
frequency penalty computation — including the O(batch * vocab) count
|
||||
tensor allocation — is skipped.
|
||||
|
||||
Args:
|
||||
logits: Raw logits ``[batch, vocab_size]``.
|
||||
frequency_penalty: Penalty per occurrence for repeated tokens
|
||||
(0.0 disables, range -2.0~2.0).
|
||||
input_ids: Previously generated token IDs ``[batch, seq_len]``.
|
||||
input_mask: Boolean mask for ``input_ids`` padding.
|
||||
return_logprobs: If ``True``, also return the log-probability
|
||||
of each sampled token under the (post-strategy) sampling
|
||||
distribution — useful for RL rollout (PPO/GRPO importance
|
||||
ratios).
|
||||
|
||||
Returns:
|
||||
Sampled token IDs ``[batch]``, or — when ``return_logprobs`` is
|
||||
``True`` — a ``(token_ids, chosen_logprobs)`` tuple where
|
||||
``chosen_logprobs`` has shape ``[batch]``.
|
||||
"""
|
||||
has_freq = (
|
||||
(isinstance(frequency_penalty, Tensor) and (frequency_penalty != 0).any())
|
||||
if isinstance(frequency_penalty, Tensor)
|
||||
else frequency_penalty != 0
|
||||
)
|
||||
|
||||
strategies: List[BaseSamplingStrategy] = [
|
||||
TemperatureStrategy(temperature),
|
||||
TopKStrategy(top_k),
|
||||
TopPStrategy(top_p),
|
||||
]
|
||||
if has_freq:
|
||||
strategies.append(FrequencyPenaltyStrategy(frequency_penalty))
|
||||
|
||||
return SamplingPipeline(strategies).sample(
|
||||
logits,
|
||||
filter_value=filter_value,
|
||||
input_ids=input_ids,
|
||||
input_mask=input_mask,
|
||||
return_logprobs=return_logprobs,
|
||||
)
|
||||
@@ -1,178 +0,0 @@
|
||||
"""Composable sampling strategies for logit transformation.
|
||||
|
||||
Implements the Strategy pattern: each sampling technique
|
||||
(temperature, top-k, top-p) is a pluggable strategy that
|
||||
can be composed into a pipeline.
|
||||
|
||||
All strategies accept both scalar and per-sample tensor
|
||||
parameters, so a single pipeline works for any batch size.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List, Union
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
class BaseSamplingStrategy(ABC):
|
||||
"""Abstract base for a logit transformation strategy."""
|
||||
|
||||
@abstractmethod
|
||||
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
||||
"""Applies the strategy to logits.
|
||||
|
||||
Args:
|
||||
logits: Raw logits tensor (batch, vocab_size).
|
||||
filter_value: Value assigned to filtered-out positions.
|
||||
|
||||
Returns:
|
||||
Transformed logits tensor.
|
||||
"""
|
||||
|
||||
|
||||
class TemperatureStrategy(BaseSamplingStrategy):
|
||||
"""Divides logits by temperature to control randomness.
|
||||
|
||||
Args:
|
||||
temperature: Scalar or ``[batch]`` tensor.
|
||||
"""
|
||||
|
||||
def __init__(self, temperature: Union[float, Tensor] = 1.0):
|
||||
self.temperature = temperature
|
||||
|
||||
def apply(self, logits, filter_value=-float("inf")):
|
||||
t = self.temperature
|
||||
if isinstance(t, Tensor):
|
||||
if (t != 1.0).any():
|
||||
logits = logits / t.to(logits.device, non_blocking=True).view(-1, 1)
|
||||
elif t != 1.0:
|
||||
logits = logits / t
|
||||
return logits
|
||||
|
||||
|
||||
class TopKStrategy(BaseSamplingStrategy):
|
||||
"""Keeps only the top-k logits, setting the rest to filter_value.
|
||||
|
||||
Args:
|
||||
top_k: Scalar or ``[batch]`` tensor (0 disables).
|
||||
"""
|
||||
|
||||
def __init__(self, top_k: Union[int, Tensor] = 0):
|
||||
self.top_k = top_k
|
||||
|
||||
def apply(self, logits, filter_value=-float("inf")):
|
||||
tk = self.top_k
|
||||
if isinstance(tk, Tensor):
|
||||
max_k = int(tk.max().item())
|
||||
if max_k <= 0:
|
||||
return logits
|
||||
k = min(max_k, logits.size(-1))
|
||||
elif tk > 0:
|
||||
k = min(tk, logits.size(-1))
|
||||
else:
|
||||
return logits
|
||||
thresholds = torch.topk(logits, k, dim=-1)[0][..., -1:]
|
||||
logits[logits < thresholds] = filter_value
|
||||
return logits
|
||||
|
||||
|
||||
class TopPStrategy(BaseSamplingStrategy):
|
||||
"""Nucleus (top-p) filtering: keeps the smallest set of tokens whose
|
||||
cumulative probability exceeds top_p.
|
||||
|
||||
Args:
|
||||
top_p: Scalar or ``[batch]`` tensor (1.0 disables).
|
||||
"""
|
||||
|
||||
def __init__(self, top_p: Union[float, Tensor] = 1.0):
|
||||
self.top_p = top_p
|
||||
|
||||
def _apply(self, logits, top_p, filter_value):
|
||||
sorted_logits, sorted_indices = torch.sort(logits, descending=True, dim=-1)
|
||||
cum_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
|
||||
remove = cum_probs > top_p
|
||||
remove[..., 1:] = remove[..., :-1].clone()
|
||||
remove[..., 0] = False
|
||||
mask = torch.zeros_like(logits, dtype=torch.bool)
|
||||
mask.scatter_(1, sorted_indices, remove)
|
||||
logits[mask] = filter_value
|
||||
return logits
|
||||
|
||||
def apply(self, logits, filter_value=-float("inf")):
|
||||
tp = self.top_p
|
||||
if isinstance(tp, Tensor):
|
||||
tp = tp.to(logits.device, non_blocking=True)
|
||||
if (tp < 1.0).any():
|
||||
logits = self._apply(logits, tp.view(-1, 1), filter_value)
|
||||
elif tp < 1.0:
|
||||
logits = self._apply(logits, tp, filter_value)
|
||||
return logits
|
||||
|
||||
|
||||
class SamplingPipeline(BaseSamplingStrategy):
|
||||
"""Composes multiple sampling strategies into a single transformation.
|
||||
|
||||
Strategies are applied sequentially in the order they are provided,
|
||||
matching the original temperature -> top-k -> top-p ordering.
|
||||
|
||||
Usage::
|
||||
|
||||
pipeline = SamplingPipeline([
|
||||
TemperatureStrategy(0.8),
|
||||
TopKStrategy(50),
|
||||
TopPStrategy(0.95),
|
||||
])
|
||||
logits = pipeline.apply(logits)
|
||||
token = pipeline.sample(logits) # softmax + multinomial
|
||||
"""
|
||||
|
||||
def __init__(self, strategies: List[BaseSamplingStrategy]):
|
||||
self.strategies = strategies
|
||||
|
||||
def apply(self, logits, filter_value=-float("inf")):
|
||||
for strategy in self.strategies:
|
||||
logits = strategy.apply(logits, filter_value)
|
||||
return logits
|
||||
|
||||
@torch.no_grad()
|
||||
def sample(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
||||
"""Apply strategies then sample (softmax + multinomial).
|
||||
|
||||
Args:
|
||||
logits: Raw logits ``[batch, vocab_size]``.
|
||||
|
||||
Returns:
|
||||
Sampled token IDs ``[batch]``.
|
||||
"""
|
||||
return torch.multinomial(
|
||||
torch.softmax(self.apply(logits, filter_value), dim=-1),
|
||||
num_samples=1,
|
||||
).squeeze(-1)
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def sample(
|
||||
logits: Tensor,
|
||||
temperature: Union[float, Tensor] = 1.0,
|
||||
top_k: Union[int, Tensor] = 0,
|
||||
top_p: Union[float, Tensor] = 1.0,
|
||||
filter_value: float = -float("inf"),
|
||||
) -> Tensor:
|
||||
"""Apply sampling strategies then sample (softmax + multinomial).
|
||||
|
||||
Shortcut for ``SamplingPipeline(...).sample(logits)``.
|
||||
|
||||
Args:
|
||||
logits: Raw logits ``[batch, vocab_size]``.
|
||||
|
||||
Returns:
|
||||
Sampled token IDs ``[batch]``.
|
||||
"""
|
||||
return SamplingPipeline(
|
||||
[
|
||||
TemperatureStrategy(temperature),
|
||||
TopKStrategy(top_k),
|
||||
TopPStrategy(top_p),
|
||||
]
|
||||
).sample(logits, filter_value)
|
||||
+339
-342
@@ -1,85 +1,29 @@
|
||||
"""Inference scheduler for single-GPU continuous batching with paged KV cache."""
|
||||
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from enum import Enum
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple
|
||||
from contextlib import nullcontext
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.inference.cache import STOP, PagedCache
|
||||
from astrai.inference.sampling import sample
|
||||
from astrai.extension import (
|
||||
ATTN_BACKEND,
|
||||
AttentionBackend,
|
||||
attn_backend,
|
||||
get_backend,
|
||||
)
|
||||
from astrai.inference.cache import PagePool, TaskCacheManager
|
||||
from astrai.inference.metrics import MetricsCollector
|
||||
from astrai.inference.runtime.executor import Executor
|
||||
from astrai.inference.task import STOP, Task, TaskManager, TaskStatus
|
||||
from astrai.model.automodel import AutoModel
|
||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class TaskStatus(Enum):
|
||||
"""Task states in the continuous batching lifecycle."""
|
||||
|
||||
PENDING = "pending"
|
||||
RUNNING = "running"
|
||||
FINISHED = "finished"
|
||||
ABORTED = "aborted"
|
||||
|
||||
|
||||
class Task:
|
||||
"""Represents a single generation request with paged KV cache tracking."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
task_id: str,
|
||||
prompt_ids: List[int],
|
||||
max_tokens: int = 1024,
|
||||
temperature: float = 1.0,
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
stream_callback: Optional[Callable[[str], None]] = None,
|
||||
):
|
||||
self.task_id = task_id
|
||||
self.prompt_ids = prompt_ids
|
||||
self.max_tokens = max_tokens
|
||||
self.temperature = temperature
|
||||
self.top_p = top_p
|
||||
self.top_k = top_k
|
||||
|
||||
self.status = TaskStatus.PENDING
|
||||
self.output_ids: List[int] = []
|
||||
self.input_tokens: int = 0
|
||||
self.output_tokens: int = 0
|
||||
self.page_table: List[int] = []
|
||||
self.n_pages: int = 0
|
||||
self._prefix_cached_tokens: int = 0
|
||||
self.arrival_time = time.time()
|
||||
self.finish_time: Optional[float] = None
|
||||
self.stream_callback = stream_callback
|
||||
self._pages_freed: bool = False
|
||||
|
||||
@property
|
||||
def next_pos(self) -> int:
|
||||
return self.input_tokens + len(self.output_ids)
|
||||
|
||||
def is_finished(self, stop_ids: List[int]) -> bool:
|
||||
if self.output_tokens >= self.max_tokens:
|
||||
return True
|
||||
if self.output_ids and self.output_ids[-1] in stop_ids:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
class InferenceScheduler:
|
||||
"""Continuous batching scheduler with paged KV cache.
|
||||
|
||||
Runs a background generation loop with four phases per iteration:
|
||||
1. Cleanup finished tasks and release resources.
|
||||
2. Refill active batch from the waiting queue.
|
||||
3. Prefill newly activated tasks.
|
||||
4. Decode the largest same-position group of active tasks.
|
||||
"""
|
||||
"""Continuous batching loop: cleanup -> refill -> prefill -> decode (all groups)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -87,325 +31,378 @@ class InferenceScheduler:
|
||||
tokenizer: AutoTokenizer,
|
||||
max_batch_size: int = 16,
|
||||
max_seq_len: Optional[int] = None,
|
||||
max_prompt_len: int = 512,
|
||||
page_size: int = 64,
|
||||
device: Optional[str] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
cache: Optional[PagePool] = None,
|
||||
enable_cuda_graph: bool = True,
|
||||
backend: Optional[Union[str, ATTN_BACKEND, AttentionBackend, type]] = None,
|
||||
):
|
||||
config = model.config
|
||||
|
||||
self.model = model
|
||||
self.tokenizer = tokenizer
|
||||
self.max_batch_size = max_batch_size
|
||||
self.max_seq_len = max_seq_len or config.max_len
|
||||
self.max_prompt_len = max_prompt_len
|
||||
self.page_size = page_size
|
||||
if max_seq_len is not None:
|
||||
self.max_seq_len = max_seq_len
|
||||
elif config.max_position_embeddings is not None:
|
||||
self.max_seq_len = config.max_position_embeddings
|
||||
else:
|
||||
raise ValueError(
|
||||
"max_seq_len must be provided either as argument "
|
||||
"or in model config (config.max_position_embeddings)"
|
||||
)
|
||||
self.device = device or next(model.parameters()).device
|
||||
self.dtype = dtype or next(model.parameters()).dtype
|
||||
|
||||
n_kv_heads = config.n_kv_heads
|
||||
head_dim = config.dim // config.n_heads
|
||||
n_layers = config.n_layers
|
||||
n_pages = (
|
||||
max_batch_size * (self.max_seq_len + page_size) + page_size - 1
|
||||
) // page_size
|
||||
head_dim = config.hidden_size // config.num_attention_heads
|
||||
|
||||
self.page_cache = PagedCache(
|
||||
n_layers,
|
||||
n_pages,
|
||||
page_size,
|
||||
n_kv_heads,
|
||||
head_dim,
|
||||
self.device,
|
||||
self.dtype,
|
||||
if cache is not None:
|
||||
self._cache = cache
|
||||
else:
|
||||
self._cache = PagePool(
|
||||
n_layers=config.num_hidden_layers,
|
||||
n_kv_heads=config.num_key_value_heads,
|
||||
head_dim=head_dim,
|
||||
max_batch_size=max_batch_size,
|
||||
max_seq_len=self.max_seq_len,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
|
||||
self._metrics = MetricsCollector()
|
||||
|
||||
self._task_cache = TaskCacheManager(self._cache)
|
||||
|
||||
self._task_mgr = TaskManager(
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=max_batch_size,
|
||||
max_seq_len=self.max_seq_len,
|
||||
metrics=self._metrics,
|
||||
)
|
||||
|
||||
self.waiting_queue: List[Task] = []
|
||||
self.active_tasks: List[Task] = []
|
||||
if backend is None:
|
||||
self._backend = None
|
||||
default_backend = get_backend()
|
||||
self._backend_name = type(default_backend).__name__
|
||||
with attn_backend(default_backend):
|
||||
self._executor = Executor(
|
||||
model=model,
|
||||
kv_cache=self._cache,
|
||||
task_cache=self._task_cache,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
enable_cuda_graph=enable_cuda_graph,
|
||||
)
|
||||
else:
|
||||
with attn_backend(backend):
|
||||
self._backend = get_backend()
|
||||
self._backend_name = type(self._backend).__name__
|
||||
self._executor = Executor(
|
||||
model=model,
|
||||
kv_cache=self._cache,
|
||||
task_cache=self._task_cache,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
enable_cuda_graph=enable_cuda_graph,
|
||||
)
|
||||
|
||||
self._running = False
|
||||
self._task_event = threading.Event()
|
||||
self._lock = threading.Lock()
|
||||
self._stop_event = threading.Event()
|
||||
self._loop_thread: Optional[threading.Thread] = None
|
||||
|
||||
self._total_tasks = 0
|
||||
self._total_tokens = 0
|
||||
def add_task(self, prompt: str, **kwargs) -> str:
|
||||
return self._task_mgr.add_task(prompt, **kwargs)
|
||||
|
||||
def _n_pages_for(self, n_tokens: int) -> int:
|
||||
return (n_tokens + self.page_size - 1) // self.page_size
|
||||
def remove_task(self, task_id: str):
|
||||
for task in self._task_mgr.remove_task(task_id):
|
||||
self._task_cache.task_free(task.task_id)
|
||||
|
||||
def add_task(
|
||||
self,
|
||||
prompt: str,
|
||||
max_tokens: int = 1024,
|
||||
temperature: float = 1.0,
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
stream_callback: Optional[Callable[[str], None]] = None,
|
||||
) -> str:
|
||||
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
|
||||
prompt_ids = self.tokenizer.encode(prompt)
|
||||
if len(prompt_ids) > self.max_prompt_len:
|
||||
prompt_ids = prompt_ids[-self.max_prompt_len :]
|
||||
def get_stats(self) -> Dict[str, Any]:
|
||||
return self._task_mgr.get_stats()
|
||||
|
||||
task = Task(
|
||||
task_id=task_id,
|
||||
prompt_ids=prompt_ids,
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
stream_callback=stream_callback,
|
||||
)
|
||||
@property
|
||||
def backend_name(self) -> str:
|
||||
return self._backend_name
|
||||
|
||||
with self._lock:
|
||||
self.waiting_queue.append(task)
|
||||
self._total_tasks += 1
|
||||
@property
|
||||
def cuda_graph_enabled(self) -> bool:
|
||||
return self._executor.cuda_graph_enabled
|
||||
|
||||
self._task_event.set()
|
||||
return task_id
|
||||
def _backend_context(self):
|
||||
if self._backend is None:
|
||||
return nullcontext()
|
||||
return attn_backend(self._backend)
|
||||
|
||||
def remove_task(self, task_id: str) -> None:
|
||||
with self._lock:
|
||||
removed_active = [t for t in self.active_tasks if t.task_id == task_id]
|
||||
self.waiting_queue = [t for t in self.waiting_queue if t.task_id != task_id]
|
||||
self.active_tasks = [t for t in self.active_tasks if t.task_id != task_id]
|
||||
@staticmethod
|
||||
def _task_backend_groups(tasks: List[Task]):
|
||||
groups = {}
|
||||
for task in tasks:
|
||||
groups.setdefault(task.backend, (task.backend, []))[1].append(task)
|
||||
return groups.values()
|
||||
|
||||
for task in removed_active:
|
||||
if not task._pages_freed:
|
||||
self._free_pages(task.page_table)
|
||||
task.page_table.clear()
|
||||
task.n_pages = 0
|
||||
task._pages_freed = True
|
||||
def _step(
|
||||
self, tasks: List[Task], return_logprobs: bool = False
|
||||
) -> Tuple[List[Task], List[Task]]:
|
||||
"""Advance every active task by one token (prefill + decode).
|
||||
|
||||
def _free_pages(self, indices: List[int]) -> None:
|
||||
for idx in indices:
|
||||
self.page_cache.free(idx)
|
||||
Single shared primitive for both the continuous-batching loop and
|
||||
the synchronous ``run_batch`` path, so the two cannot drift.
|
||||
|
||||
def _record_page_hashes(self, task: Task, start_logical_page: int = 0) -> None:
|
||||
full_pages = len(task.prompt_ids) // self.page_size
|
||||
for i in range(start_logical_page, full_pages):
|
||||
self.page_cache.record_page(task.page_table[i], task.prompt_ids, i)
|
||||
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.
|
||||
|
||||
def _remove_finished_tasks(self) -> None:
|
||||
finished = []
|
||||
for task in self.active_tasks:
|
||||
if task.is_finished(self.tokenizer.stop_ids):
|
||||
task.status = TaskStatus.FINISHED
|
||||
task.finish_time = time.time()
|
||||
finished.append(task)
|
||||
self._total_tokens += task.output_tokens
|
||||
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``.
|
||||
|
||||
for task in finished:
|
||||
if not task._pages_freed:
|
||||
self._free_pages(task.page_table)
|
||||
task.page_table.clear()
|
||||
task.n_pages = 0
|
||||
task._pages_freed = True
|
||||
Returns:
|
||||
``(decoded, aborted)``: tasks that produced a new token (its ID
|
||||
already appended to ``output_ids``) and tasks that hit the
|
||||
sequence cap and were marked ``ABORTED``.
|
||||
"""
|
||||
to_prefill = [t for t in tasks if not t.prefill_done and t.prompt_ids]
|
||||
prefilled_ids = set()
|
||||
produced: List[Task] = []
|
||||
if to_prefill:
|
||||
for t in to_prefill:
|
||||
t.input_tokens = len(t.prompt_ids)
|
||||
|
||||
self.active_tasks = [
|
||||
t for t in self.active_tasks if t.status != TaskStatus.FINISHED
|
||||
]
|
||||
groups: Dict[Tuple[int, int, Optional[AttentionBackend]], List[Task]] = {}
|
||||
for t in to_prefill:
|
||||
start_pos = min(
|
||||
self._task_cache.task_cached(t.task_id), len(t.prompt_ids) - 1
|
||||
)
|
||||
groups.setdefault((len(t.prompt_ids), start_pos, t.backend), []).append(
|
||||
t
|
||||
)
|
||||
|
||||
def _refill_active_batch(self) -> None:
|
||||
available = self.max_batch_size - len(self.active_tasks)
|
||||
if available <= 0:
|
||||
return
|
||||
for (prompt_len, start_pos, _), group in groups.items():
|
||||
backend = group[0].backend
|
||||
backend_context = (
|
||||
attn_backend(backend) if backend is not None else nullcontext()
|
||||
)
|
||||
with (
|
||||
backend_context,
|
||||
self._metrics.record([t.task_id for t in group], "prefill"),
|
||||
):
|
||||
prefilled, step_out = self._executor.execute_prefill(
|
||||
group, prompt_len, start_pos, return_logprobs=return_logprobs
|
||||
)
|
||||
|
||||
to_add: List[Task] = []
|
||||
with self._lock:
|
||||
n = min(available, len(self.waiting_queue))
|
||||
for _ in range(n):
|
||||
to_add.append(self.waiting_queue.pop(0))
|
||||
for t, out in zip(prefilled, step_out):
|
||||
t.output_ids.append(out[0] if return_logprobs else out)
|
||||
t.output_tokens += 1
|
||||
t.mark_prefill_done()
|
||||
prefilled_ids.add(t.task_id)
|
||||
produced.append(t)
|
||||
|
||||
failed: List[Task] = []
|
||||
for task in to_add:
|
||||
prompt_len = len(task.prompt_ids)
|
||||
start_logical_page = start_pos // self._cache.page_size
|
||||
for t in group:
|
||||
self._task_cache.task_record_hashes(
|
||||
t.task_id, t.prompt_ids, start_logical_page
|
||||
)
|
||||
|
||||
hit_pages = self.page_cache.lookup_prefix(task.prompt_ids)
|
||||
cached_tokens = len(hit_pages) * self.page_size
|
||||
for p in hit_pages:
|
||||
self.page_cache.inc_ref(p)
|
||||
|
||||
remaining = prompt_len - cached_tokens
|
||||
n_new = self._n_pages_for(remaining) if remaining > 0 else 0
|
||||
new_pages = self.page_cache.alloc_n(n_new) if n_new > 0 else []
|
||||
|
||||
if remaining > 0 and not new_pages:
|
||||
for p in hit_pages:
|
||||
self.page_cache.free(p)
|
||||
failed.append(task)
|
||||
decoded: List[Task] = []
|
||||
aborted: List[Task] = []
|
||||
for t in tasks:
|
||||
if t.task_id in prefilled_ids:
|
||||
continue
|
||||
if self._task_cache.task_extend(t.task_id, t.next_pos):
|
||||
decoded.append(t)
|
||||
else:
|
||||
t.status = TaskStatus.ABORTED
|
||||
aborted.append(t)
|
||||
|
||||
task.page_table = hit_pages + new_pages
|
||||
task.n_pages = len(task.page_table)
|
||||
task._prefix_cached_tokens = cached_tokens
|
||||
task.status = TaskStatus.RUNNING
|
||||
self.active_tasks.append(task)
|
||||
|
||||
if failed:
|
||||
with self._lock:
|
||||
self.waiting_queue[:0] = failed
|
||||
|
||||
def _execute_prefill(
|
||||
self, tasks: List[Task], prompt_len: int, start_pos: int = 0
|
||||
) -> None:
|
||||
tasks = sorted(tasks, key=lambda t: t.task_id)
|
||||
batch_sz = len(tasks)
|
||||
|
||||
seq_len = prompt_len - start_pos
|
||||
input_ids = torch.empty(batch_sz, seq_len, dtype=torch.long, device=self.device)
|
||||
input_mask = torch.ones(batch_sz, seq_len, dtype=torch.bool, device=self.device)
|
||||
|
||||
for i, t in enumerate(tasks):
|
||||
input_ids[i] = torch.tensor(
|
||||
t.prompt_ids[start_pos:prompt_len], device=self.device
|
||||
for backend, group in self._task_backend_groups(decoded):
|
||||
backend_context = (
|
||||
attn_backend(backend) if backend is not None else nullcontext()
|
||||
)
|
||||
with (
|
||||
backend_context,
|
||||
self._metrics.record([t.task_id for t in group], "decode"),
|
||||
):
|
||||
step_out = self._executor.execute_decode(
|
||||
group, return_logprobs=return_logprobs
|
||||
)
|
||||
for t, out in zip(group, step_out):
|
||||
t.output_ids.append(out[0] if return_logprobs else out)
|
||||
t.output_tokens += 1
|
||||
t.advance_kv()
|
||||
produced.append(t)
|
||||
|
||||
page_tables = self._make_page_table_tensor(tasks)
|
||||
return produced, aborted
|
||||
|
||||
with torch.inference_mode():
|
||||
self.model(
|
||||
input_ids,
|
||||
input_mask=input_mask,
|
||||
start_pos=start_pos,
|
||||
paged_cache=self.page_cache.bind(page_tables, total_len=prompt_len),
|
||||
)
|
||||
|
||||
start_logical_page = start_pos // self.page_size
|
||||
for t in tasks:
|
||||
self._record_page_hashes(t, start_logical_page=start_logical_page)
|
||||
|
||||
def _execute_decode(self, tasks: List[Task], start_pos: int) -> None:
|
||||
if not tasks:
|
||||
return
|
||||
|
||||
tasks = sorted(tasks, key=lambda t: t.task_id)
|
||||
batch_sz = len(tasks)
|
||||
|
||||
for t in tasks:
|
||||
self._maybe_alloc_page(t, start_pos)
|
||||
|
||||
input_ids = torch.tensor(
|
||||
[t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks],
|
||||
dtype=torch.long,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
active_mask = torch.ones((batch_sz, 1), dtype=torch.bool, device=self.device)
|
||||
|
||||
page_tables = self._make_page_table_tensor(tasks)
|
||||
total_len = start_pos + 1
|
||||
|
||||
temperatures = torch.tensor([t.temperature for t in tasks], device=self.device)
|
||||
top_ks = torch.tensor([t.top_k for t in tasks], device=self.device)
|
||||
top_ps = torch.tensor([t.top_p for t in tasks], device=self.device)
|
||||
|
||||
with torch.inference_mode():
|
||||
outputs = self.model(
|
||||
input_ids.unsqueeze(1),
|
||||
input_mask=active_mask,
|
||||
paged_cache=self.page_cache.bind(page_tables, total_len=total_len),
|
||||
start_pos=start_pos,
|
||||
)
|
||||
logits = outputs["logits"][:, -1, :]
|
||||
|
||||
next_tokens = sample(
|
||||
logits,
|
||||
temperature=temperatures,
|
||||
top_k=top_ks,
|
||||
top_p=top_ps,
|
||||
).tolist()
|
||||
|
||||
for t, ntok in zip(tasks, next_tokens):
|
||||
t.output_ids.append(ntok)
|
||||
t.output_tokens += 1
|
||||
pos = t.input_tokens + t.output_tokens
|
||||
self._maybe_alloc_page(t, pos)
|
||||
if t.stream_callback:
|
||||
t.stream_callback(self.tokenizer.decode([ntok]))
|
||||
|
||||
for t in tasks:
|
||||
if t.is_finished(self.tokenizer.stop_ids):
|
||||
if t.stream_callback:
|
||||
t.stream_callback(STOP)
|
||||
|
||||
def _make_page_table_tensor(self, tasks: List[Task]) -> Tensor:
|
||||
max_pages = max(t.n_pages for t in tasks)
|
||||
rows = [t.page_table + [-1] * (max_pages - t.n_pages) for t in tasks]
|
||||
return torch.tensor(rows, dtype=torch.long, device=self.device)
|
||||
|
||||
def _maybe_alloc_page(self, task: Task, pos: int) -> None:
|
||||
needed = self._n_pages_for(pos + 1)
|
||||
while task.n_pages < needed:
|
||||
p = self.page_cache.alloc()
|
||||
if p < 0:
|
||||
break
|
||||
task.page_table.append(p)
|
||||
task.n_pages += 1
|
||||
|
||||
def _run_generation_loop(self) -> None:
|
||||
def _run_generation_loop(self):
|
||||
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||
try:
|
||||
while self._running:
|
||||
self._remove_finished_tasks()
|
||||
self._refill_active_batch()
|
||||
with self._backend_context():
|
||||
while not self._stop_event.is_set():
|
||||
finished = self._task_mgr.remove_finished_tasks(stop_ids)
|
||||
for task in finished:
|
||||
if task.status == TaskStatus.FINISHED:
|
||||
self._task_cache.task_record_hashes(
|
||||
task.task_id,
|
||||
self._task_cache.task_cacheable_ids(
|
||||
task.task_id, task.prompt_ids, task.output_ids
|
||||
),
|
||||
)
|
||||
self._task_cache.task_free(task.task_id)
|
||||
|
||||
if not self.active_tasks and not self.waiting_queue:
|
||||
self._task_event.clear()
|
||||
self._task_event.wait(timeout=1.0)
|
||||
continue
|
||||
active = self._task_mgr.get_active_tasks()
|
||||
available = self._task_mgr.max_batch_size - len(active)
|
||||
if available > 0:
|
||||
candidates = self._task_mgr.pull_candidates(available)
|
||||
failed = []
|
||||
for task in candidates:
|
||||
if self._task_cache.task_alloc(
|
||||
task.task_id, task.prompt_ids
|
||||
):
|
||||
self._task_mgr.activate(task)
|
||||
else:
|
||||
failed.append(task)
|
||||
if failed:
|
||||
self._task_mgr.return_to_waiting(failed)
|
||||
|
||||
to_prefill = [t for t in self.active_tasks if t.output_tokens == 0]
|
||||
if to_prefill:
|
||||
for t in to_prefill:
|
||||
t.input_tokens = len(t.prompt_ids)
|
||||
if not self._task_mgr.has_work():
|
||||
self._task_mgr.wait_for_tasks(timeout=1.0)
|
||||
continue
|
||||
|
||||
groups: Dict[Tuple[int, int], List[Task]] = {}
|
||||
for t in to_prefill:
|
||||
key = (len(t.prompt_ids), t._prefix_cached_tokens)
|
||||
groups.setdefault(key, []).append(t)
|
||||
active = self._task_mgr.get_active_tasks()
|
||||
|
||||
for (prompt_len, start_pos), group in groups.items():
|
||||
if start_pos < prompt_len:
|
||||
self._execute_prefill(group, prompt_len, start_pos)
|
||||
decoded, aborted = self._step(active)
|
||||
|
||||
pos_groups: Dict[int, List[Task]] = {}
|
||||
for t in self.active_tasks:
|
||||
pos_groups.setdefault(t.next_pos, []).append(t)
|
||||
for t in aborted:
|
||||
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||
|
||||
for t in decoded:
|
||||
new_text = t.decode_new_token(self._task_mgr.tokenizer)
|
||||
if new_text:
|
||||
self._task_mgr.invoke_callback(t.task_id, new_text)
|
||||
if t.is_finished(stop_ids):
|
||||
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||
|
||||
if pos_groups:
|
||||
best_pos = max(pos_groups, key=lambda p: len(pos_groups[p]))
|
||||
self._execute_decode(pos_groups[best_pos], best_pos)
|
||||
except Exception as e:
|
||||
self._stop_event.set()
|
||||
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
||||
for task in self.active_tasks:
|
||||
if task.stream_callback:
|
||||
task.stream_callback(STOP)
|
||||
for task in self.waiting_queue:
|
||||
if task.stream_callback:
|
||||
task.stream_callback(STOP)
|
||||
raise
|
||||
for task in self._task_mgr.get_active_tasks():
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._task_cache.task_free(task.task_id)
|
||||
for task in self._task_mgr.get_waiting_tasks():
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._task_mgr.clear_queues()
|
||||
|
||||
def start(self) -> None:
|
||||
if not self._running:
|
||||
self._running = True
|
||||
t = threading.Thread(target=self._run_generation_loop, daemon=True)
|
||||
t.start()
|
||||
self._loop_thread = t
|
||||
def start(self):
|
||||
if self._loop_thread is not None and self._loop_thread.is_alive():
|
||||
return
|
||||
self._stop_event.clear()
|
||||
t = threading.Thread(target=self._run_generation_loop, daemon=True)
|
||||
t.start()
|
||||
self._loop_thread = t
|
||||
|
||||
def stop(self) -> None:
|
||||
self._running = False
|
||||
self._task_event.set()
|
||||
if hasattr(self, "_loop_thread"):
|
||||
def stop(self):
|
||||
self._stop_event.set()
|
||||
self._task_mgr.wake()
|
||||
if self._loop_thread is not None:
|
||||
self._loop_thread.join(timeout=2.0)
|
||||
self.waiting_queue.clear()
|
||||
self.active_tasks.clear()
|
||||
self._loop_thread = None
|
||||
for task in self._task_mgr.get_active_tasks():
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._task_cache.task_free(task.task_id)
|
||||
for task in self._task_mgr.get_waiting_tasks():
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._task_cache.task_free(task.task_id)
|
||||
self._task_mgr.clear_queues()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def get_stats(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"total_tasks": self._total_tasks,
|
||||
"total_tokens": self._total_tokens,
|
||||
"active_tasks": len(self.active_tasks),
|
||||
"waiting_queue": len(self.waiting_queue),
|
||||
}
|
||||
def run_batch(
|
||||
self,
|
||||
prompt_ids_list: List[List[int]],
|
||||
*,
|
||||
max_tokens: Optional[int] = None,
|
||||
temperature: float = 1.0,
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
return_logprobs: bool = False,
|
||||
) -> List[List[int]]:
|
||||
"""Synchronous batch generation without the scheduler thread.
|
||||
|
||||
Accepts already-tokenized prompts (no string round-trip) and runs
|
||||
prefill + decode to completion on the calling thread. Designed for
|
||||
RL rollout, where logprobs of the behaviour policy must be collected
|
||||
alongside generated tokens.
|
||||
|
||||
Args:
|
||||
prompt_ids_list: ``B`` prompts, each a list of token IDs.
|
||||
max_tokens: Maximum tokens to generate per prompt. ``None``
|
||||
uses ``self.max_seq_len - len(prompt_ids)``.
|
||||
temperature/top_p/top_k/frequency_penalty/rep_window: Sampling
|
||||
parameters (uniform across the batch).
|
||||
return_logprobs: If ``True``, return ``(token_ids, logprobs)``
|
||||
tuples per prompt (logprobs aligned 1-to-1 with token_ids).
|
||||
|
||||
Returns:
|
||||
``List[List[int]]`` of generated token IDs per prompt, or —
|
||||
when ``return_logprobs`` is ``True`` —
|
||||
``List[Tuple[List[int], List[float]]]``.
|
||||
"""
|
||||
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||
seq_cap = self.max_seq_len
|
||||
request_backend = get_backend(use_default=False)
|
||||
|
||||
tasks: List[Task] = []
|
||||
for ids in prompt_ids_list:
|
||||
if len(ids) >= seq_cap:
|
||||
tasks.append(None)
|
||||
continue
|
||||
t_max = max_tokens
|
||||
if t_max is None:
|
||||
t_max = seq_cap - len(ids)
|
||||
else:
|
||||
t_max = min(t_max, seq_cap - len(ids))
|
||||
if t_max <= 0:
|
||||
tasks.append(None)
|
||||
continue
|
||||
task = Task(
|
||||
task_id=f"batch_{uuid.uuid4().hex[:8]}",
|
||||
prompt_ids=list(ids),
|
||||
max_tokens=t_max,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
frequency_penalty=frequency_penalty,
|
||||
rep_window=rep_window,
|
||||
backend=request_backend,
|
||||
)
|
||||
if not self._task_cache.task_alloc(task.task_id, task.prompt_ids):
|
||||
tasks.append(None)
|
||||
continue
|
||||
task.input_tokens = len(task.prompt_ids)
|
||||
self._metrics.register(task.task_id)
|
||||
tasks.append(task)
|
||||
|
||||
try:
|
||||
live = [t for t in tasks if t is not None]
|
||||
|
||||
with self._backend_context():
|
||||
while live:
|
||||
decoded, _ = self._step(live, return_logprobs=return_logprobs)
|
||||
live = [t for t in decoded if not t.is_finished(stop_ids)]
|
||||
finally:
|
||||
for t in tasks:
|
||||
if t is not None:
|
||||
self._metrics.mark_finished(
|
||||
t.task_id, t.input_tokens, t.output_tokens
|
||||
)
|
||||
self._task_cache.task_free(t.task_id)
|
||||
|
||||
results: List[Any] = []
|
||||
for t in tasks:
|
||||
if t is None:
|
||||
results.append(([], []) if return_logprobs else [])
|
||||
elif return_logprobs:
|
||||
results.append((list(t.output_ids), list(t.output_logprobs)))
|
||||
else:
|
||||
results.append(list(t.output_ids))
|
||||
return results
|
||||
|
||||
@@ -1,486 +0,0 @@
|
||||
"""
|
||||
OpenAI / Anthropic-compatible chat completion server backed by continuous-batching inference.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import torch
|
||||
import uvicorn
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
from astrai.model import AutoModel
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_project_root = Path(__file__).parent.parent.parent
|
||||
|
||||
|
||||
class ServerState:
|
||||
def __init__(self):
|
||||
self.engine: Optional[InferenceEngine] = None
|
||||
self.config: Dict[str, Any] = {
|
||||
"device": "cuda",
|
||||
"dtype": torch.bfloat16,
|
||||
"param_path": None,
|
||||
"max_batch_size": 16,
|
||||
}
|
||||
|
||||
|
||||
_state = ServerState()
|
||||
|
||||
|
||||
class ChatMessage(BaseModel):
|
||||
role: str
|
||||
content: str
|
||||
|
||||
|
||||
class ChatCompletionRequest(BaseModel):
|
||||
"""OpenAI Chat Completion API request body."""
|
||||
|
||||
model: str = "astrai"
|
||||
messages: List[ChatMessage]
|
||||
temperature: Optional[float] = Field(default=1.0, ge=0.0, le=2.0)
|
||||
top_p: Optional[float] = Field(default=1.0, ge=0.0, le=1.0)
|
||||
top_k: Optional[int] = Field(default=50, ge=1)
|
||||
stream: Optional[bool] = False
|
||||
stop: Optional[Union[str, List[str]]] = None
|
||||
max_tokens: Optional[int] = Field(default=2048, ge=1)
|
||||
n: Optional[int] = Field(default=1, ge=1)
|
||||
presence_penalty: Optional[float] = Field(default=0.0, ge=-2.0, le=2.0)
|
||||
frequency_penalty: Optional[float] = Field(default=0.0, ge=-2.0, le=2.0)
|
||||
logit_bias: Optional[Dict[int, float]] = None
|
||||
user: Optional[str] = None
|
||||
|
||||
|
||||
class AnthropicMessage(BaseModel):
|
||||
role: str
|
||||
content: Union[str, List[Dict[str, Any]]]
|
||||
|
||||
|
||||
class MessagesRequest(BaseModel):
|
||||
"""Anthropic Messages API request body."""
|
||||
|
||||
model: str = "astrai"
|
||||
max_tokens: int = Field(default=1024, ge=1)
|
||||
messages: List[AnthropicMessage]
|
||||
system: Optional[str] = None
|
||||
temperature: Optional[float] = Field(default=1.0, ge=0.0, le=2.0)
|
||||
top_p: Optional[float] = Field(default=1.0, ge=0.0, le=1.0)
|
||||
top_k: Optional[int] = Field(default=50, ge=1)
|
||||
stream: Optional[bool] = False
|
||||
stop_sequences: Optional[List[str]] = None
|
||||
|
||||
|
||||
def configure_server(
|
||||
device: str = "cuda",
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
param_path: Optional[Path] = None,
|
||||
max_batch_size: int = 16,
|
||||
):
|
||||
_state.config.update(
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
param_path=param_path,
|
||||
max_batch_size=max_batch_size,
|
||||
)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
try:
|
||||
load_model(
|
||||
param_path=_state.config["param_path"],
|
||||
device=_state.config["device"],
|
||||
dtype=_state.config["dtype"],
|
||||
max_batch_size=_state.config["max_batch_size"],
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to load model: {e}")
|
||||
raise
|
||||
yield
|
||||
if _state.engine:
|
||||
_state.engine.shutdown()
|
||||
logger.info("Inference engine shutdown complete")
|
||||
|
||||
|
||||
app = FastAPI(title="AstrAI Inference Server", version="0.2.0", lifespan=lifespan)
|
||||
|
||||
|
||||
def load_model(
|
||||
param_path: Optional[Path] = None,
|
||||
device: str = "cuda",
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
max_batch_size: int = 16,
|
||||
):
|
||||
if param_path is None:
|
||||
param_path = _project_root / "params"
|
||||
if not param_path.exists():
|
||||
raise FileNotFoundError(f"Parameter directory not found: {param_path}")
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||
model = AutoModel.from_pretrained(param_path)
|
||||
model.to(device=device, dtype=dtype)
|
||||
logger.info(f"Model loaded on {device} with dtype {dtype}")
|
||||
|
||||
_state.engine = InferenceEngine(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=max_batch_size,
|
||||
)
|
||||
logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}")
|
||||
|
||||
|
||||
def _get_engine() -> InferenceEngine:
|
||||
if _state.engine is None:
|
||||
raise HTTPException(status_code=503, detail="Engine not initialized")
|
||||
return _state.engine
|
||||
|
||||
|
||||
def _make_chunk(
|
||||
delta: Dict[str, str],
|
||||
finish_reason: Optional[str] = None,
|
||||
*,
|
||||
resp_id: str,
|
||||
created: int,
|
||||
model: str,
|
||||
index: int = 0,
|
||||
) -> str:
|
||||
"""Build a single SSE ``data:`` chunk matching OpenAI streaming format."""
|
||||
data = {
|
||||
"id": resp_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": created,
|
||||
"model": model,
|
||||
"choices": [
|
||||
{
|
||||
"index": index,
|
||||
"delta": delta,
|
||||
"finish_reason": finish_reason,
|
||||
}
|
||||
],
|
||||
}
|
||||
return f"data: {json.dumps(data, ensure_ascii=False)}\n\n"
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
async def health():
|
||||
return {
|
||||
"status": "ok",
|
||||
"model_loaded": _state.engine is not None,
|
||||
}
|
||||
|
||||
|
||||
@app.get("/stats")
|
||||
async def get_stats():
|
||||
return _get_engine().get_stats()
|
||||
|
||||
|
||||
@app.post("/v1/chat/completions")
|
||||
async def chat_completion(request: ChatCompletionRequest):
|
||||
"""OpenAI-compatible chat completion endpoint (streaming + non-streaming)."""
|
||||
engine = _get_engine()
|
||||
resp_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
||||
created = int(time.time())
|
||||
model = request.model
|
||||
|
||||
prompt = engine.tokenizer.apply_chat_template(
|
||||
[{"role": m.role, "content": m.content} for m in request.messages],
|
||||
tokenize=False,
|
||||
)
|
||||
prompt_tokens = len(engine.tokenizer.encode(prompt))
|
||||
|
||||
if request.stream:
|
||||
agen = engine.generate_async(
|
||||
prompt=prompt,
|
||||
max_tokens=request.max_tokens,
|
||||
temperature=request.temperature,
|
||||
top_p=request.top_p,
|
||||
top_k=request.top_k,
|
||||
)
|
||||
|
||||
async def event_stream():
|
||||
yield _make_chunk(
|
||||
{"role": "assistant"},
|
||||
finish_reason=None,
|
||||
resp_id=resp_id,
|
||||
created=created,
|
||||
model=model,
|
||||
)
|
||||
|
||||
completion_tokens = 0
|
||||
async for token in agen:
|
||||
yield _make_chunk(
|
||||
{"content": token},
|
||||
finish_reason=None,
|
||||
resp_id=resp_id,
|
||||
created=created,
|
||||
model=model,
|
||||
)
|
||||
completion_tokens += 1
|
||||
|
||||
yield _make_chunk(
|
||||
{},
|
||||
finish_reason="stop",
|
||||
resp_id=resp_id,
|
||||
created=created,
|
||||
model=model,
|
||||
)
|
||||
|
||||
usage = {
|
||||
"prompt_tokens": prompt_tokens,
|
||||
"completion_tokens": completion_tokens,
|
||||
"total_tokens": prompt_tokens + completion_tokens,
|
||||
}
|
||||
yield f"data: {json.dumps(usage, ensure_ascii=False)}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
return StreamingResponse(
|
||||
event_stream(),
|
||||
media_type="text/event-stream",
|
||||
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
||||
)
|
||||
|
||||
completion_tokens = 0
|
||||
chunks: List[str] = []
|
||||
agen = engine.generate_async(
|
||||
prompt=prompt,
|
||||
max_tokens=request.max_tokens,
|
||||
temperature=request.temperature,
|
||||
top_p=request.top_p,
|
||||
top_k=request.top_k,
|
||||
)
|
||||
async for token in agen:
|
||||
chunks.append(token)
|
||||
completion_tokens += 1
|
||||
content = "".join(chunks)
|
||||
|
||||
return {
|
||||
"id": resp_id,
|
||||
"object": "chat.completion",
|
||||
"created": created,
|
||||
"model": model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": content},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": prompt_tokens,
|
||||
"completion_tokens": completion_tokens,
|
||||
"total_tokens": prompt_tokens + completion_tokens,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _make_anthropic_sse(event: str, data: Dict[str, Any]) -> str:
|
||||
return f"event: {event}\ndata: {json.dumps(data, ensure_ascii=False)}\n\n"
|
||||
|
||||
|
||||
def _check_stop_sequence(text: str, stop_sequences: List[str]) -> Optional[str]:
|
||||
for seq in stop_sequences:
|
||||
if seq and seq in text:
|
||||
return seq
|
||||
return None
|
||||
|
||||
|
||||
def _extract_text_content(content: Union[str, List[Dict[str, Any]]]) -> str:
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
for block in content:
|
||||
if isinstance(block, dict) and block.get("type") == "text":
|
||||
return block.get("text", "")
|
||||
return ""
|
||||
|
||||
|
||||
def _build_anthropic_messages(
|
||||
messages: List[AnthropicMessage], system: Optional[str]
|
||||
) -> List[Dict[str, str]]:
|
||||
result: List[Dict[str, str]] = []
|
||||
if system:
|
||||
result.append({"role": "system", "content": system})
|
||||
for m in messages:
|
||||
content = _extract_text_content(m.content)
|
||||
if content:
|
||||
result.append({"role": m.role, "content": content})
|
||||
return result
|
||||
|
||||
|
||||
@app.post("/v1/messages")
|
||||
async def create_message(request: MessagesRequest):
|
||||
"""Anthropic-compatible Messages API endpoint (streaming + non-streaming)."""
|
||||
engine = _get_engine()
|
||||
resp_id = f"msg_{uuid.uuid4().hex[:24]}"
|
||||
model = request.model
|
||||
|
||||
chat_messages = _build_anthropic_messages(request.messages, request.system)
|
||||
prompt = engine.tokenizer.apply_chat_template(chat_messages, tokenize=False)
|
||||
prompt_tokens = len(engine.tokenizer.encode(prompt))
|
||||
|
||||
stop_sequences = request.stop_sequences or []
|
||||
|
||||
if request.stream:
|
||||
agen = engine.generate_async(
|
||||
prompt=prompt,
|
||||
max_tokens=request.max_tokens,
|
||||
temperature=request.temperature,
|
||||
top_p=request.top_p,
|
||||
top_k=request.top_k,
|
||||
)
|
||||
|
||||
async def event_stream():
|
||||
yield _make_anthropic_sse(
|
||||
"message_start",
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": resp_id,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": model,
|
||||
"content": [],
|
||||
"usage": {"input_tokens": prompt_tokens},
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
yield _make_anthropic_sse(
|
||||
"content_block_start",
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 0,
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
},
|
||||
)
|
||||
|
||||
completion_tokens = 0
|
||||
accumulated = ""
|
||||
stopped_seq: Optional[str] = None
|
||||
async for token in agen:
|
||||
accumulated += token
|
||||
completion_tokens += 1
|
||||
|
||||
matched = _check_stop_sequence(accumulated, stop_sequences)
|
||||
if matched:
|
||||
text = accumulated[: accumulated.rfind(matched)]
|
||||
stopped_seq = matched
|
||||
if text:
|
||||
yield _make_anthropic_sse(
|
||||
"content_block_delta",
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": text},
|
||||
},
|
||||
)
|
||||
break
|
||||
|
||||
yield _make_anthropic_sse(
|
||||
"content_block_delta",
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": token},
|
||||
},
|
||||
)
|
||||
|
||||
yield _make_anthropic_sse(
|
||||
"content_block_stop",
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
)
|
||||
|
||||
stop_reason = "stop_sequence" if stopped_seq else "end_turn"
|
||||
yield _make_anthropic_sse(
|
||||
"message_delta",
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": stop_reason, "stop_sequence": stopped_seq},
|
||||
"usage": {"output_tokens": completion_tokens},
|
||||
},
|
||||
)
|
||||
|
||||
yield _make_anthropic_sse(
|
||||
"message_stop",
|
||||
{"type": "message_stop"},
|
||||
)
|
||||
|
||||
return StreamingResponse(
|
||||
event_stream(),
|
||||
media_type="text/event-stream",
|
||||
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
||||
)
|
||||
|
||||
completion_tokens = 0
|
||||
chunks: List[str] = []
|
||||
agen = engine.generate_async(
|
||||
prompt=prompt,
|
||||
max_tokens=request.max_tokens,
|
||||
temperature=request.temperature,
|
||||
top_p=request.top_p,
|
||||
top_k=request.top_k,
|
||||
)
|
||||
stopped_seq: Optional[str] = None
|
||||
accumulated = ""
|
||||
async for token in agen:
|
||||
chunks.append(token)
|
||||
completion_tokens += 1
|
||||
accumulated += token
|
||||
matched = _check_stop_sequence(accumulated, stop_sequences)
|
||||
if matched:
|
||||
stopped_seq = matched
|
||||
break
|
||||
|
||||
content = "".join(chunks)
|
||||
if stopped_seq:
|
||||
idx = content.rfind(stopped_seq)
|
||||
if idx != -1:
|
||||
content = content[:idx]
|
||||
|
||||
return {
|
||||
"id": resp_id,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": model,
|
||||
"content": [{"type": "text", "text": content}],
|
||||
"stop_reason": "stop_sequence" if stopped_seq else "end_turn",
|
||||
"stop_sequence": stopped_seq,
|
||||
"usage": {
|
||||
"input_tokens": prompt_tokens,
|
||||
"output_tokens": completion_tokens,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def run_server(
|
||||
host: str = "0.0.0.0",
|
||||
port: int = 8000,
|
||||
reload: bool = False,
|
||||
device: str = "cuda",
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
param_path: Optional[Path] = None,
|
||||
max_batch_size: int = 16,
|
||||
):
|
||||
configure_server(
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
param_path=param_path,
|
||||
max_batch_size=max_batch_size,
|
||||
)
|
||||
uvicorn.run(
|
||||
"astrai.inference.server:app",
|
||||
host=host,
|
||||
port=port,
|
||||
reload=reload,
|
||||
)
|
||||
@@ -0,0 +1,290 @@
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from collections import deque
|
||||
from enum import Enum
|
||||
from typing import TYPE_CHECKING, Any, Callable, Deque, Dict, List, Optional
|
||||
|
||||
from tokenizers.decoders import DecodeStream
|
||||
|
||||
from astrai.inference.metrics import MetricsCollector
|
||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from astrai.extension import AttentionBackend
|
||||
|
||||
STOP = object()
|
||||
|
||||
|
||||
class StreamDecoder:
|
||||
"""Incremental decoder backed by the tokenizers library's DecodeStream.
|
||||
|
||||
Delegates to the Rust-native streaming decoder which maintains an
|
||||
O(1) bounded token buffer internally (via prefix drain), avoiding
|
||||
the O(n²) cost of re-decoding the full history on each step.
|
||||
|
||||
Multi-byte UTF-8 sequences split across token boundaries are
|
||||
buffered until complete; ``push`` returns "" while the trailing
|
||||
sequence is still incomplete.
|
||||
"""
|
||||
|
||||
__slots__ = ("_stream", "_tok")
|
||||
|
||||
def __init__(self, tokenizer: AutoTokenizer):
|
||||
self._tok = tokenizer._tokenizer
|
||||
self._stream = DecodeStream(skip_special_tokens=True)
|
||||
|
||||
def push(self, token_id: int) -> str:
|
||||
"""Append a token ID and return newly completed text.
|
||||
|
||||
Returns "" while a multi-byte character is still incomplete.
|
||||
"""
|
||||
chunk = self._stream.step(self._tok, token_id)
|
||||
return chunk or ""
|
||||
|
||||
|
||||
class TaskStatus(Enum):
|
||||
"""Task lifecycle states."""
|
||||
|
||||
PENDING = "pending"
|
||||
RUNNING = "running"
|
||||
FINISHED = "finished"
|
||||
ABORTED = "aborted"
|
||||
|
||||
|
||||
class Task:
|
||||
"""Single generation request: prompt, sampling params, output state."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
task_id: str,
|
||||
prompt_ids: List[int],
|
||||
max_tokens: Optional[int] = None,
|
||||
temperature: float = 1.0,
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
backend: Optional["AttentionBackend"] = None,
|
||||
):
|
||||
self.task_id = task_id
|
||||
self.prompt_ids = prompt_ids
|
||||
self.max_tokens = max_tokens
|
||||
self.temperature = temperature
|
||||
self.top_p = top_p
|
||||
self.top_k = top_k
|
||||
self.frequency_penalty = frequency_penalty
|
||||
self.rep_window = rep_window
|
||||
self.backend = backend
|
||||
|
||||
self.status = TaskStatus.PENDING
|
||||
self.output_ids: List[int] = []
|
||||
self.output_logprobs: List[float] = []
|
||||
self.input_tokens: int = 0
|
||||
self.output_tokens: int = 0
|
||||
self._kv_len: int = 0
|
||||
self._decoder: Optional[StreamDecoder] = None
|
||||
|
||||
def mark_prefill_done(self):
|
||||
"""Prompt KV is materialized by prefill; first output sampled but
|
||||
not yet written to KV."""
|
||||
self._kv_len = self.input_tokens
|
||||
|
||||
def advance_kv(self):
|
||||
"""One more position written to KV (after a decode forward)."""
|
||||
self._kv_len += 1
|
||||
|
||||
def decode_new_token(self, tokenizer: AutoTokenizer) -> str:
|
||||
"""Decode the last appended output token, buffering incomplete
|
||||
multi-byte sequences across calls.
|
||||
|
||||
Lazily creates a :class:`StreamDecoder` on first use.
|
||||
"""
|
||||
if self._decoder is None:
|
||||
self._decoder = StreamDecoder(tokenizer)
|
||||
return self._decoder.push(self.output_ids[-1])
|
||||
|
||||
@property
|
||||
def next_pos(self) -> int:
|
||||
"""KV position where the next decode step will write."""
|
||||
return self._kv_len
|
||||
|
||||
@property
|
||||
def prefill_done(self) -> bool:
|
||||
"""True when all prompt KV entries are materialized."""
|
||||
return self._kv_len >= self.input_tokens > 0
|
||||
|
||||
def is_finished(self, stop_ids: List[int]) -> bool:
|
||||
if self.max_tokens is not None and self.output_tokens >= self.max_tokens:
|
||||
return True
|
||||
if self.output_ids and self.output_ids[-1] in stop_ids:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
class TaskManager:
|
||||
"""Thread-safe task queues and lifecycle transitions (no page ops)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: AutoTokenizer,
|
||||
max_batch_size: int = 16,
|
||||
max_seq_len: int = 8192,
|
||||
metrics: Optional["MetricsCollector"] = None,
|
||||
):
|
||||
self.tokenizer = tokenizer
|
||||
self.max_batch_size = max_batch_size
|
||||
self.max_seq_len = max_seq_len
|
||||
|
||||
self.waiting_queue: Deque[Task] = deque()
|
||||
self.active_tasks: List[Task] = []
|
||||
self._callbacks: Dict[str, Callable[[str], None]] = {}
|
||||
|
||||
self._task_event = threading.Event()
|
||||
self._lock = threading.Lock()
|
||||
|
||||
self._total_tasks = 0
|
||||
self._total_tokens = 0
|
||||
|
||||
self._metrics = metrics
|
||||
|
||||
def add_task(
|
||||
self,
|
||||
prompt: str,
|
||||
max_tokens: Optional[int] = None,
|
||||
temperature: float = 1.0,
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
backend: Optional["AttentionBackend"] = None,
|
||||
stream_callback: Optional[Callable[[str], None]] = None,
|
||||
) -> str:
|
||||
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
|
||||
prompt_ids = self.tokenizer.encode(prompt)
|
||||
if len(prompt_ids) > self.max_seq_len:
|
||||
prompt_ids = prompt_ids[-self.max_seq_len :]
|
||||
|
||||
if max_tokens is None:
|
||||
max_tokens = self.max_seq_len - len(prompt_ids)
|
||||
else:
|
||||
max_tokens = min(max_tokens, self.max_seq_len - len(prompt_ids))
|
||||
|
||||
task = Task(
|
||||
task_id=task_id,
|
||||
prompt_ids=prompt_ids,
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
frequency_penalty=frequency_penalty,
|
||||
rep_window=rep_window,
|
||||
backend=backend,
|
||||
)
|
||||
|
||||
with self._lock:
|
||||
self.waiting_queue.append(task)
|
||||
self._total_tasks += 1
|
||||
if stream_callback:
|
||||
self._callbacks[task_id] = stream_callback
|
||||
|
||||
if self._metrics is not None:
|
||||
self._metrics.register(task_id)
|
||||
|
||||
self._task_event.set()
|
||||
return task_id
|
||||
|
||||
def remove_task(self, task_id: str) -> List[Task]:
|
||||
with self._lock:
|
||||
removed_active = [t for t in self.active_tasks if t.task_id == task_id]
|
||||
self.waiting_queue = deque(
|
||||
t for t in self.waiting_queue if t.task_id != task_id
|
||||
)
|
||||
self.active_tasks = [t for t in self.active_tasks if t.task_id != task_id]
|
||||
self._callbacks.pop(task_id, None)
|
||||
return removed_active
|
||||
|
||||
def invoke_callback(self, task_id: str, token: str):
|
||||
cb = self._callbacks.get(task_id)
|
||||
if cb:
|
||||
cb(token)
|
||||
|
||||
def get_stats(self) -> Dict[str, Any]:
|
||||
stats: Dict[str, Any] = {
|
||||
"total_tasks": self._total_tasks,
|
||||
"total_tokens": self._total_tokens,
|
||||
"active_tasks": len(self.active_tasks),
|
||||
"waiting_queue": len(self.waiting_queue),
|
||||
}
|
||||
if self._metrics is not None:
|
||||
stats.update(self._metrics.get_stats())
|
||||
return stats
|
||||
|
||||
def remove_finished_tasks(self, stop_ids: List[int]) -> List[Task]:
|
||||
with self._lock:
|
||||
finished = []
|
||||
for task in self.active_tasks:
|
||||
if task.status == TaskStatus.ABORTED:
|
||||
finished.append(task)
|
||||
elif task.is_finished(stop_ids):
|
||||
task.status = TaskStatus.FINISHED
|
||||
finished.append(task)
|
||||
self._total_tokens += task.output_tokens
|
||||
|
||||
if self._metrics is not None:
|
||||
for task in finished:
|
||||
self._metrics.mark_finished(
|
||||
task.task_id, task.input_tokens, task.output_tokens
|
||||
)
|
||||
|
||||
self.active_tasks = [
|
||||
t
|
||||
for t in self.active_tasks
|
||||
if t.status not in (TaskStatus.FINISHED, TaskStatus.ABORTED)
|
||||
]
|
||||
return finished
|
||||
|
||||
def pull_candidates(self, n: int) -> List[Task]:
|
||||
to_add: List[Task] = []
|
||||
with self._lock:
|
||||
take = min(n, len(self.waiting_queue))
|
||||
for _ in range(take):
|
||||
to_add.append(self.waiting_queue.popleft())
|
||||
return to_add
|
||||
|
||||
def activate(self, task: Task):
|
||||
task.status = TaskStatus.RUNNING
|
||||
with self._lock:
|
||||
self.active_tasks.append(task)
|
||||
|
||||
def return_to_waiting(self, tasks: List[Task]):
|
||||
with self._lock:
|
||||
for task in reversed(tasks):
|
||||
self.waiting_queue.appendleft(task)
|
||||
|
||||
def has_work(self) -> bool:
|
||||
return bool(self.active_tasks or self.waiting_queue)
|
||||
|
||||
def wait_for_tasks(self, timeout: float = 1.0):
|
||||
with self._lock:
|
||||
if self.waiting_queue or self.active_tasks:
|
||||
return
|
||||
self._task_event.clear()
|
||||
self._task_event.wait(timeout=timeout)
|
||||
|
||||
def get_active_tasks(self) -> List[Task]:
|
||||
with self._lock:
|
||||
return list(self.active_tasks)
|
||||
|
||||
def get_waiting_tasks(self) -> List[Task]:
|
||||
with self._lock:
|
||||
return list(self.waiting_queue)
|
||||
|
||||
def clear_queues(self):
|
||||
with self._lock:
|
||||
self.waiting_queue.clear()
|
||||
self.active_tasks.clear()
|
||||
self._callbacks.clear()
|
||||
|
||||
def wake(self):
|
||||
self._task_event.set()
|
||||
@@ -0,0 +1,159 @@
|
||||
"""Pre-allocated buffers for the inference decode hot path.
|
||||
|
||||
Mirrors FlashInfer / SGLang's global workspace pattern: all per-step tensors
|
||||
are allocated eagerly at init (nothing is lazy), so the decode step
|
||||
reads/writes fixed-address tensors with zero ``torch.empty`` calls during
|
||||
the hot loop — a prerequisite for CUDA-graph capture.
|
||||
"""
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
_MAX_SPLITS = 32
|
||||
Q_TILE_ROWS = 64
|
||||
|
||||
|
||||
class InferenceWorkspace:
|
||||
"""Reusable fixed-shape per-step buffers for decode.
|
||||
|
||||
Families of buffers, all sized to ``max_batch_size`` / ``max_seq_len``
|
||||
and sliced via views each step:
|
||||
|
||||
- ``decode_mask``: a ``[B, 1, total_len]`` validity mask, the RHS
|
||||
``arange`` pre-computed so only a single ``torch.ge(out=)`` runs per
|
||||
step.
|
||||
- ``input_ids``: per-step token IDs filled from host (pinned, double-
|
||||
buffered so an in-flight async H2D copy never races the next fill).
|
||||
- KV-cache bind metadata (``req_pool_indices``, ``seq_lens``,
|
||||
``kv_indptr``, ``inc``, ``out_cache_loc``), written by
|
||||
``PagePool.bind_tasks`` when the Executor passes this workspace.
|
||||
- ``decode_o_part`` / ``decode_ml_part``: split-KV partial result buffers
|
||||
(mirrors FlashInfer's workspace). One global alloc, reused by every
|
||||
decode step across all layers. Sliced views are passed to the CUDA
|
||||
attention kernel so its internal ``torch.empty`` hot-path alloc goes
|
||||
through a stable address (CUDA-graph capturable).
|
||||
|
||||
No re-allocation while the server's bounds are respected.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
max_batch_size: int,
|
||||
max_seq_len: int,
|
||||
max_q_heads: int,
|
||||
head_dim: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
self.max_batch_size = max_batch_size
|
||||
self.max_seq_len = max_seq_len
|
||||
self.max_q_heads = max_q_heads
|
||||
self.head_dim = head_dim
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
|
||||
# ``position_ids[:, None, None] >= arange`` RHS, reused every step.
|
||||
self.arange = torch.arange(max_seq_len, device=device)
|
||||
# Decode validity mask: [max_batch, 1, max_seq_len] bool.
|
||||
self.input_mask = torch.empty(
|
||||
(max_batch_size, 1, max_seq_len), dtype=torch.bool, device=device
|
||||
)
|
||||
|
||||
# Per-step token IDs. Values come from host Python lists every
|
||||
# step, so the device buffer is pre-allocated (stable address for
|
||||
# CUDA-graph capture) and filled via a host staging buffer. A
|
||||
# double buffer keeps a copy in flight from being overwritten by
|
||||
# the next fill.
|
||||
self.input_ids = torch.empty((max_batch_size,), dtype=torch.long, device=device)
|
||||
self._pin = [
|
||||
torch.empty((max_batch_size,), dtype=torch.long),
|
||||
torch.empty((max_batch_size,), dtype=torch.long),
|
||||
]
|
||||
self._pin_idx = 0
|
||||
|
||||
# KV-cache bind metadata (fixed shape, written by ``PagePool.bind_tasks``
|
||||
# when the Executor passes this workspace). Stable addresses make the
|
||||
# decode forward CUDA-graph capturable.
|
||||
self.req_pool_indices = torch.empty(
|
||||
(max_batch_size,), dtype=torch.int32, device=device
|
||||
)
|
||||
self.seq_lens = torch.empty((max_batch_size,), dtype=torch.long, device=device)
|
||||
self.kv_indptr = torch.empty(
|
||||
(max_batch_size + 1,), dtype=torch.int32, device=device
|
||||
)
|
||||
self.qo_indptr = torch.empty(
|
||||
(max_batch_size + 1,), dtype=torch.int32, device=device
|
||||
)
|
||||
max_q_tiles = max_batch_size * ((max_seq_len + Q_TILE_ROWS - 1) // Q_TILE_ROWS)
|
||||
self.q_tile_to_batch = torch.empty(
|
||||
(max_q_tiles,), dtype=torch.int32, device=device
|
||||
)
|
||||
self.q_tile_to_index = torch.empty(
|
||||
(max_q_tiles,), dtype=torch.int32, device=device
|
||||
)
|
||||
self.inc = torch.arange(max_batch_size + 1, dtype=torch.int32, device=device)
|
||||
self.out_cache_loc = torch.empty(
|
||||
(max_batch_size, 1), dtype=torch.int32, device=device
|
||||
)
|
||||
|
||||
# Per-step position IDs (must be at a fixed address for CUDA-graph capture).
|
||||
self.position_ids = torch.empty(
|
||||
(max_batch_size,), dtype=torch.long, device=device
|
||||
)
|
||||
|
||||
# Split-KV partial-result buffers for decode (persistent, one global
|
||||
# alloc per process — mirrors FlashInfer's workspace pattern).
|
||||
# Shape: [max_batch_size, max_q_heads, _MAX_SPLITS, head_dim] (o_part)
|
||||
# [max_batch_size, max_q_heads, _MAX_SPLITS, 2] (ml_part)
|
||||
self.decode_o_part = torch.empty(
|
||||
(max_batch_size, max_q_heads, _MAX_SPLITS, head_dim),
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
self.decode_ml_part = torch.empty(
|
||||
(max_batch_size, max_q_heads, _MAX_SPLITS, 2),
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
|
||||
# Decode output buffer (graph-safe pre-alloc). Shape matches the
|
||||
# decode kernel's output: [batch, q_head, head_dim].
|
||||
self.decode_out = torch.empty(
|
||||
(max_batch_size, max_q_heads, head_dim),
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
)
|
||||
|
||||
def decode_buffers(self, batch: int, q_heads: int):
|
||||
"""Return ``(o_part, ml_part)`` view sliced to live dimensions."""
|
||||
return (
|
||||
self.decode_o_part[:batch, :q_heads],
|
||||
self.decode_ml_part[:batch, :q_heads],
|
||||
)
|
||||
|
||||
def fill_input_ids(self, ids: "list[int]") -> Tensor:
|
||||
"""Write ``ids`` into the device buffer and return ``[B]``.
|
||||
|
||||
Host values are staged through the double buffer and copied into the
|
||||
stable device buffer (``copy_`` without pinning is synchronous, so
|
||||
the alternating buffers guard against an in-flight transfer).
|
||||
"""
|
||||
b = len(ids)
|
||||
pin = self._pin[self._pin_idx]
|
||||
self._pin_idx ^= 1
|
||||
for i, v in enumerate(ids):
|
||||
pin[i] = v
|
||||
self.input_ids[:b].copy_(pin[:b])
|
||||
return self.input_ids[:b]
|
||||
|
||||
def decode_mask(self, position_ids: Tensor, total_len: int) -> Tensor:
|
||||
"""Return the ``[B, 1, total_len]`` validity mask for this step.
|
||||
|
||||
Written into the pre-allocated buffer via ``torch.ge(out=)`` — no
|
||||
new tensor is allocated. ``position_ids`` is the current step's
|
||||
``[B]`` positions; ``total_len`` must not exceed ``max_seq_len``.
|
||||
"""
|
||||
b = position_ids.size(0)
|
||||
out = self.input_mask[:b, :, :total_len]
|
||||
torch.ge(position_ids[:, None, None], self.arange[:total_len], out=out)
|
||||
return out
|
||||
@@ -0,0 +1,27 @@
|
||||
import logging
|
||||
import os
|
||||
|
||||
|
||||
def setup_logging(level: str = "INFO"):
|
||||
"""Attach a StreamHandler to the ``astrai`` logger (idempotent).
|
||||
|
||||
Call once per process at the top of CLI scripts.
|
||||
Set ``ASTR_LOG_LEVEL`` env var to override the default level.
|
||||
|
||||
Level names: ``DEBUG``, ``INFO``, ``WARNING``, ``ERROR``, ``CRITICAL``.
|
||||
``DEBUG`` enables per-step prefill/decode timing logs
|
||||
(:func:`astrai.inference.runtime.executor.timed`).
|
||||
"""
|
||||
logger = logging.getLogger("astrai")
|
||||
if logger.handlers:
|
||||
return
|
||||
level_name = os.environ.get("ASTR_LOG_LEVEL", level).upper()
|
||||
logger.setLevel(getattr(logging, level_name, logging.INFO))
|
||||
handler = logging.StreamHandler()
|
||||
handler.setFormatter(
|
||||
logging.Formatter(
|
||||
"%(asctime)s | %(levelname)-7s | %(name)s | %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
)
|
||||
)
|
||||
logger.addHandler(handler)
|
||||
@@ -1,21 +1,35 @@
|
||||
from astrai.model.automodel import AutoModel
|
||||
from astrai.model.module import (
|
||||
GQA,
|
||||
MLP,
|
||||
DecoderBlock,
|
||||
Linear,
|
||||
RMSNorm,
|
||||
from astrai.model.components.attention import GQA
|
||||
from astrai.model.components.decoder_block import DecoderBlock
|
||||
from astrai.model.components.linear import Linear
|
||||
from astrai.model.components.lora import (
|
||||
LoRAConfig,
|
||||
inject_lora,
|
||||
load_lora,
|
||||
merge_lora,
|
||||
save_lora,
|
||||
)
|
||||
from astrai.model.transformer import Transformer
|
||||
from astrai.model.components.mlp import MLP, DeepSeekMoE
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
from astrai.model.encoder import EmbeddingEncoder
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
|
||||
__all__ = [
|
||||
# Modules
|
||||
"Linear",
|
||||
"RMSNorm",
|
||||
"MLP",
|
||||
"DeepSeekMoE",
|
||||
"GQA",
|
||||
"DecoderBlock",
|
||||
# Models
|
||||
"Transformer",
|
||||
"AutoRegressiveLM",
|
||||
"EmbeddingEncoder",
|
||||
"AutoModel",
|
||||
# LoRA
|
||||
"LoRAConfig",
|
||||
"inject_lora",
|
||||
"merge_lora",
|
||||
"save_lora",
|
||||
"load_lora",
|
||||
]
|
||||
|
||||
+34
-67
@@ -4,18 +4,22 @@ AutoModel base class for model loading and saving.
|
||||
|
||||
from contextlib import contextmanager
|
||||
from pathlib import Path
|
||||
from typing import Self, Type, Union
|
||||
from typing import Self, Union
|
||||
|
||||
import safetensors.torch as st
|
||||
import torch.nn as nn
|
||||
|
||||
from astrai.config import ModelConfig
|
||||
from astrai.factory import Registry
|
||||
from astrai.config.model_config import BaseModelConfig, ConfigFactory
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.serialization import load_model_config, load_model_weights, save_model
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _disable_random_init(enable: bool = True):
|
||||
init_functions = [
|
||||
if not enable:
|
||||
yield
|
||||
return
|
||||
|
||||
names = (
|
||||
"xavier_normal_",
|
||||
"xavier_uniform_",
|
||||
"kaiming_normal_",
|
||||
@@ -25,60 +29,28 @@ def _disable_random_init(enable: bool = True):
|
||||
"constant_",
|
||||
"normal_",
|
||||
"uniform_",
|
||||
]
|
||||
original_funcs = {}
|
||||
for name in init_functions:
|
||||
if enable and hasattr(nn.init, name):
|
||||
original_funcs[name] = getattr(nn.init, name)
|
||||
setattr(nn.init, name, lambda *args, **kwargs: None)
|
||||
)
|
||||
orig = {n: getattr(nn.init, n) for n in names if hasattr(nn.init, n)}
|
||||
for n in orig:
|
||||
setattr(nn.init, n, lambda *a, **kw: None)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
if enable:
|
||||
for name, orig_func in original_funcs.items():
|
||||
setattr(nn.init, name, orig_func)
|
||||
for n, fn in orig.items():
|
||||
setattr(nn.init, n, fn)
|
||||
|
||||
|
||||
class ModelFactory(BaseFactory[nn.Module]):
|
||||
"""Pure factory for model dispatch, separated from nn.Module state."""
|
||||
|
||||
|
||||
class AutoModel(nn.Module):
|
||||
"""
|
||||
Autoregressive language model base class.
|
||||
Provides model loading/saving and generation capabilities.
|
||||
"""
|
||||
"""Model base class with loading/saving and generation."""
|
||||
|
||||
_registry = Registry()
|
||||
|
||||
def __init__(self, config: ModelConfig):
|
||||
def __init__(self, config: BaseModelConfig):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
|
||||
@classmethod
|
||||
def register(cls, model_type: str):
|
||||
"""
|
||||
Class method decorator to register model type.
|
||||
|
||||
Usage:
|
||||
@AutoModel.register('transformer')
|
||||
class Transformer(AutoModel):
|
||||
...
|
||||
"""
|
||||
|
||||
def decorator(sub_cls: Type["AutoModel"]) -> Type["AutoModel"]:
|
||||
cls._registry.register(model_type.lower(), sub_cls)
|
||||
return sub_cls
|
||||
|
||||
return decorator
|
||||
|
||||
@classmethod
|
||||
def get_model_class(cls, model_type: str) -> Type["AutoModel"]:
|
||||
"""Get model class by model_type string."""
|
||||
model_type = model_type.lower()
|
||||
if not cls._registry.contains(model_type):
|
||||
available = cls._registry.list_names()
|
||||
raise ValueError(
|
||||
f"Unknown model_type: {model_type}. Available: {available}"
|
||||
)
|
||||
return cls._registry.get(model_type)
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(
|
||||
cls,
|
||||
@@ -89,24 +61,22 @@ class AutoModel(nn.Module):
|
||||
|
||||
model_path = Path(path)
|
||||
|
||||
# Load config
|
||||
config = ModelConfig()
|
||||
config_path = model_path / "config.json"
|
||||
if config_path.exists():
|
||||
config.load(str(config_path))
|
||||
else:
|
||||
if not config_path.exists():
|
||||
raise FileNotFoundError(f"Config file not found: {config_path}")
|
||||
|
||||
model_type = config.model_type or "transformer"
|
||||
actual_cls = cls.get_model_class(model_type)
|
||||
raw = load_model_config(str(model_path))
|
||||
config = ConfigFactory.load(raw)
|
||||
model_type = config.model_type or "autoregressive_lm"
|
||||
|
||||
actual_cls = ModelFactory.get_component_class(model_type)
|
||||
|
||||
with _disable_random_init(enable=disable_random_init):
|
||||
model = actual_cls(config)
|
||||
|
||||
# Load weights
|
||||
weights_path = model_path / "model.safetensors"
|
||||
if weights_path.exists():
|
||||
state_dict = st.load_file(str(weights_path))
|
||||
state_dict = load_model_weights(str(model_path))
|
||||
model.load_state_dict(state_dict, strict=strict)
|
||||
|
||||
return model
|
||||
@@ -114,15 +84,12 @@ class AutoModel(nn.Module):
|
||||
def save_pretrained(
|
||||
self,
|
||||
save_directory: Union[str, Path],
|
||||
) -> None:
|
||||
save_path = Path(save_directory)
|
||||
save_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Save config
|
||||
self.config.save(str(save_path / "config.json"))
|
||||
|
||||
# Save weights
|
||||
st.save_file(self.state_dict(), str(save_path / "model.safetensors"))
|
||||
):
|
||||
save_model(
|
||||
config=self.config.to_dict(),
|
||||
state_dict=self.state_dict(),
|
||||
save_directory=str(save_directory),
|
||||
)
|
||||
|
||||
def to(self, *args, **kwargs) -> Self:
|
||||
"""Move model to device/dtype."""
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
from astrai.extension.backend.rotary import apply_rotary_emb
|
||||
from astrai.model.components.attention import GQA, MLA
|
||||
from astrai.model.components.decoder_block import DecoderBlock
|
||||
from astrai.model.components.embedding import Embedding
|
||||
from astrai.model.components.linear import Linear
|
||||
from astrai.model.components.mlp import MLP, DeepSeekMoE
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
from astrai.model.components.rope import (
|
||||
RotaryEmbedding,
|
||||
get_rotary_emb,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Linear",
|
||||
"RMSNorm",
|
||||
"MLP",
|
||||
"DeepSeekMoE",
|
||||
"Embedding",
|
||||
"GQA",
|
||||
"MLA",
|
||||
"DecoderBlock",
|
||||
"RotaryEmbedding",
|
||||
"apply_rotary_emb",
|
||||
"get_rotary_emb",
|
||||
]
|
||||
@@ -0,0 +1,181 @@
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.extension.backend import apply_rotary_emb, attention
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.inference.cache import KVCache
|
||||
from astrai.model.components.linear import Linear
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
|
||||
|
||||
class AttnFactory(BaseFactory[nn.Module]):
|
||||
pass
|
||||
|
||||
|
||||
@AttnFactory.register("gqa")
|
||||
class GQA(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
n_heads: int,
|
||||
n_kv_heads: int,
|
||||
use_qk_norm: bool,
|
||||
norm_eps: float,
|
||||
use_gated_attention: bool,
|
||||
layer_id: int,
|
||||
n_layers: int = 1,
|
||||
):
|
||||
super().__init__()
|
||||
assert dim % n_heads == 0
|
||||
assert n_heads % n_kv_heads == 0
|
||||
|
||||
self.head_dim = dim // n_heads
|
||||
self.layer_id = layer_id
|
||||
self.dim = dim
|
||||
self.n_heads = n_heads
|
||||
self.n_kv_heads = n_kv_heads
|
||||
self.n_rep = n_heads // n_kv_heads
|
||||
self.use_qk_norm = use_qk_norm
|
||||
self.use_gated_attention = use_gated_attention
|
||||
|
||||
self.q_proj = Linear(dim, n_heads * self.head_dim)
|
||||
self.k_proj = Linear(dim, n_kv_heads * self.head_dim)
|
||||
self.v_proj = Linear(dim, n_kv_heads * self.head_dim)
|
||||
self.o_proj = Linear(dim, dim, init_std=0.02 / (2 * n_layers) ** 0.5)
|
||||
|
||||
if self.use_qk_norm:
|
||||
self.q_norm = RMSNorm(self.head_dim, norm_eps)
|
||||
self.k_norm = RMSNorm(self.head_dim, norm_eps)
|
||||
|
||||
if self.use_gated_attention:
|
||||
self.gate = Linear(dim, dim)
|
||||
|
||||
def _split_heads(self, x: Tensor, n_heads) -> Tensor:
|
||||
return x.reshape(*x.shape[:-1], n_heads, self.head_dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
rotary_emb: Tensor,
|
||||
attn_mask: Tensor = None,
|
||||
kv_cache: Optional[KVCache] = None,
|
||||
is_causal: bool = False,
|
||||
fwd: Optional[str] = None,
|
||||
) -> Tensor:
|
||||
q = self._split_heads(self.q_proj(x), self.n_heads)
|
||||
k = self._split_heads(self.k_proj(x), self.n_kv_heads)
|
||||
v = self._split_heads(self.v_proj(x), self.n_kv_heads)
|
||||
q, k = apply_rotary_emb(q, rotary_emb), apply_rotary_emb(k, rotary_emb)
|
||||
|
||||
if self.use_qk_norm:
|
||||
q, k = self.q_norm(q), self.k_norm(k)
|
||||
|
||||
sdqa_out = attention(
|
||||
q, k, v, kv_cache, self.layer_id, attn_mask, is_causal, fwd
|
||||
).reshape(*x.shape[:-1], self.dim)
|
||||
|
||||
if self.use_gated_attention:
|
||||
sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
|
||||
|
||||
out = self.o_proj(sdqa_out)
|
||||
return out
|
||||
|
||||
|
||||
@AttnFactory.register("mla")
|
||||
class MLA(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
n_heads: int,
|
||||
n_kv_heads: int,
|
||||
kv_lora_rank: int,
|
||||
qk_nope_head_dim: int,
|
||||
qk_rope_head_dim: int,
|
||||
norm_eps: float,
|
||||
use_qk_norm: bool,
|
||||
use_gated_attention: bool,
|
||||
layer_id: int,
|
||||
n_layers: int = 1,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.n_heads = n_heads
|
||||
self.n_kv_heads = n_kv_heads
|
||||
self.kv_lora_rank = kv_lora_rank
|
||||
self.qk_nope_head_dim = qk_nope_head_dim
|
||||
self.qk_rope_head_dim = qk_rope_head_dim
|
||||
self.head_dim = qk_nope_head_dim + qk_rope_head_dim
|
||||
self.layer_id = layer_id
|
||||
self.n_rep = n_heads // n_kv_heads
|
||||
self.use_qk_norm = use_qk_norm
|
||||
self.use_gated_attention = use_gated_attention
|
||||
|
||||
self.q_proj = Linear(dim, n_heads * self.head_dim, bias=False)
|
||||
|
||||
if self.use_qk_norm:
|
||||
self.q_norm = RMSNorm(self.head_dim, norm_eps)
|
||||
self.k_norm = RMSNorm(self.head_dim, norm_eps)
|
||||
self.kv_a_proj = Linear(dim, kv_lora_rank, bias=False)
|
||||
self.kv_norm = RMSNorm(kv_lora_rank, norm_eps)
|
||||
|
||||
self.kv_b_proj = Linear(
|
||||
kv_lora_rank,
|
||||
n_kv_heads * (2 * self.head_dim),
|
||||
)
|
||||
|
||||
self.o_proj = Linear(
|
||||
dim, dim, bias=False, init_std=0.02 / (2 * n_layers) ** 0.5
|
||||
)
|
||||
|
||||
if use_gated_attention:
|
||||
self.gate = Linear(dim, dim, bias=False)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
rotary_emb: Tensor,
|
||||
attn_mask: Tensor = None,
|
||||
kv_cache: Optional[KVCache] = None,
|
||||
is_causal: bool = False,
|
||||
fwd: Optional[str] = None,
|
||||
) -> Tensor:
|
||||
q = self.q_proj(x)
|
||||
q = q.reshape(*x.shape[:-1], self.n_heads, self.head_dim)
|
||||
|
||||
kv_compressed = self.kv_a_proj(x)
|
||||
kv_compressed = self.kv_norm(kv_compressed)
|
||||
|
||||
kv = self.kv_b_proj(kv_compressed)
|
||||
kv = kv.reshape(*x.shape[:-1], self.n_kv_heads, -1)
|
||||
|
||||
k_nope, k_rope, v = torch.split(
|
||||
kv, [self.qk_nope_head_dim, self.qk_rope_head_dim, self.head_dim], dim=-1
|
||||
)
|
||||
|
||||
q_nope, q_rope = (
|
||||
q[..., : self.qk_nope_head_dim],
|
||||
q[..., self.qk_nope_head_dim :],
|
||||
)
|
||||
q_rope = apply_rotary_emb(q_rope, rotary_emb)
|
||||
k_rope = apply_rotary_emb(k_rope, rotary_emb)
|
||||
|
||||
q = torch.cat([q_nope, q_rope], dim=-1)
|
||||
k = torch.cat([k_nope, k_rope], dim=-1)
|
||||
|
||||
if self.use_qk_norm:
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
|
||||
attn_out = attention(
|
||||
q, k, v, kv_cache, self.layer_id, attn_mask, is_causal, fwd
|
||||
).reshape(*x.shape[:-1], self.dim)
|
||||
|
||||
if self.use_gated_attention:
|
||||
attn_out = attn_out * F.sigmoid(self.gate(x))
|
||||
|
||||
out = self.o_proj(attn_out)
|
||||
return out
|
||||
@@ -0,0 +1,76 @@
|
||||
from dataclasses import asdict
|
||||
from typing import Optional, TypedDict
|
||||
|
||||
import torch.nn as nn
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.inference.cache import KVCache
|
||||
from astrai.model.components.attention import AttnFactory
|
||||
from astrai.model.components.mlp import FFNFactory, RouterStats
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
|
||||
|
||||
class DecoderOutput(TypedDict):
|
||||
hidden_states: Tensor
|
||||
aux_loss: Optional[Tensor]
|
||||
router_stats: Optional[RouterStats]
|
||||
|
||||
|
||||
class DecoderBlock(nn.Module):
|
||||
def __init__(self, config, layer_id: int):
|
||||
super().__init__()
|
||||
cfg = asdict(config)
|
||||
cfg.update(
|
||||
dim=config.hidden_size,
|
||||
dim_ffn=config.intermediate_size,
|
||||
n_layers=config.num_hidden_layers,
|
||||
n_heads=config.num_attention_heads,
|
||||
n_kv_heads=config.num_key_value_heads,
|
||||
norm_eps=config.rms_norm_eps,
|
||||
down_init_std=0.02 / (2 * config.num_hidden_layers) ** 0.5,
|
||||
)
|
||||
self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id)
|
||||
self.input_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||
self.post_attention_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||
ffn_type = self._resolve_ffn_type(config, layer_id)
|
||||
self.mlp = FFNFactory.create(ffn_type, **cfg)
|
||||
|
||||
@staticmethod
|
||||
def _resolve_ffn_type(config, layer_id: int) -> str:
|
||||
if config.ffn_type != "moe":
|
||||
return config.ffn_type
|
||||
mlp_only = config.mlp_only_layers or []
|
||||
if layer_id in mlp_only:
|
||||
return "mlp"
|
||||
if config.decoder_sparse_step > 1:
|
||||
if (layer_id + 1) % config.decoder_sparse_step != 0:
|
||||
return "mlp"
|
||||
return "moe"
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
rotary_emb: Tensor,
|
||||
attention_mask: Optional[Tensor] = None,
|
||||
kv_cache: Optional[KVCache] = None,
|
||||
is_causal: bool = False,
|
||||
fwd: Optional[str] = None,
|
||||
) -> DecoderOutput:
|
||||
attn_output = self.attention(
|
||||
self.input_norm(x),
|
||||
rotary_emb,
|
||||
attention_mask,
|
||||
kv_cache,
|
||||
is_causal,
|
||||
fwd,
|
||||
)
|
||||
x = attn_output + x
|
||||
normalized = self.post_attention_norm(x)
|
||||
mlp_output = self.mlp(normalized)
|
||||
x = mlp_output["hidden_states"] + x
|
||||
|
||||
return {
|
||||
"hidden_states": x,
|
||||
"aux_loss": mlp_output["aux_loss"],
|
||||
"router_stats": mlp_output.get("router_stats"),
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
class Embedding(nn.Module):
|
||||
def __init__(self, vocab_size: int, embedding_dim: int, neftune_alpha: float = 0.0):
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.empty((vocab_size, embedding_dim)))
|
||||
self.neftune_noise_alpha = neftune_alpha
|
||||
|
||||
def set_neftune_alpha(self, alpha: float):
|
||||
self.neftune_noise_alpha = alpha
|
||||
|
||||
def reset_parameters(self):
|
||||
nn.init.normal_(self.weight, mean=0.0, std=0.02)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
out = F.embedding(x, self.weight)
|
||||
if self.training and self.neftune_noise_alpha > 0.0:
|
||||
eps = self.neftune_noise_alpha / math.sqrt(out.size(1))
|
||||
out = out + eps * torch.randn_like(out)
|
||||
return out
|
||||
@@ -0,0 +1,24 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
class Linear(nn.Module):
|
||||
def __init__(
|
||||
self, in_dim: int, out_dim: int, bias: bool = False, init_std: float = 0.02
|
||||
):
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.empty((out_dim, in_dim)))
|
||||
self.bias = nn.Parameter(torch.zeros(out_dim)) if bias else None
|
||||
self.init_std = init_std
|
||||
|
||||
def reset_parameters(self):
|
||||
nn.init.normal_(self.weight, mean=0.0, std=self.init_std)
|
||||
if self.bias is not None:
|
||||
fan_in, _ = nn.init._calculate_fan_in_and_fan_out(self.weight)
|
||||
bound = 1 / (fan_in**0.5)
|
||||
nn.init.uniform_(self.bias, -bound, bound)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return F.linear(x, self.weight, self.bias)
|
||||
@@ -0,0 +1,199 @@
|
||||
import logging
|
||||
from dataclasses import asdict
|
||||
from pathlib import Path
|
||||
from typing import Optional, Set
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from pydantic.dataclasses import dataclass
|
||||
|
||||
from astrai.model.components.linear import Linear
|
||||
from astrai.serialization import (
|
||||
load_json,
|
||||
load_safetensors,
|
||||
save_json,
|
||||
save_safetensors,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
TARGET_MODULES_ATTN = {"q_proj", "k_proj", "v_proj", "o_proj"}
|
||||
TARGET_MODULES_FFN = {"up", "gate", "down"}
|
||||
|
||||
|
||||
@dataclass
|
||||
class LoRAConfig:
|
||||
r: int = 16
|
||||
alpha: int = 32
|
||||
target_modules: tuple = ("q_proj", "v_proj")
|
||||
|
||||
|
||||
class LoRALinear(nn.Module):
|
||||
def __init__(self, base: Linear, r: int = 16, alpha: int = 32):
|
||||
super().__init__()
|
||||
self.register_parameter("weight", base.weight)
|
||||
self.weight.requires_grad_(False)
|
||||
self.bias = base.bias
|
||||
if self.bias is not None:
|
||||
self.bias.requires_grad_(False)
|
||||
|
||||
self.r = r
|
||||
self.scaling = alpha / r
|
||||
device = self.weight.device
|
||||
dtype = self.weight.dtype
|
||||
lora_a = torch.randn(r, self.weight.shape[1], device=device, dtype=dtype) / r
|
||||
lora_b = torch.zeros(self.weight.shape[0], r, device=device, dtype=dtype)
|
||||
self.lora_A = nn.Parameter(lora_a)
|
||||
self.lora_B = nn.Parameter(lora_b)
|
||||
self._merged = False
|
||||
|
||||
def forward(self, x):
|
||||
out = F.linear(x, self.weight, self.bias)
|
||||
if not self._merged:
|
||||
out += (F.linear(x, self.lora_A) @ self.lora_B.T) * self.scaling
|
||||
return out
|
||||
|
||||
def merge(self):
|
||||
if self._merged:
|
||||
return
|
||||
self.weight.data += (self.lora_B @ self.lora_A) * self.scaling
|
||||
self._merged = True
|
||||
del self.lora_A
|
||||
del self.lora_B
|
||||
|
||||
|
||||
def _collect_lora_info(model: nn.Module) -> dict:
|
||||
names = {}
|
||||
for n, m in model.named_modules():
|
||||
if isinstance(m, Linear):
|
||||
_, _, child = n.rpartition(".")
|
||||
names.setdefault(child, []).append(n)
|
||||
return names
|
||||
|
||||
|
||||
def _get_lora_count(model: nn.Module) -> int:
|
||||
return sum(1 for m in model.modules() if isinstance(m, LoRALinear))
|
||||
|
||||
|
||||
def inject_lora(
|
||||
model: nn.Module,
|
||||
r: int = 16,
|
||||
alpha: int = 32,
|
||||
target_modules: Optional[Set[str]] = None,
|
||||
) -> LoRAConfig:
|
||||
if target_modules is None:
|
||||
target_modules = TARGET_MODULES_ATTN
|
||||
|
||||
available = _collect_lora_info(model)
|
||||
injected = 0
|
||||
|
||||
for name, module in list(model.named_modules()):
|
||||
if not isinstance(module, Linear):
|
||||
continue
|
||||
parent_name, _, child_name = name.rpartition(".")
|
||||
if child_name not in target_modules:
|
||||
continue
|
||||
parent = model.get_submodule(parent_name) if parent_name else model
|
||||
setattr(parent, child_name, LoRALinear(module, r=r, alpha=alpha))
|
||||
injected += 1
|
||||
|
||||
if injected == 0:
|
||||
logger.warning(
|
||||
"No LoRA layers injected. Available Linear child names: %s. "
|
||||
"target_modules: %s. Check model type and target_modules.",
|
||||
sorted(available),
|
||||
sorted(target_modules),
|
||||
)
|
||||
else:
|
||||
logger.info("LoRA injected: %d layers (r=%d, alpha=%d)", injected, r, alpha)
|
||||
|
||||
return LoRAConfig(r=r, alpha=alpha, target_modules=tuple(target_modules))
|
||||
|
||||
|
||||
def merge_lora(model: nn.Module):
|
||||
n = 0
|
||||
for module in model.modules():
|
||||
if isinstance(module, LoRALinear):
|
||||
module.merge()
|
||||
n += 1
|
||||
if n == 0:
|
||||
logger.warning("No LoRA layers to merge.")
|
||||
else:
|
||||
logger.info("Merged %d LoRA layers", n)
|
||||
|
||||
|
||||
def save_lora(model: nn.Module, save_dir: str, config: LoRAConfig):
|
||||
lora_sd = {
|
||||
k: v
|
||||
for k, v in model.state_dict().items()
|
||||
if k.endswith((".lora_A", ".lora_B"))
|
||||
}
|
||||
if not lora_sd:
|
||||
raise RuntimeError(
|
||||
"No LoRA parameters found in model. "
|
||||
"The model may not have been injected or was already merged."
|
||||
)
|
||||
|
||||
path = Path(save_dir)
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
save_safetensors(lora_sd, path / "adapter_model.safetensors")
|
||||
save_json(asdict(config), path / "adapter_config.json")
|
||||
logger.info("LoRA adapter saved to %s (%d keys)", save_dir, len(lora_sd))
|
||||
|
||||
|
||||
def load_lora(model: nn.Module, load_dir: str) -> LoRAConfig:
|
||||
path = Path(load_dir)
|
||||
raw = load_json(path / "adapter_config.json")
|
||||
config = LoRAConfig(
|
||||
r=raw["r"], alpha=raw["alpha"], target_modules=tuple(raw["target_modules"])
|
||||
)
|
||||
|
||||
existing = _get_lora_count(model)
|
||||
if existing > 0:
|
||||
logger.warning(
|
||||
"Model already has %d LoRA layers. Skipping injection, "
|
||||
"loading weights onto existing layers only.",
|
||||
existing,
|
||||
)
|
||||
else:
|
||||
inject_lora(
|
||||
model,
|
||||
r=config.r,
|
||||
alpha=config.alpha,
|
||||
target_modules=set(config.target_modules),
|
||||
)
|
||||
|
||||
weights = load_safetensors(path / "adapter_model.safetensors")
|
||||
try:
|
||||
missing, unexpected = model.load_state_dict(weights, strict=False)
|
||||
except RuntimeError as e:
|
||||
msg = str(e)
|
||||
if "size mismatch" in msg:
|
||||
raise RuntimeError(
|
||||
f"LoRA weight shapes do not match the model. "
|
||||
f"The adapter config (r={config.r}) may not match the injected layers. "
|
||||
f"Original error: {msg}"
|
||||
) from e
|
||||
raise
|
||||
|
||||
injected = _get_lora_count(model)
|
||||
if injected == 0:
|
||||
raise RuntimeError(
|
||||
"No LoRA layers found after loading. "
|
||||
"Inject LoRA before calling load_lora, or check the adapter config."
|
||||
)
|
||||
|
||||
if missing:
|
||||
lora_missing = [k for k in missing if "lora" in k]
|
||||
if lora_missing:
|
||||
raise RuntimeError(
|
||||
f"LoRA weight keys not found in model: {lora_missing}. "
|
||||
f"The adapter config (r={config.r}) may not match the model."
|
||||
)
|
||||
logger.debug("LoRA load: %d missing base-weight keys (expected)", len(missing))
|
||||
if unexpected:
|
||||
logger.warning("LoRA load: %d unexpected keys", len(unexpected))
|
||||
|
||||
logger.info("LoRA adapter loaded from %s", load_dir)
|
||||
return config
|
||||
@@ -0,0 +1,178 @@
|
||||
from typing import Optional, TypedDict
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.model.components.linear import Linear
|
||||
|
||||
|
||||
class FFNFactory(BaseFactory[nn.Module]):
|
||||
pass
|
||||
|
||||
|
||||
class RouterStats(TypedDict):
|
||||
"""Per-layer MoE routing statistics for training diagnostics.
|
||||
|
||||
Both tensors are detached monitoring data produced during forward.
|
||||
"""
|
||||
|
||||
probs: Tensor
|
||||
topk_indices: Tensor
|
||||
|
||||
|
||||
class FFNOutput(TypedDict):
|
||||
hidden_states: Tensor
|
||||
aux_loss: Optional[Tensor]
|
||||
router_stats: Optional[RouterStats]
|
||||
|
||||
|
||||
class RoutedOutput(TypedDict):
|
||||
hidden_states: Tensor
|
||||
aux_loss: Optional[Tensor]
|
||||
router_stats: Optional[RouterStats]
|
||||
|
||||
|
||||
@FFNFactory.register("mlp")
|
||||
class MLP(nn.Module):
|
||||
def __init__(self, dim: int, dim_ffn: int, down_init_std: float = 0.02):
|
||||
super().__init__()
|
||||
self.up = Linear(dim, dim_ffn)
|
||||
self.gate = Linear(dim, dim_ffn)
|
||||
self.down = Linear(dim_ffn, dim, init_std=down_init_std)
|
||||
|
||||
def forward(self, x: Tensor) -> FFNOutput:
|
||||
gated = self.up(x) * F.silu(self.gate(x))
|
||||
out = self.down(gated)
|
||||
return {"hidden_states": out, "aux_loss": None, "router_stats": None}
|
||||
|
||||
|
||||
@FFNFactory.register("moe")
|
||||
class DeepSeekMoE(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
dim_ffn: int,
|
||||
n_routed_experts: int,
|
||||
n_shared_experts: int = 1,
|
||||
n_activated_experts: int = 2,
|
||||
topk_method: str = "greedy",
|
||||
n_layers: int = 1,
|
||||
moe_intermediate_size: Optional[int] = None,
|
||||
shared_expert_intermediate_size: Optional[int] = None,
|
||||
norm_topk_prob: bool = True,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.n_routed_experts = n_routed_experts
|
||||
self.n_shared_experts = n_shared_experts
|
||||
self.n_activated_experts = n_activated_experts
|
||||
self.topk_method = topk_method
|
||||
self.norm_topk_prob = norm_topk_prob
|
||||
|
||||
expert_dim_ffn = (
|
||||
moe_intermediate_size if moe_intermediate_size is not None else dim_ffn
|
||||
)
|
||||
shared_dim_ffn = (
|
||||
shared_expert_intermediate_size
|
||||
if shared_expert_intermediate_size is not None
|
||||
else dim_ffn
|
||||
)
|
||||
|
||||
self.router = Linear(dim, n_routed_experts, bias=False)
|
||||
moe_scale = 1 / max(n_shared_experts, 1) + 1 / n_activated_experts
|
||||
down_init_std = 0.02 / (2 * n_layers * moe_scale) ** 0.5
|
||||
|
||||
self.shared_experts = nn.ModuleList(
|
||||
[
|
||||
MLP(dim, shared_dim_ffn, down_init_std=down_init_std)
|
||||
for _ in range(n_shared_experts)
|
||||
]
|
||||
)
|
||||
self.routed_experts = nn.ModuleList(
|
||||
[
|
||||
MLP(dim, expert_dim_ffn, down_init_std=down_init_std)
|
||||
for _ in range(n_routed_experts)
|
||||
]
|
||||
)
|
||||
|
||||
def forward(self, x: Tensor) -> FFNOutput:
|
||||
include_aux_loss = self.training and torch.is_grad_enabled()
|
||||
shape = x.shape
|
||||
dim = shape[-1]
|
||||
x_flat = x.view(-1, dim)
|
||||
|
||||
shared_out = self._shared_forward(x_flat)
|
||||
routed_output = self._routed_forward(x_flat, include_aux_loss)
|
||||
|
||||
out = (shared_out + routed_output["hidden_states"]).view(shape)
|
||||
return {
|
||||
"hidden_states": out,
|
||||
"aux_loss": routed_output["aux_loss"],
|
||||
"router_stats": routed_output["router_stats"],
|
||||
}
|
||||
|
||||
def _shared_forward(self, x: Tensor) -> Tensor:
|
||||
if self.n_shared_experts == 0:
|
||||
return torch.zeros_like(x)
|
||||
return (
|
||||
sum(e(x)["hidden_states"] for e in self.shared_experts)
|
||||
/ self.n_shared_experts
|
||||
)
|
||||
|
||||
def _routed_forward(self, x: Tensor, include_aux_loss: bool) -> RoutedOutput:
|
||||
N, D = x.shape
|
||||
K = self.n_activated_experts
|
||||
E = self.n_routed_experts
|
||||
|
||||
router_logits = self.router(x)
|
||||
router_probs = torch.softmax(router_logits.float(), dim=-1).to(x.dtype)
|
||||
|
||||
topk_weights, topk_indices = torch.topk(router_probs, K, dim=-1, sorted=False)
|
||||
if self.norm_topk_prob:
|
||||
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
||||
|
||||
aux_loss = None
|
||||
router_stats = None
|
||||
if include_aux_loss:
|
||||
expert_load = F.one_hot(topk_indices, num_classes=E).float()
|
||||
expert_load = expert_load.mean(dim=(0, 1))
|
||||
router_prob = router_probs.float().mean(dim=0)
|
||||
aux_loss = E * (expert_load * router_prob).sum()
|
||||
router_stats = {
|
||||
"probs": router_probs.detach(),
|
||||
"topk_indices": topk_indices,
|
||||
}
|
||||
|
||||
# Grouped dispatch: sort (token, slot) pairs by expert so each expert
|
||||
# consumes one contiguous slice instead of a per-expert mask scan.
|
||||
flat_experts = topk_indices.reshape(-1)
|
||||
sorted_experts, order = torch.sort(flat_experts)
|
||||
flat_tokens = x.repeat_interleave(K, dim=0)[order]
|
||||
flat_weights = topk_weights.reshape(-1, 1)[order]
|
||||
boundaries = torch.cumsum(
|
||||
torch.bincount(sorted_experts, minlength=E), dim=0
|
||||
).tolist()
|
||||
|
||||
output = torch.zeros(N, D, device=x.device, dtype=x.dtype)
|
||||
start = 0
|
||||
for expert_idx, end in enumerate(boundaries):
|
||||
if end == start:
|
||||
continue
|
||||
expert_output = self.routed_experts[expert_idx](flat_tokens[start:end])[
|
||||
"hidden_states"
|
||||
]
|
||||
output.index_add_(
|
||||
0,
|
||||
order[start:end] // K,
|
||||
expert_output * flat_weights[start:end],
|
||||
)
|
||||
start = end
|
||||
|
||||
return {
|
||||
"hidden_states": output,
|
||||
"aux_loss": aux_loss,
|
||||
"router_stats": router_stats,
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
def __init__(self, dim, norm_eps):
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
self.normalized_shape = (dim,)
|
||||
self.norm_eps = norm_eps
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return F.rms_norm(x, self.normalized_shape, self.weight, self.norm_eps)
|
||||
@@ -0,0 +1,76 @@
|
||||
from typing import Dict, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
def get_rotary_emb(
|
||||
dim: int,
|
||||
max_len: int,
|
||||
base: float = 10000,
|
||||
device: Optional[torch.device] = None,
|
||||
) -> Tensor:
|
||||
"""Precompute cos/sin tables for rotary embedding.
|
||||
|
||||
Returns:
|
||||
[max_len, dim/2, 2] (f32) — [cos, sin] pairs.
|
||||
"""
|
||||
theta = base ** (-torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim)
|
||||
t = torch.arange(0, max_len, dtype=torch.float64, device=device)
|
||||
freqs = torch.outer(t, theta).float()
|
||||
cos = torch.cos(freqs)
|
||||
sin = torch.sin(freqs)
|
||||
return torch.stack([cos, sin], dim=-1)
|
||||
|
||||
|
||||
def ntk_base(base: float, dim: int, factor: float) -> float:
|
||||
return base * (factor ** (dim / (dim - 2)))
|
||||
|
||||
|
||||
class RotaryEmbedding(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
max_len: int,
|
||||
base: float = 10000,
|
||||
rope_scaling: Optional[Dict] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.max_len = max_len
|
||||
self.base = base
|
||||
self.rope_scaling = rope_scaling
|
||||
|
||||
if rope_scaling is not None:
|
||||
scaling_type = rope_scaling.get("type", "ntk")
|
||||
factor = rope_scaling.get("factor", 1.0)
|
||||
if scaling_type == "ntk":
|
||||
self.base = ntk_base(base, dim, factor)
|
||||
|
||||
self._set_rotary_buffer(self.max_len)
|
||||
|
||||
def _set_rotary_buffer(self, max_len: int):
|
||||
freqs_cis = get_rotary_emb(self.dim, max_len, self.base)
|
||||
self.register_buffer("freqs_cis", freqs_cis, persistent=False)
|
||||
|
||||
def forward(self, x: Tensor, position_ids: Optional[Tensor] = None) -> Tensor:
|
||||
"""Lookup cos/sin for the given positions.
|
||||
|
||||
Args:
|
||||
x: [batch, seq_len, ...] — only batch and seq_len are used.
|
||||
position_ids: [batch, seq_len] optional position indices.
|
||||
|
||||
Returns:
|
||||
[batch, seq_len, dim/2, 2] (f32) — [cos, sin] pairs.
|
||||
"""
|
||||
if position_ids is None:
|
||||
if x.ndim == 2:
|
||||
position_ids = torch.arange(x.size(0), device=x.device)
|
||||
else:
|
||||
position_ids = (
|
||||
torch.arange(x.size(1), device=x.device)
|
||||
.unsqueeze(0)
|
||||
.expand(x.size(0), -1)
|
||||
)
|
||||
return self.freqs_cis[position_ids].float()
|
||||
@@ -0,0 +1,97 @@
|
||||
from typing import Any, Mapping, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.config.model_config import EncoderConfig
|
||||
from astrai.model.automodel import AutoModel, ModelFactory
|
||||
from astrai.model.components.decoder_block import DecoderBlock
|
||||
from astrai.model.components.embedding import Embedding
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
from astrai.model.components.rope import RotaryEmbedding
|
||||
from astrai.model.transformer import process_attention_mask
|
||||
|
||||
|
||||
@ModelFactory.register("embedding")
|
||||
class EmbeddingEncoder(AutoModel):
|
||||
def __init__(self, config: EncoderConfig):
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
rope_dim = config.hidden_size // config.num_attention_heads
|
||||
rope_base = config.rope_theta if config.rope_theta is not None else 10000
|
||||
self.rotary_embedding = RotaryEmbedding(
|
||||
rope_dim,
|
||||
config.max_position_embeddings,
|
||||
rope_base,
|
||||
rope_scaling=config.rope_scaling,
|
||||
)
|
||||
self.embed_tokens = Embedding(
|
||||
config.vocab_size,
|
||||
config.hidden_size,
|
||||
neftune_alpha=config.neftune_alpha,
|
||||
)
|
||||
|
||||
self.layers = nn.ModuleList(
|
||||
[
|
||||
DecoderBlock(config, layer_id)
|
||||
for layer_id in range(config.num_hidden_layers)
|
||||
]
|
||||
)
|
||||
|
||||
self.norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||
|
||||
self.pooling_type = config.pooling_type or "mean"
|
||||
self.normalize_embeddings = config.normalize_embeddings or False
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, module):
|
||||
if hasattr(module, "reset_parameters"):
|
||||
module.reset_parameters()
|
||||
|
||||
def load_state_dict(self, state_dict: Mapping[str, Any], strict=True, assign=False):
|
||||
state_dict = dict(state_dict)
|
||||
state_dict.pop("lm_head.weight", None)
|
||||
return super().load_state_dict(state_dict, strict=strict, assign=assign)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Tensor,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
position_ids: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
assert input_ids.ndim == 2
|
||||
B, S = input_ids.shape
|
||||
|
||||
x = self.embed_tokens(input_ids)
|
||||
|
||||
rotary_emb = self.rotary_embedding(x, position_ids)
|
||||
attn_mask = process_attention_mask(input_mask)
|
||||
|
||||
for layer in self.layers:
|
||||
x = layer(x, rotary_emb, attn_mask)["hidden_states"]
|
||||
|
||||
hidden_states = self.norm(x)
|
||||
|
||||
if self.pooling_type == "cls":
|
||||
pooled = hidden_states[:, 0]
|
||||
elif self.pooling_type == "last":
|
||||
if input_mask is not None:
|
||||
lengths = input_mask.sum(dim=1) - 1
|
||||
pooled = hidden_states[torch.arange(B, device=x.device), lengths]
|
||||
else:
|
||||
pooled = hidden_states[:, -1]
|
||||
else:
|
||||
if input_mask is not None:
|
||||
mask = input_mask.unsqueeze(-1).to(dtype=hidden_states.dtype)
|
||||
pooled = (hidden_states * mask).sum(dim=1) / mask.sum(dim=1).clamp(
|
||||
min=1.0
|
||||
)
|
||||
else:
|
||||
pooled = hidden_states.mean(dim=1)
|
||||
|
||||
if self.normalize_embeddings:
|
||||
pooled = torch.nn.functional.normalize(pooled, p=2, dim=-1)
|
||||
|
||||
return pooled
|
||||
@@ -1,337 +0,0 @@
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.inference.cache import CacheView
|
||||
|
||||
|
||||
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
||||
"""Repeat KV heads n_rep times for GQA."""
|
||||
bs, slen, n_heads, head_dim = x.shape
|
||||
if n_rep == 1:
|
||||
return x
|
||||
return (
|
||||
x[:, :, :, None, :]
|
||||
.expand(bs, slen, n_heads, n_rep, head_dim)
|
||||
.reshape(bs, slen, n_heads * n_rep, head_dim)
|
||||
)
|
||||
|
||||
|
||||
def get_rotary_emb(
|
||||
dim: int,
|
||||
max_len: int,
|
||||
base: float = 10000,
|
||||
device: Optional[torch.device] = None,
|
||||
) -> Tuple[Tensor, Tensor]:
|
||||
"""Precompute cos/sin for RoPE."""
|
||||
theta = base ** (-torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim)
|
||||
t = torch.arange(0, max_len, dtype=torch.float64, device=device)
|
||||
freqs = torch.outer(t, theta)
|
||||
return torch.cos(freqs).float(), torch.sin(freqs).float()
|
||||
|
||||
|
||||
def apply_rotary_emb(x: torch.Tensor, rotary_emb: Tuple[Tensor, Tensor]) -> Tensor:
|
||||
"""Apply rotary embedding via cos/sin (shape-preserving)."""
|
||||
dtype = x.dtype
|
||||
cos, sin = rotary_emb
|
||||
cos = cos.unsqueeze(0).unsqueeze(2)
|
||||
sin = sin.unsqueeze(0).unsqueeze(2)
|
||||
x_real = x[..., 0::2]
|
||||
x_imag = x[..., 1::2]
|
||||
x_real_rot = x_real * cos - x_imag * sin
|
||||
x_imag_rot = x_real * sin + x_imag * cos
|
||||
x_out = torch.stack([x_real_rot, x_imag_rot], dim=-1)
|
||||
x_out = x_out.view(*x_out.shape[:-2], -1)
|
||||
return x_out.to(dtype)
|
||||
|
||||
|
||||
class RotaryEmbedding(nn.Module):
|
||||
def __init__(self, dim: int, max_len: int, base: int = 10000):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.max_len = max_len
|
||||
self.base = base
|
||||
self.max_len_cached = None
|
||||
self._set_rotary_buffer(self.max_len, None)
|
||||
|
||||
def _set_rotary_buffer(self, max_len: int, device: Optional[torch.device] = None):
|
||||
cos_cached, sin_cached = get_rotary_emb(self.dim, max_len, self.base, device)
|
||||
self.register_buffer("cos_cached", cos_cached, persistent=False)
|
||||
self.register_buffer("sin_cached", sin_cached, persistent=False)
|
||||
self.max_len_cached = max_len
|
||||
|
||||
def forward(self, x: Tensor, start_pos: int = 0) -> Tuple[Tensor, Tensor]:
|
||||
seq_len = x.size(1)
|
||||
if self.max_len_cached < seq_len + start_pos:
|
||||
self._set_rotary_buffer(self.max_len_cached * 2, x.device)
|
||||
cos = self.cos_cached[start_pos : start_pos + seq_len]
|
||||
sin = self.sin_cached[start_pos : start_pos + seq_len]
|
||||
return (cos, sin)
|
||||
|
||||
|
||||
class Linear(nn.Module):
|
||||
def __init__(self, in_dim: int, out_dim: int, bias: bool = False):
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.empty((out_dim, in_dim)))
|
||||
self.bias = nn.Parameter(torch.zeros(out_dim)) if bias else None
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return F.linear(x, self.weight, self.bias)
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
def __init__(self, dim, norm_eps):
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
self.normalized_shape = (dim,)
|
||||
self.norm_eps = norm_eps
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return F.rms_norm(x, self.normalized_shape, self.weight, self.norm_eps)
|
||||
|
||||
|
||||
class MLP(nn.Module):
|
||||
def __init__(self, dim: int, dim_feed_forward: int):
|
||||
super().__init__()
|
||||
self.up = Linear(dim, dim_feed_forward)
|
||||
self.gate = Linear(dim, dim_feed_forward)
|
||||
self.down = Linear(dim_feed_forward, dim)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
gated = self.up(x) * F.silu(self.gate(x))
|
||||
out = self.down(gated)
|
||||
return out
|
||||
|
||||
|
||||
class GQA(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
n_heads: int,
|
||||
n_kv_heads: int,
|
||||
use_qk_norm: bool,
|
||||
norm_eps: float,
|
||||
use_gated_attention: bool,
|
||||
layer_id: int,
|
||||
):
|
||||
super().__init__()
|
||||
assert dim % n_heads == 0
|
||||
assert n_heads % n_kv_heads == 0
|
||||
|
||||
self.head_dim = dim // n_heads
|
||||
self.layer_id = layer_id
|
||||
self.dim = dim
|
||||
self.n_heads = n_heads
|
||||
self.n_kv_heads = n_kv_heads
|
||||
self.n_rep = n_heads // n_kv_heads
|
||||
self.use_qk_norm = use_qk_norm
|
||||
self.use_gated_attention = use_gated_attention
|
||||
|
||||
self.q_proj = Linear(dim, n_heads * self.head_dim)
|
||||
self.k_proj = Linear(dim, n_kv_heads * self.head_dim)
|
||||
self.v_proj = Linear(dim, n_kv_heads * self.head_dim)
|
||||
self.o_proj = Linear(dim, dim)
|
||||
|
||||
if self.use_qk_norm:
|
||||
self.q_norm = RMSNorm(self.head_dim, norm_eps)
|
||||
self.k_norm = RMSNorm(self.head_dim, norm_eps)
|
||||
|
||||
if self.use_gated_attention:
|
||||
self.gate = Linear(dim, dim)
|
||||
|
||||
def _split_heads(self, x: Tensor, n_heads) -> Tensor:
|
||||
batch_size, seq_len, _ = x.shape
|
||||
x = x.reshape(batch_size, seq_len, n_heads, self.head_dim)
|
||||
return x
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
rotary_emb: Tuple[Tensor, Tensor],
|
||||
mask: Tensor = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
start_pos: int = 0,
|
||||
) -> Tensor:
|
||||
bsz, seq_len, _ = x.size()
|
||||
is_causal = mask is None
|
||||
|
||||
# (bsz, seq_len, dim) -> (bsz, seq_len, n_heads, head_dim)
|
||||
q = self._split_heads(self.q_proj(x), self.n_heads)
|
||||
k = self._split_heads(self.k_proj(x), self.n_kv_heads)
|
||||
v = self._split_heads(self.v_proj(x), self.n_kv_heads)
|
||||
q, k = apply_rotary_emb(q, rotary_emb), apply_rotary_emb(k, rotary_emb)
|
||||
|
||||
if self.use_qk_norm:
|
||||
q, k = self.q_norm(q), self.k_norm(k)
|
||||
|
||||
if paged_cache is not None:
|
||||
paged_cache.write(self.layer_id, start_pos, k, v)
|
||||
k, v = paged_cache.gather(self.layer_id)
|
||||
|
||||
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)
|
||||
|
||||
# (bsz, seq_len, n_heads, head_dim) -> (bsz, n_heads, seq_len, head_dim)
|
||||
q, k, v = q.permute(0, 2, 1, 3), k.permute(0, 2, 1, 3), v.permute(0, 2, 1, 3)
|
||||
sdqa_out = (
|
||||
F.scaled_dot_product_attention(q, k, v, mask, is_causal=is_causal)
|
||||
.permute(0, 2, 1, 3)
|
||||
.contiguous()
|
||||
.flatten(2)
|
||||
)
|
||||
|
||||
if self.use_gated_attention:
|
||||
sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
|
||||
|
||||
out = self.o_proj(sdqa_out)
|
||||
return out
|
||||
|
||||
|
||||
class MLA(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
n_heads: int,
|
||||
n_kv_heads: int,
|
||||
kv_lora_rank: int,
|
||||
qk_nope_head_dim: int,
|
||||
qk_rope_head_dim: int,
|
||||
norm_eps: float,
|
||||
use_gated_attention: bool,
|
||||
layer_id: int,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.n_heads = n_heads
|
||||
self.n_kv_heads = n_kv_heads
|
||||
self.kv_lora_rank = kv_lora_rank
|
||||
self.qk_nope_head_dim = qk_nope_head_dim
|
||||
self.qk_rope_head_dim = qk_rope_head_dim
|
||||
self.head_dim = qk_nope_head_dim + qk_rope_head_dim
|
||||
self.layer_id = layer_id
|
||||
self.n_rep = n_heads // n_kv_heads
|
||||
self.use_gated_attention = use_gated_attention
|
||||
|
||||
self.q_proj = Linear(dim, n_heads * self.head_dim, bias=False)
|
||||
self.kv_a_proj = Linear(dim, kv_lora_rank, bias=False)
|
||||
self.kv_norm = RMSNorm(kv_lora_rank, norm_eps)
|
||||
|
||||
# fused KV: (k_nope, k_rope, v)
|
||||
self.kv_b_proj = Linear(
|
||||
kv_lora_rank,
|
||||
n_kv_heads * (self.head_dim + qk_rope_head_dim + self.head_dim),
|
||||
)
|
||||
|
||||
self.o_proj = Linear(dim, dim, bias=False)
|
||||
|
||||
if use_gated_attention:
|
||||
self.gate = Linear(dim, dim, bias=False)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
rotary_emb: Tuple[Tensor, Tensor],
|
||||
mask: Tensor = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
start_pos: int = 0,
|
||||
) -> Tensor:
|
||||
bsz, seq_len, _ = x.size()
|
||||
is_causal = mask is None
|
||||
|
||||
q = self.q_proj(x)
|
||||
q = q.view(bsz, seq_len, self.n_heads, self.head_dim)
|
||||
|
||||
kv_compressed = self.kv_a_proj(x)
|
||||
kv_compressed = self.kv_norm(kv_compressed)
|
||||
|
||||
kv = self.kv_b_proj(kv_compressed)
|
||||
kv = kv.view(bsz, seq_len, self.n_kv_heads, -1)
|
||||
|
||||
k_nope, k_rope, v = torch.split(
|
||||
kv, [self.qk_nope_head_dim, self.qk_rope_head_dim, self.head_dim], dim=-1
|
||||
)
|
||||
|
||||
q_nope, q_rope = (
|
||||
q[..., : self.qk_nope_head_dim],
|
||||
q[..., self.qk_rope_head_dim :],
|
||||
)
|
||||
q_rope = apply_rotary_emb(q_rope, rotary_emb)
|
||||
k_rope = apply_rotary_emb(k_rope, rotary_emb)
|
||||
|
||||
q = torch.cat([q_nope, q_rope], dim=-1)
|
||||
k = torch.cat([k_nope, k_rope], dim=-1)
|
||||
|
||||
if paged_cache is not None:
|
||||
paged_cache.write(self.layer_id, start_pos, k, v)
|
||||
k, v = paged_cache.gather(self.layer_id)
|
||||
|
||||
q = q.permute(0, 2, 1, 3)
|
||||
k = k.permute(0, 2, 1, 3)
|
||||
v = v.permute(0, 2, 1, 3)
|
||||
|
||||
attn_out = F.scaled_dot_product_attention(q, k, v, mask, is_causal=is_causal)
|
||||
attn_out = attn_out.permute(0, 2, 1, 3).contiguous().flatten(2)
|
||||
|
||||
if self.use_gated_attention:
|
||||
attn_out = attn_out * F.sigmoid(self.gate(x))
|
||||
|
||||
out = self.o_proj(attn_out)
|
||||
return out
|
||||
|
||||
|
||||
class DecoderBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
n_heads: int,
|
||||
dim_ffn: int,
|
||||
n_kv_heads: int,
|
||||
norm_eps: int,
|
||||
use_qk_norm: bool,
|
||||
use_gated_attention: bool,
|
||||
layer_id: int,
|
||||
):
|
||||
super().__init__()
|
||||
self.attention = GQA(
|
||||
dim,
|
||||
n_heads,
|
||||
n_kv_heads,
|
||||
use_qk_norm,
|
||||
norm_eps,
|
||||
use_gated_attention,
|
||||
layer_id,
|
||||
)
|
||||
self.input_norm = RMSNorm(dim, norm_eps)
|
||||
self.mlp = MLP(dim, dim_ffn)
|
||||
self.post_attention_norm = RMSNorm(dim, norm_eps)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
rotary_emb: Tuple[Tensor, Tensor],
|
||||
attention_mask: Optional[Tensor] = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
start_pos: int = 0,
|
||||
) -> Tensor:
|
||||
attn_output = self.attention(
|
||||
self.input_norm(x),
|
||||
rotary_emb,
|
||||
attention_mask,
|
||||
paged_cache,
|
||||
start_pos,
|
||||
)
|
||||
x = attn_output + x
|
||||
|
||||
x = self.mlp(self.post_attention_norm(x)) + x
|
||||
return x
|
||||
|
||||
|
||||
class Embedding(nn.Module):
|
||||
def __init__(self, vocab_size: int, embedding_dim: int):
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.empty((vocab_size, embedding_dim)))
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return F.embedding(x, self.weight)
|
||||
+87
-81
@@ -1,98 +1,74 @@
|
||||
from typing import Any, Mapping, Optional
|
||||
from typing import Any, Dict, Mapping, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.config.model_config import ModelConfig
|
||||
from astrai.inference.cache import CacheView
|
||||
from astrai.model.automodel import AutoModel
|
||||
from astrai.model.module import (
|
||||
DecoderBlock,
|
||||
Embedding,
|
||||
Linear,
|
||||
RMSNorm,
|
||||
RotaryEmbedding,
|
||||
)
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.inference.cache import KVCache
|
||||
from astrai.model.automodel import AutoModel, ModelFactory
|
||||
from astrai.model.components.decoder_block import DecoderBlock
|
||||
from astrai.model.components.embedding import Embedding
|
||||
from astrai.model.components.linear import Linear
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
from astrai.model.components.rope import RotaryEmbedding
|
||||
|
||||
|
||||
def process_attention_mask(
|
||||
seq_mask: Tensor,
|
||||
input_tensor: Tensor,
|
||||
start_pos: int = 0,
|
||||
is_causal: bool = False,
|
||||
) -> Tensor:
|
||||
"""Build 4D attention mask from 2D seq_mask, with optional causal masking."""
|
||||
device = input_tensor.device
|
||||
dtype = input_tensor.dtype
|
||||
seq_len = input_tensor.size(1)
|
||||
|
||||
if seq_mask is None:
|
||||
if start_pos != 0:
|
||||
seq_mask = torch.ones((1, seq_len), dtype=torch.bool, device=device)
|
||||
else:
|
||||
return None
|
||||
|
||||
if seq_mask.dim() > 2:
|
||||
return seq_mask
|
||||
|
||||
batch_size = seq_mask.size(0)
|
||||
seq_mask = seq_mask[:, : start_pos + seq_len].to(device=device, dtype=torch.bool)
|
||||
expanded_mask = seq_mask.unsqueeze(1).expand(
|
||||
batch_size, seq_len, start_pos + seq_len
|
||||
)
|
||||
|
||||
if is_causal:
|
||||
expanded_mask = torch.tril(expanded_mask, diagonal=start_pos)
|
||||
|
||||
attention_mask = torch.zeros_like(expanded_mask, dtype=dtype, device=device)
|
||||
attention_mask = attention_mask.masked_fill_(
|
||||
~expanded_mask, -torch.finfo(dtype).max / 2
|
||||
).unsqueeze(1)
|
||||
|
||||
return attention_mask
|
||||
input_mask: Optional[Tensor],
|
||||
) -> Optional[Tensor]:
|
||||
if input_mask is None:
|
||||
return None
|
||||
if input_mask.dim() == 2:
|
||||
return input_mask[:, None, None, :]
|
||||
if input_mask.dim() == 3:
|
||||
return input_mask[:, None, :, :]
|
||||
return input_mask
|
||||
|
||||
|
||||
@AutoModel.register("transformer")
|
||||
class Transformer(AutoModel):
|
||||
"""Transformer language model with paged KV cache."""
|
||||
@ModelFactory.register("autoregressive_lm")
|
||||
class AutoRegressiveLM(AutoModel):
|
||||
"""Autoregressive language model with paged KV cache."""
|
||||
|
||||
def __init__(self, config: ModelConfig):
|
||||
def __init__(self, config: AutoRegressiveLMConfig):
|
||||
super().__init__(config)
|
||||
self.config = config
|
||||
rope_dim = (
|
||||
config.qk_rope_head_dim
|
||||
if config.attn_type == "mla"
|
||||
else config.hidden_size // config.num_attention_heads
|
||||
)
|
||||
rope_base = config.rope_theta if config.rope_theta is not None else 10000
|
||||
self.rotary_embedding = RotaryEmbedding(
|
||||
config.dim // config.n_heads, config.max_len
|
||||
rope_dim,
|
||||
config.max_position_embeddings,
|
||||
rope_base,
|
||||
rope_scaling=config.rope_scaling,
|
||||
)
|
||||
self.embed_tokens = Embedding(
|
||||
config.vocab_size,
|
||||
config.hidden_size,
|
||||
neftune_alpha=config.neftune_alpha,
|
||||
)
|
||||
self.embed_tokens = Embedding(config.vocab_size, config.dim)
|
||||
|
||||
self.layers = nn.ModuleList(
|
||||
[
|
||||
DecoderBlock(
|
||||
config.dim,
|
||||
config.n_heads,
|
||||
config.dim_ffn,
|
||||
config.n_kv_heads,
|
||||
config.norm_eps,
|
||||
config.use_qk_norm,
|
||||
config.use_gated_attention,
|
||||
layer_id,
|
||||
)
|
||||
for layer_id in range(config.n_layers)
|
||||
DecoderBlock(config, layer_id)
|
||||
for layer_id in range(config.num_hidden_layers)
|
||||
]
|
||||
)
|
||||
|
||||
self.norm = RMSNorm(config.dim, config.norm_eps)
|
||||
self.lm_head = Linear(config.dim, config.vocab_size)
|
||||
self.norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||
self.lm_head = Linear(config.hidden_size, config.vocab_size)
|
||||
|
||||
if self.config.tie_weight:
|
||||
if self.config.tie_word_embeddings is True:
|
||||
self.lm_head.weight = self.embed_tokens.weight
|
||||
|
||||
self._init_weights()
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self):
|
||||
for param in self.parameters():
|
||||
if param.dim() > 1:
|
||||
nn.init.normal_(param, mean=0.0, std=0.006)
|
||||
def _init_weights(self, module):
|
||||
if hasattr(module, "reset_parameters"):
|
||||
module.reset_parameters()
|
||||
|
||||
def load_state_dict(self, state_dict: Mapping[str, Any], strict=True, assign=False):
|
||||
lm_head_key = "lm_head.weight"
|
||||
@@ -100,7 +76,7 @@ class Transformer(AutoModel):
|
||||
|
||||
state_dict = dict(state_dict)
|
||||
|
||||
if self.config.tie_weight:
|
||||
if self.config.tie_word_embeddings is True:
|
||||
# same tensor for embed and lm_head
|
||||
if embed_key in state_dict:
|
||||
state_dict[lm_head_key] = state_dict[embed_key]
|
||||
@@ -116,7 +92,7 @@ class Transformer(AutoModel):
|
||||
destination=destination, prefix=prefix, keep_vars=keep_vars
|
||||
)
|
||||
|
||||
if self.config.tie_weight:
|
||||
if self.config.tie_word_embeddings is True:
|
||||
lm_head_key = prefix + "lm_head.weight"
|
||||
if lm_head_key in state_dict:
|
||||
del state_dict[lm_head_key]
|
||||
@@ -127,20 +103,50 @@ class Transformer(AutoModel):
|
||||
self,
|
||||
input_ids: Tensor,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
start_pos: int = 0,
|
||||
) -> Tensor:
|
||||
assert input_ids.ndim == 2
|
||||
kv_cache: Optional[KVCache] = None,
|
||||
position_ids: Optional[Tensor] = None,
|
||||
fwd: Optional[str] = None,
|
||||
) -> Dict[str, Tensor]:
|
||||
if fwd is None:
|
||||
if input_ids.ndim != 2:
|
||||
raise ValueError("training input_ids must be [batch, seq_len]")
|
||||
if kv_cache is not None:
|
||||
raise ValueError("training forward does not accept a KV cache")
|
||||
elif fwd in ("prefill", "decode"):
|
||||
if input_ids.ndim != 1:
|
||||
raise ValueError("inference input_ids must be packed [tokens]")
|
||||
if kv_cache is None:
|
||||
raise ValueError("inference forward requires a KV cache")
|
||||
else:
|
||||
raise ValueError(f"unsupported forward mode: {fwd}")
|
||||
|
||||
x = self.embed_tokens(input_ids)
|
||||
rotary_emb = self.rotary_embedding(x, start_pos)
|
||||
|
||||
attn_mask = process_attention_mask(input_mask, x, start_pos, is_causal=True)
|
||||
rotary_emb = self.rotary_embedding(x, position_ids)
|
||||
attn_mask = process_attention_mask(input_mask)
|
||||
use_sdpa_causal_mask = attn_mask is None
|
||||
|
||||
aux_losses = []
|
||||
router_stats_list = []
|
||||
for layer in self.layers:
|
||||
x = layer(x, rotary_emb, attn_mask, paged_cache, start_pos)
|
||||
layer_output = layer(
|
||||
x,
|
||||
rotary_emb,
|
||||
attn_mask,
|
||||
kv_cache,
|
||||
use_sdpa_causal_mask,
|
||||
fwd,
|
||||
)
|
||||
x = layer_output["hidden_states"]
|
||||
stats = layer_output.get("router_stats")
|
||||
if stats is not None:
|
||||
aux_losses.append(layer_output["aux_loss"])
|
||||
router_stats_list.append(stats)
|
||||
|
||||
hidden_states = self.norm(x)
|
||||
logits = self.lm_head(hidden_states)
|
||||
|
||||
return {"logits": logits, "hidden_states": hidden_states}
|
||||
output = {"logits": logits, "hidden_states": hidden_states}
|
||||
if aux_losses:
|
||||
output["aux_loss"] = torch.stack(aux_losses).mean()
|
||||
output["router_stats"] = router_stats_list
|
||||
return output
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
"""Optimizer implementations and factory registration."""
|
||||
|
||||
from astrai.optim.composite import (
|
||||
OptimizerFactory,
|
||||
composite_state_dict,
|
||||
composite_step,
|
||||
composite_zero_grad,
|
||||
refresh_param_groups,
|
||||
)
|
||||
from astrai.optim.mano_adamw import Mano, ManoAdamW
|
||||
from astrai.optim.muon_adamw import MuonAdamW
|
||||
from astrai.optim.nora_nadamw import (
|
||||
NAdamW,
|
||||
Nora,
|
||||
NoraNAdamW,
|
||||
OptimizerParameterGroups,
|
||||
nora_direction,
|
||||
nora_lr_scale,
|
||||
partition_optimizer_parameters,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Mano",
|
||||
"ManoAdamW",
|
||||
"MuonAdamW",
|
||||
"NAdamW",
|
||||
"Nora",
|
||||
"NoraNAdamW",
|
||||
"OptimizerFactory",
|
||||
"OptimizerParameterGroups",
|
||||
"composite_state_dict",
|
||||
"composite_step",
|
||||
"composite_zero_grad",
|
||||
"nora_direction",
|
||||
"nora_lr_scale",
|
||||
"partition_optimizer_parameters",
|
||||
"refresh_param_groups",
|
||||
]
|
||||
@@ -0,0 +1,71 @@
|
||||
"""Shared infrastructure for the optim package.
|
||||
|
||||
This module hosts two things:
|
||||
|
||||
* ``OptimizerFactory`` — the registry for built-in optimizers. Defining it
|
||||
here (rather than in ``__init__.py``) lets each optimizer module import it
|
||||
and register itself with a decorator, avoiding circular imports.
|
||||
* Composite-optimizer helpers — ``step``/``zero_grad``/``state_dict``/
|
||||
``param_groups`` delegation shared by every optimizer that routes different
|
||||
parameter groups through distinct sub-optimizers.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch.optim import Optimizer
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
|
||||
class OptimizerFactory(BaseFactory[Optimizer]):
|
||||
"""Factory for built-in training optimizers."""
|
||||
|
||||
|
||||
def composite_step(
|
||||
sub_optimizers: list[Optimizer],
|
||||
closure=None,
|
||||
) -> torch.Tensor | None:
|
||||
"""Run ``step`` on every sub-optimizer, invoking the closure once.
|
||||
|
||||
The closure (if given) is executed inside ``torch.enable_grad`` exactly
|
||||
once before any sub-optimizer steps, matching the contract of a single
|
||||
``Optimizer.step``. Sub-optimizers receive ``None`` so they do not
|
||||
re-execute it.
|
||||
"""
|
||||
loss = None
|
||||
if closure is not None:
|
||||
with torch.enable_grad():
|
||||
loss = closure()
|
||||
for sub in sub_optimizers:
|
||||
sub.step()
|
||||
return loss
|
||||
|
||||
|
||||
def composite_zero_grad(
|
||||
sub_optimizers: list[Optimizer],
|
||||
set_to_none: bool = True,
|
||||
) -> None:
|
||||
for sub in sub_optimizers:
|
||||
sub.zero_grad(set_to_none=set_to_none)
|
||||
|
||||
|
||||
def composite_state_dict(
|
||||
named_sub_optimizers: dict[str, Optimizer | None],
|
||||
) -> dict[str, Any]:
|
||||
"""Serialize sub-optimizers, preserving ``None`` slots."""
|
||||
return {
|
||||
name: sub.state_dict() if sub is not None else None
|
||||
for name, sub in named_sub_optimizers.items()
|
||||
}
|
||||
|
||||
|
||||
def refresh_param_groups(
|
||||
sub_optimizers: list[Optimizer],
|
||||
) -> list[dict]:
|
||||
"""Concatenate param_groups from every non-None sub-optimizer."""
|
||||
groups: list[dict] = []
|
||||
for sub in sub_optimizers:
|
||||
if sub is not None:
|
||||
groups.extend(sub.param_groups)
|
||||
return groups
|
||||
@@ -0,0 +1,214 @@
|
||||
"""Mano manifold optimizer combined with AdamW.
|
||||
|
||||
Mano projects the momentum onto the tangent space of the Oblique manifold
|
||||
(axis-wise tangent projection) and normalizes it, replacing the expensive
|
||||
Newton-Schulz iteration in Muon with a cheaper manifold normalization.
|
||||
|
||||
Reference: https://arxiv.org/abs/2601.23000
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
from torch import nn, optim
|
||||
from torch.optim import Optimizer
|
||||
|
||||
from astrai.optim.composite import (
|
||||
OptimizerFactory,
|
||||
composite_state_dict,
|
||||
composite_step,
|
||||
composite_zero_grad,
|
||||
refresh_param_groups,
|
||||
)
|
||||
from astrai.optim.nora_nadamw import partition_optimizer_parameters
|
||||
|
||||
|
||||
class Mano(Optimizer):
|
||||
"""Manifold Normalized Optimizer for two-dimensional matrices.
|
||||
|
||||
Each step alternates the projection axis (dim 0 / dim 1) to restrike the
|
||||
manifold along both rows and columns. The tangent momentum is computed
|
||||
without normalizing the parameter itself (v2 simplification) and the
|
||||
epsilon is added (not clamped) to the norm denominator.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
params,
|
||||
lr: float = 1e-3,
|
||||
weight_decay: float = 0.1,
|
||||
momentum: float = 0.95,
|
||||
nesterov: bool = True,
|
||||
eps: float = 1e-8,
|
||||
):
|
||||
if lr < 0:
|
||||
raise ValueError(f"Invalid learning rate: {lr}")
|
||||
if weight_decay < 0:
|
||||
raise ValueError(f"Invalid weight decay: {weight_decay}")
|
||||
if not 0 <= momentum <= 1:
|
||||
raise ValueError(f"Invalid momentum: {momentum}")
|
||||
if eps <= 0:
|
||||
raise ValueError(f"Invalid epsilon: {eps}")
|
||||
|
||||
defaults = {
|
||||
"lr": lr,
|
||||
"weight_decay": weight_decay,
|
||||
"momentum": momentum,
|
||||
"nesterov": nesterov,
|
||||
"eps": eps,
|
||||
"steps": 0,
|
||||
}
|
||||
super().__init__(params, defaults)
|
||||
for group in self.param_groups:
|
||||
for param in group["params"]:
|
||||
if param.ndim != 2:
|
||||
raise ValueError(
|
||||
f"Mano only supports 2D matrices, got shape {tuple(param.shape)}"
|
||||
)
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
loss = None
|
||||
if closure is not None:
|
||||
with torch.enable_grad():
|
||||
loss = closure()
|
||||
|
||||
for group in self.param_groups:
|
||||
lr = group["lr"]
|
||||
weight_decay = group["weight_decay"]
|
||||
momentum = group["momentum"]
|
||||
nesterov = group["nesterov"]
|
||||
eps = group["eps"]
|
||||
dim = int(group["steps"] % 2)
|
||||
|
||||
for param in group["params"]:
|
||||
if param.grad is None:
|
||||
continue
|
||||
if param.grad.is_sparse:
|
||||
raise RuntimeError("Mano does not support sparse gradients")
|
||||
|
||||
grad = param.grad
|
||||
state = self.state[param]
|
||||
momentum_buffer = state.get("momentum_buffer")
|
||||
if momentum_buffer is None:
|
||||
momentum_buffer = torch.zeros_like(grad)
|
||||
momentum_buffer.mul_(momentum).add_(grad)
|
||||
update = (
|
||||
grad.add(momentum_buffer, alpha=momentum)
|
||||
if nesterov
|
||||
else momentum_buffer
|
||||
)
|
||||
|
||||
tangent = update - (
|
||||
torch.sum(update * param.data, dim=dim, keepdim=True) * param.data
|
||||
)
|
||||
direction = tangent / (
|
||||
torch.norm(tangent, p=2, dim=dim, keepdim=True) + eps
|
||||
)
|
||||
|
||||
if weight_decay != 0:
|
||||
param.mul_(1 - lr * weight_decay)
|
||||
adjusted_lr = lr * 0.2 * math.sqrt(direction.shape[dim])
|
||||
param.add_(direction, alpha=-adjusted_lr)
|
||||
state["momentum_buffer"] = momentum_buffer
|
||||
|
||||
group["steps"] += 1
|
||||
|
||||
return loss
|
||||
|
||||
|
||||
@OptimizerFactory.register("mano_adamw")
|
||||
class ManoAdamW(Optimizer):
|
||||
"""Mano for internal linear weights and AdamW for remaining parameters."""
|
||||
|
||||
optimizer_name = "mano_adamw"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: nn.Module,
|
||||
lr: float = 3e-4,
|
||||
weight_decay: float = 0.1,
|
||||
momentum: float = 0.95,
|
||||
nesterov: bool = True,
|
||||
):
|
||||
groups = partition_optimizer_parameters(model)
|
||||
all_params = [
|
||||
*groups.nora,
|
||||
*groups.nadamw_decay,
|
||||
*groups.nadamw_no_decay,
|
||||
]
|
||||
if not all_params:
|
||||
raise ValueError(
|
||||
"Cannot build an optimizer for a model with no trainable parameters"
|
||||
)
|
||||
super().__init__(all_params, {})
|
||||
|
||||
self.mano = (
|
||||
Mano(
|
||||
groups.nora,
|
||||
lr=lr,
|
||||
weight_decay=weight_decay,
|
||||
momentum=momentum,
|
||||
nesterov=nesterov,
|
||||
)
|
||||
if groups.nora
|
||||
else None
|
||||
)
|
||||
|
||||
adamw_groups = []
|
||||
if groups.nadamw_decay:
|
||||
adamw_groups.append(
|
||||
{"params": groups.nadamw_decay, "weight_decay": weight_decay}
|
||||
)
|
||||
if groups.nadamw_no_decay:
|
||||
adamw_groups.append({"params": groups.nadamw_no_decay, "weight_decay": 0.0})
|
||||
self.adamw = (
|
||||
optim.AdamW(
|
||||
adamw_groups,
|
||||
lr=lr,
|
||||
betas=(0.9, 0.95),
|
||||
fused=True,
|
||||
)
|
||||
if adamw_groups
|
||||
else None
|
||||
)
|
||||
self.param_groups = refresh_param_groups([self.mano, self.adamw])
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
return composite_step(
|
||||
[opt for opt in (self.mano, self.adamw) if opt is not None],
|
||||
closure,
|
||||
)
|
||||
|
||||
def zero_grad(self, set_to_none: bool = True):
|
||||
composite_zero_grad(
|
||||
[opt for opt in (self.mano, self.adamw) if opt is not None],
|
||||
set_to_none,
|
||||
)
|
||||
|
||||
def state_dict(self) -> dict:
|
||||
return composite_state_dict({"mano": self.mano, "adamw": self.adamw})
|
||||
|
||||
def load_state_dict(self, state_dict: dict):
|
||||
if "muon" in state_dict or "nora" in state_dict:
|
||||
raise ValueError(
|
||||
"Checkpoint uses a different optimizer; select the matching "
|
||||
"--optimizer to resume it"
|
||||
)
|
||||
if "mano" not in state_dict or "adamw" not in state_dict:
|
||||
raise ValueError(
|
||||
"Checkpoint optimizer state is not compatible with mano_adamw"
|
||||
)
|
||||
|
||||
saved_mano = state_dict["mano"]
|
||||
saved_adamw = state_dict["adamw"]
|
||||
if (self.mano is None) != (saved_mano is None):
|
||||
raise ValueError("Checkpoint Mano parameter groups do not match the model")
|
||||
if (self.adamw is None) != (saved_adamw is None):
|
||||
raise ValueError("Checkpoint AdamW parameter groups do not match the model")
|
||||
if self.mano is not None:
|
||||
self.mano.load_state_dict(saved_mano)
|
||||
if self.adamw is not None:
|
||||
self.adamw.load_state_dict(saved_adamw)
|
||||
self.param_groups = refresh_param_groups([self.mano, self.adamw])
|
||||
@@ -0,0 +1,95 @@
|
||||
"""Legacy Muon + AdamW combined optimizer."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch import Tensor, nn, optim
|
||||
|
||||
from astrai.optim.composite import (
|
||||
OptimizerFactory,
|
||||
composite_state_dict,
|
||||
composite_step,
|
||||
composite_zero_grad,
|
||||
refresh_param_groups,
|
||||
)
|
||||
|
||||
|
||||
@OptimizerFactory.register("muon_adamw")
|
||||
class MuonAdamW(optim.Optimizer):
|
||||
"""Combined Muon (matrix) + AdamW (non-matrix) optimizer."""
|
||||
|
||||
optimizer_name = "muon_adamw"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: nn.Module,
|
||||
lr: float = 3e-4,
|
||||
weight_decay: float = 0.1,
|
||||
momentum: float = 0.95,
|
||||
nesterov: bool = True,
|
||||
ns_steps: int = 5,
|
||||
adjust_lr_fn: str = "match_rms_adamw",
|
||||
):
|
||||
defaults = {
|
||||
"lr": lr,
|
||||
"weight_decay": weight_decay,
|
||||
"momentum": momentum,
|
||||
"nesterov": nesterov,
|
||||
"ns_steps": ns_steps,
|
||||
"adjust_lr_fn": adjust_lr_fn,
|
||||
}
|
||||
params = [param for param in model.parameters() if param.requires_grad]
|
||||
super().__init__(params, defaults)
|
||||
|
||||
matrix_params: list[Tensor] = []
|
||||
other_params: list[Tensor] = []
|
||||
for name, param in model.named_parameters():
|
||||
if not param.requires_grad:
|
||||
continue
|
||||
if (
|
||||
param.dim() >= 2
|
||||
and "norm" not in name
|
||||
and "bias" not in name
|
||||
and "embed" not in name
|
||||
and "lm_head" not in name
|
||||
):
|
||||
matrix_params.append(param)
|
||||
else:
|
||||
other_params.append(param)
|
||||
|
||||
self.muon = optim.Muon(
|
||||
matrix_params,
|
||||
lr=lr,
|
||||
weight_decay=weight_decay,
|
||||
momentum=momentum,
|
||||
nesterov=nesterov,
|
||||
ns_steps=ns_steps,
|
||||
adjust_lr_fn=adjust_lr_fn,
|
||||
)
|
||||
self.adamw = optim.AdamW(
|
||||
[{"params": other_params, "weight_decay": 0.0}],
|
||||
lr=lr,
|
||||
betas=(0.9, 0.95),
|
||||
fused=True,
|
||||
)
|
||||
|
||||
self.param_groups = refresh_param_groups([self.muon, self.adamw])
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
return composite_step([self.muon, self.adamw], closure)
|
||||
|
||||
def zero_grad(self, set_to_none: bool = True):
|
||||
composite_zero_grad([self.muon, self.adamw], set_to_none)
|
||||
|
||||
def state_dict(self) -> dict[str, Any]:
|
||||
return composite_state_dict({"muon": self.muon, "adamw": self.adamw})
|
||||
|
||||
def load_state_dict(self, state_dict: dict[str, Any]):
|
||||
if "muon" not in state_dict or "adamw" not in state_dict:
|
||||
raise ValueError(
|
||||
"Checkpoint optimizer state is not compatible with muon_adamw"
|
||||
)
|
||||
self.muon.load_state_dict(state_dict["muon"])
|
||||
self.adamw.load_state_dict(state_dict["adamw"])
|
||||
self.param_groups = refresh_param_groups([self.muon, self.adamw])
|
||||
@@ -0,0 +1,372 @@
|
||||
"""Nora matrix optimizer combined with Nesterov AdamW."""
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch import Tensor, nn
|
||||
from torch.distributed.tensor import DTensor, Shard
|
||||
from torch.optim import Optimizer
|
||||
|
||||
from astrai.model.components.embedding import Embedding
|
||||
from astrai.model.components.linear import Linear
|
||||
from astrai.model.components.lora import LoRALinear
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
from astrai.optim.composite import (
|
||||
OptimizerFactory,
|
||||
composite_state_dict,
|
||||
composite_step,
|
||||
composite_zero_grad,
|
||||
refresh_param_groups,
|
||||
)
|
||||
|
||||
NORA_EPS = 1e-10
|
||||
|
||||
|
||||
def _row_normalize(tensor: Tensor, eps: float) -> Tensor:
|
||||
return tensor / tensor.norm(dim=-1, keepdim=True).clamp(min=eps)
|
||||
|
||||
|
||||
def nora_direction(update: Tensor, param: Tensor, eps: float = NORA_EPS) -> Tensor:
|
||||
"""Project an update onto each parameter row's tangent space and normalize."""
|
||||
theta_hat = _row_normalize(param.to(torch.float32), eps)
|
||||
update_fp32 = update.to(torch.float32)
|
||||
radial = (update_fp32 * theta_hat).sum(dim=-1, keepdim=True) * theta_hat
|
||||
direction = _row_normalize(update_fp32 - radial, eps)
|
||||
return direction.to(update.dtype)
|
||||
|
||||
|
||||
def nora_lr_scale(lr: float, shape: torch.Size) -> float:
|
||||
"""Scale Nora's LR for tall ``[d_out, d_in]`` linear weights."""
|
||||
return lr * math.sqrt(max(1.0, shape[-2] / shape[-1]))
|
||||
|
||||
|
||||
def _validate_complete_rows(param: Tensor) -> None:
|
||||
if not isinstance(param, DTensor):
|
||||
return
|
||||
last_dim = param.ndim - 1
|
||||
for placement in param.placements:
|
||||
if isinstance(placement, Shard) and placement.dim % param.ndim == last_dim:
|
||||
raise ValueError(
|
||||
"Nora requires complete parameter rows, but this DTensor is sharded "
|
||||
"along its last dimension"
|
||||
)
|
||||
|
||||
|
||||
class Nora(Optimizer):
|
||||
"""Normalized Orthogonal Row Alignment for two-dimensional matrices."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
params,
|
||||
lr: float = 5e-3,
|
||||
weight_decay: float = 0.0,
|
||||
momentum: float = 0.95,
|
||||
beta: float = 0.95,
|
||||
nesterov: bool = True,
|
||||
eps: float = NORA_EPS,
|
||||
):
|
||||
if lr < 0:
|
||||
raise ValueError(f"Invalid learning rate: {lr}")
|
||||
if weight_decay < 0:
|
||||
raise ValueError(f"Invalid weight decay: {weight_decay}")
|
||||
if not 0 <= momentum <= 1:
|
||||
raise ValueError(f"Invalid momentum: {momentum}")
|
||||
if not 0 <= beta < 1:
|
||||
raise ValueError(f"Invalid beta: {beta}")
|
||||
if eps <= 0:
|
||||
raise ValueError(f"Invalid epsilon: {eps}")
|
||||
|
||||
defaults = {
|
||||
"lr": lr,
|
||||
"weight_decay": weight_decay,
|
||||
"momentum": momentum,
|
||||
"beta": beta,
|
||||
"nesterov": nesterov,
|
||||
"eps": eps,
|
||||
}
|
||||
super().__init__(params, defaults)
|
||||
for group in self.param_groups:
|
||||
for param in group["params"]:
|
||||
if param.ndim != 2:
|
||||
raise ValueError(
|
||||
f"Nora only supports 2D matrices, got shape {tuple(param.shape)}"
|
||||
)
|
||||
_validate_complete_rows(param)
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
loss = None
|
||||
if closure is not None:
|
||||
with torch.enable_grad():
|
||||
loss = closure()
|
||||
|
||||
for group in self.param_groups:
|
||||
lr = group["lr"]
|
||||
weight_decay = group["weight_decay"]
|
||||
momentum = group["momentum"]
|
||||
beta = group["beta"]
|
||||
nesterov = group["nesterov"]
|
||||
eps = group["eps"]
|
||||
for param in group["params"]:
|
||||
if param.grad is None:
|
||||
continue
|
||||
if param.grad.is_sparse:
|
||||
raise RuntimeError("Nora does not support sparse gradients")
|
||||
|
||||
grad = param.grad
|
||||
state = self.state[param]
|
||||
momentum_buffer = state.get("momentum_buffer")
|
||||
if momentum_buffer is None:
|
||||
momentum_buffer = torch.zeros_like(grad)
|
||||
momentum_buffer.lerp_(grad, 1 - beta)
|
||||
update = (
|
||||
grad.lerp(momentum_buffer, momentum)
|
||||
if nesterov
|
||||
else momentum_buffer
|
||||
)
|
||||
direction = nora_direction(update, param, eps)
|
||||
|
||||
if weight_decay != 0:
|
||||
param.mul_(1 - lr * weight_decay)
|
||||
param.add_(direction, alpha=-nora_lr_scale(lr, param.shape))
|
||||
state["momentum_buffer"] = momentum_buffer
|
||||
|
||||
return loss
|
||||
|
||||
|
||||
class NAdamW(Optimizer):
|
||||
"""AdamW using the reference Nesterov first-moment update."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
params,
|
||||
lr: float = 3e-4,
|
||||
betas: tuple[float, float] = (0.9, 0.999),
|
||||
eps: float = 1e-8,
|
||||
weight_decay: float = 0.1,
|
||||
):
|
||||
beta1, beta2 = betas
|
||||
if lr < 0:
|
||||
raise ValueError(f"Invalid learning rate: {lr}")
|
||||
if not 0 <= beta1 < 1 or not 0 <= beta2 < 1:
|
||||
raise ValueError(f"Invalid betas: {betas}")
|
||||
if eps <= 0:
|
||||
raise ValueError(f"Invalid epsilon: {eps}")
|
||||
if weight_decay < 0:
|
||||
raise ValueError(f"Invalid weight decay: {weight_decay}")
|
||||
defaults = {
|
||||
"lr": lr,
|
||||
"betas": betas,
|
||||
"eps": eps,
|
||||
"weight_decay": weight_decay,
|
||||
}
|
||||
super().__init__(params, defaults)
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
loss = None
|
||||
if closure is not None:
|
||||
with torch.enable_grad():
|
||||
loss = closure()
|
||||
|
||||
for group in self.param_groups:
|
||||
beta1, beta2 = group["betas"]
|
||||
eps = group["eps"]
|
||||
lr = group["lr"]
|
||||
weight_decay = group["weight_decay"]
|
||||
for param in group["params"]:
|
||||
if param.grad is None:
|
||||
continue
|
||||
if param.grad.is_sparse:
|
||||
raise RuntimeError("NAdamW does not support sparse gradients")
|
||||
|
||||
grad = param.grad
|
||||
state = self.state[param]
|
||||
if not state:
|
||||
state["step"] = 0
|
||||
state["m"] = torch.zeros_like(param)
|
||||
state["v"] = torch.zeros_like(param)
|
||||
|
||||
state["step"] += 1
|
||||
first_moment = state["m"]
|
||||
second_moment = state["v"]
|
||||
first_moment.mul_(beta1).add_(grad, alpha=1 - beta1)
|
||||
second_moment.mul_(beta2).addcmul_(grad, grad, value=1 - beta2)
|
||||
|
||||
bias_correction1 = 1 - beta1 ** state["step"]
|
||||
bias_correction2 = 1 - beta2 ** state["step"]
|
||||
nesterov_moment = (
|
||||
beta1 * first_moment + (1 - beta1) * grad
|
||||
) / bias_correction1
|
||||
corrected_second_moment = second_moment / bias_correction2
|
||||
|
||||
if weight_decay != 0:
|
||||
param.mul_(1 - lr * weight_decay)
|
||||
param.addcdiv_(
|
||||
nesterov_moment,
|
||||
corrected_second_moment.sqrt().add_(eps),
|
||||
value=-lr,
|
||||
)
|
||||
|
||||
return loss
|
||||
|
||||
|
||||
@dataclass
|
||||
class OptimizerParameterGroups:
|
||||
nora: list[Tensor]
|
||||
nadamw_decay: list[Tensor]
|
||||
nadamw_no_decay: list[Tensor]
|
||||
|
||||
|
||||
def partition_optimizer_parameters(model: nn.Module) -> OptimizerParameterGroups:
|
||||
"""Partition trainable parameters by module role and parameter identity."""
|
||||
nora_ids: set[int] = set()
|
||||
no_decay_ids: set[int] = set()
|
||||
|
||||
for module_name, module in model.named_modules():
|
||||
if isinstance(module, LoRALinear):
|
||||
for param in module.parameters(recurse=False):
|
||||
if param.requires_grad:
|
||||
no_decay_ids.add(id(param))
|
||||
continue
|
||||
|
||||
if isinstance(module, (Embedding, RMSNorm)):
|
||||
for param in module.parameters(recurse=False):
|
||||
if param.requires_grad:
|
||||
no_decay_ids.add(id(param))
|
||||
continue
|
||||
|
||||
if not isinstance(module, Linear):
|
||||
continue
|
||||
|
||||
if module.bias is not None and module.bias.requires_grad:
|
||||
no_decay_ids.add(id(module.bias))
|
||||
if not module.weight.requires_grad:
|
||||
continue
|
||||
if module_name.rsplit(".", 1)[-1] == "lm_head":
|
||||
no_decay_ids.add(id(module.weight))
|
||||
elif module.weight.ndim == 2:
|
||||
nora_ids.add(id(module.weight))
|
||||
|
||||
nora: list[Tensor] = []
|
||||
nadamw_decay: list[Tensor] = []
|
||||
nadamw_no_decay: list[Tensor] = []
|
||||
seen: set[int] = set()
|
||||
for param in model.parameters():
|
||||
param_id = id(param)
|
||||
if not param.requires_grad or param_id in seen:
|
||||
continue
|
||||
seen.add(param_id)
|
||||
if param_id in no_decay_ids or param.ndim <= 1:
|
||||
nadamw_no_decay.append(param)
|
||||
elif param_id in nora_ids:
|
||||
nora.append(param)
|
||||
else:
|
||||
nadamw_decay.append(param)
|
||||
|
||||
trainable_ids = {id(param) for param in model.parameters() if param.requires_grad}
|
||||
grouped_ids = {id(param) for param in [*nora, *nadamw_decay, *nadamw_no_decay]}
|
||||
if grouped_ids != trainable_ids:
|
||||
missing = len(trainable_ids - grouped_ids)
|
||||
extra = len(grouped_ids - trainable_ids)
|
||||
raise RuntimeError(
|
||||
f"Optimizer parameter partition is incomplete: missing={missing}, extra={extra}"
|
||||
)
|
||||
|
||||
return OptimizerParameterGroups(nora, nadamw_decay, nadamw_no_decay)
|
||||
|
||||
|
||||
@OptimizerFactory.register("nora_nadamw")
|
||||
class NoraNAdamW(Optimizer):
|
||||
"""Nora for internal linear weights and NAdamW for remaining parameters."""
|
||||
|
||||
optimizer_name = "nora_nadamw"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: nn.Module,
|
||||
lr: float = 3e-4,
|
||||
weight_decay: float = 0.1,
|
||||
nora_lr: float = 5e-3,
|
||||
nora_weight_decay: float = 0.0,
|
||||
nora_beta: float = 0.95,
|
||||
nora_momentum: float = 0.95,
|
||||
):
|
||||
groups = partition_optimizer_parameters(model)
|
||||
all_params = [
|
||||
*groups.nora,
|
||||
*groups.nadamw_decay,
|
||||
*groups.nadamw_no_decay,
|
||||
]
|
||||
if not all_params:
|
||||
raise ValueError(
|
||||
"Cannot build an optimizer for a model with no trainable parameters"
|
||||
)
|
||||
super().__init__(all_params, {})
|
||||
|
||||
self.nora = (
|
||||
Nora(
|
||||
groups.nora,
|
||||
lr=nora_lr,
|
||||
weight_decay=nora_weight_decay,
|
||||
momentum=nora_momentum,
|
||||
beta=nora_beta,
|
||||
)
|
||||
if groups.nora
|
||||
else None
|
||||
)
|
||||
|
||||
nadamw_groups = []
|
||||
if groups.nadamw_decay:
|
||||
nadamw_groups.append(
|
||||
{"params": groups.nadamw_decay, "weight_decay": weight_decay}
|
||||
)
|
||||
if groups.nadamw_no_decay:
|
||||
nadamw_groups.append(
|
||||
{"params": groups.nadamw_no_decay, "weight_decay": 0.0}
|
||||
)
|
||||
self.nadamw = NAdamW(nadamw_groups, lr=lr) if nadamw_groups else None
|
||||
self.param_groups = refresh_param_groups([self.nora, self.nadamw])
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, closure=None):
|
||||
return composite_step(
|
||||
[opt for opt in (self.nora, self.nadamw) if opt is not None],
|
||||
closure,
|
||||
)
|
||||
|
||||
def zero_grad(self, set_to_none: bool = True):
|
||||
composite_zero_grad(
|
||||
[opt for opt in (self.nora, self.nadamw) if opt is not None],
|
||||
set_to_none,
|
||||
)
|
||||
|
||||
def state_dict(self) -> dict[str, Any]:
|
||||
return composite_state_dict({"nora": self.nora, "nadamw": self.nadamw})
|
||||
|
||||
def load_state_dict(self, state_dict: dict[str, Any]):
|
||||
if "muon" in state_dict or "adamw" in state_dict:
|
||||
raise ValueError(
|
||||
"Checkpoint uses muon_adamw state; select optimizer='muon_adamw' "
|
||||
"to resume it"
|
||||
)
|
||||
if "nora" not in state_dict or "nadamw" not in state_dict:
|
||||
raise ValueError(
|
||||
"Checkpoint optimizer state is not compatible with nora_nadamw"
|
||||
)
|
||||
|
||||
saved_nora = state_dict["nora"]
|
||||
saved_nadamw = state_dict["nadamw"]
|
||||
if (self.nora is None) != (saved_nora is None):
|
||||
raise ValueError("Checkpoint Nora parameter groups do not match the model")
|
||||
if (self.nadamw is None) != (saved_nadamw is None):
|
||||
raise ValueError(
|
||||
"Checkpoint NAdamW parameter groups do not match the model"
|
||||
)
|
||||
if self.nora is not None:
|
||||
self.nora.load_state_dict(saved_nora)
|
||||
if self.nadamw is not None:
|
||||
self.nadamw.load_state_dict(saved_nadamw)
|
||||
self.param_groups = refresh_param_groups([self.nora, self.nadamw])
|
||||
@@ -1,4 +1,15 @@
|
||||
from astrai.parallel.module import ColumnParallelLinear, RowParallelLinear
|
||||
from astrai.parallel.executor import (
|
||||
AccumOptimizer,
|
||||
AccumScheduler,
|
||||
BaseExecutor,
|
||||
DDPExecutor,
|
||||
ExecutorFactory,
|
||||
FSDPExecutor,
|
||||
GradientState,
|
||||
NoneExecutor,
|
||||
broadcast_state_dict,
|
||||
create_ref_model,
|
||||
)
|
||||
from astrai.parallel.setup import (
|
||||
get_current_device,
|
||||
get_rank,
|
||||
@@ -15,6 +26,14 @@ __all__ = [
|
||||
"only_on_rank",
|
||||
"setup_parallel",
|
||||
"spawn_parallel_fn",
|
||||
"RowParallelLinear",
|
||||
"ColumnParallelLinear",
|
||||
"ExecutorFactory",
|
||||
"BaseExecutor",
|
||||
"GradientState",
|
||||
"AccumOptimizer",
|
||||
"AccumScheduler",
|
||||
"NoneExecutor",
|
||||
"DDPExecutor",
|
||||
"FSDPExecutor",
|
||||
"create_ref_model",
|
||||
"broadcast_state_dict",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,428 @@
|
||||
"""Unified training executor — parallel strategy + gradient accumulation."""
|
||||
|
||||
import contextlib
|
||||
import logging
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from typing import Any, Callable, Dict, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
from torch.distributed.fsdp import (
|
||||
FSDPModule,
|
||||
fully_shard,
|
||||
)
|
||||
from torch.distributed.tensor import DTensor
|
||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||
from torch.optim import Optimizer
|
||||
from torch.optim.lr_scheduler import LRScheduler
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.parallel.setup import get_rank, get_world_size
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def broadcast_state_dict(
|
||||
state_dict: Optional[Dict[str, torch.Tensor]],
|
||||
src: int = 0,
|
||||
) -> Optional[Dict[str, torch.Tensor]]:
|
||||
"""Broadcast a state_dict from *src* rank to all ranks.
|
||||
|
||||
Tensors stay on their original device (GPU) for the broadcast.
|
||||
All ranks must call this collectively.
|
||||
|
||||
On non-distributed runs, returns *state_dict* unchanged.
|
||||
"""
|
||||
if not dist.is_initialized() or dist.get_world_size() == 1:
|
||||
return state_dict
|
||||
|
||||
rank = dist.get_rank()
|
||||
|
||||
# Broadcast metadata (keys, shapes, dtypes, device) so non-src ranks
|
||||
# can allocate matching empty tensors on the correct device.
|
||||
if rank == src:
|
||||
device = next(iter(state_dict.values())).device
|
||||
metadata = [
|
||||
(k, tuple(v.shape), v.dtype, str(device)) for k, v in state_dict.items()
|
||||
]
|
||||
else:
|
||||
metadata = None
|
||||
metadata_list = [metadata]
|
||||
dist.broadcast_object_list(metadata_list, src=src)
|
||||
metadata = metadata_list[0]
|
||||
|
||||
# Non-src ranks allocate empty tensors with the broadcasted metadata.
|
||||
if rank != src:
|
||||
state_dict = {
|
||||
k: torch.empty(s, dtype=d, device=torch.device(dev))
|
||||
for k, s, d, dev in metadata
|
||||
}
|
||||
|
||||
# Broadcast each tensor in-place.
|
||||
for tensor in state_dict.values():
|
||||
dist.broadcast(tensor, src=src)
|
||||
|
||||
return state_dict
|
||||
|
||||
|
||||
def create_ref_model(
|
||||
model_fn: Callable[[], nn.Module],
|
||||
executor: Optional["BaseExecutor"] = None,
|
||||
model: Optional[nn.Module] = None,
|
||||
state_dict: Optional[Dict[str, torch.Tensor]] = None,
|
||||
device: Optional[str] = None,
|
||||
) -> Optional[nn.Module]:
|
||||
"""Create a frozen reference model from executor or state dict.
|
||||
|
||||
In distributed mode (FSDP), ``unwrap_model`` returns ``None`` on
|
||||
non-rank-0. The state_dict is broadcast from rank-0 to all ranks
|
||||
so every rank gets a complete copy.
|
||||
"""
|
||||
if state_dict is None and executor is not None and model is not None:
|
||||
state_dict = executor.unwrap_model(model)
|
||||
|
||||
# FSDP's unwrap_model returns None on non-rank-0. Broadcast from
|
||||
# rank-0 so every rank receives a complete state_dict.
|
||||
if executor is not None and executor.use_distributed:
|
||||
state_dict = broadcast_state_dict(state_dict)
|
||||
|
||||
if state_dict is None:
|
||||
return None
|
||||
|
||||
ref_model = model_fn()
|
||||
ref_model.load_state_dict(state_dict)
|
||||
ref_model.requires_grad_(False)
|
||||
ref_model.eval()
|
||||
if device is not None:
|
||||
ref_model = ref_model.to(device=device)
|
||||
return ref_model
|
||||
|
||||
|
||||
class GradientState:
|
||||
def __init__(self, grad_accum_steps: int = 1):
|
||||
self.num_steps = max(grad_accum_steps, 1)
|
||||
self._step: int = 0
|
||||
self._sync_gradients: bool = True
|
||||
|
||||
@property
|
||||
def sync_gradients(self) -> bool:
|
||||
return self._sync_gradients
|
||||
|
||||
def _do_sync(self):
|
||||
self._step += 1
|
||||
self._sync_gradients = self._step % self.num_steps == 0
|
||||
|
||||
|
||||
class AccumOptimizer:
|
||||
def __init__(self, optimizer: Optimizer, gradient_state: GradientState):
|
||||
self.optimizer = optimizer
|
||||
self.gradient_state = gradient_state
|
||||
|
||||
def step(self, closure=None):
|
||||
if self.gradient_state.sync_gradients:
|
||||
self.optimizer.step(closure)
|
||||
|
||||
def zero_grad(self):
|
||||
if self.gradient_state.sync_gradients:
|
||||
self.optimizer.zero_grad()
|
||||
|
||||
@property
|
||||
def param_groups(self):
|
||||
return self.optimizer.param_groups
|
||||
|
||||
def state_dict(self):
|
||||
return self.optimizer.state_dict()
|
||||
|
||||
def load_state_dict(self, d):
|
||||
self.optimizer.load_state_dict(d)
|
||||
|
||||
|
||||
class AccumScheduler:
|
||||
def __init__(self, scheduler: LRScheduler, gradient_state: GradientState):
|
||||
self.scheduler = scheduler
|
||||
self.gradient_state = gradient_state
|
||||
|
||||
def step(self):
|
||||
if self.gradient_state.sync_gradients:
|
||||
self.scheduler.step()
|
||||
|
||||
def state_dict(self):
|
||||
return self.scheduler.state_dict()
|
||||
|
||||
def load_state_dict(self, d):
|
||||
self.scheduler.load_state_dict(d)
|
||||
|
||||
def get_last_lr(self):
|
||||
return self.scheduler.get_last_lr()
|
||||
|
||||
|
||||
class BaseExecutor:
|
||||
def __init__(self, grad_accum_steps: int = 1):
|
||||
self.gradient_state = GradientState(grad_accum_steps)
|
||||
|
||||
def prepare(
|
||||
self,
|
||||
model_fn: Callable[[], nn.Module],
|
||||
optimizer_fn: Optional[Callable[[nn.Module], Optimizer]] = None,
|
||||
scheduler_fn: Optional[Callable[[Optimizer], LRScheduler]] = None,
|
||||
before_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
|
||||
after_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
|
||||
) -> Tuple[nn.Module, Optional[Optimizer], Optional[LRScheduler]]:
|
||||
model = model_fn()
|
||||
if before_wrap is not None:
|
||||
model = before_wrap(model)
|
||||
model = self._prepare_model(model)
|
||||
if after_wrap is not None:
|
||||
model = after_wrap(model)
|
||||
optimizer = None
|
||||
scheduler = None
|
||||
if optimizer_fn is not None:
|
||||
optimizer = optimizer_fn(model)
|
||||
if scheduler_fn is not None:
|
||||
scheduler = scheduler_fn(optimizer)
|
||||
optimizer = AccumOptimizer(optimizer, self.gradient_state)
|
||||
if scheduler is not None:
|
||||
scheduler = AccumScheduler(scheduler, self.gradient_state)
|
||||
return model, optimizer, scheduler
|
||||
|
||||
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||
return model
|
||||
|
||||
def _no_sync(self, model: nn.Module):
|
||||
return contextlib.nullcontext()
|
||||
|
||||
@contextmanager
|
||||
def accumulate(self, model: nn.Module):
|
||||
self.gradient_state._do_sync()
|
||||
if not self.gradient_state.sync_gradients:
|
||||
with self._no_sync(model):
|
||||
yield
|
||||
else:
|
||||
yield
|
||||
|
||||
def backward(self, loss: torch.Tensor):
|
||||
loss.backward()
|
||||
|
||||
def unwrap_model(self, model: nn.Module):
|
||||
return model.state_dict()
|
||||
|
||||
@contextmanager
|
||||
def checkpoint_context(self, model: nn.Module):
|
||||
if self.use_distributed:
|
||||
dist.barrier()
|
||||
state_dict = self._gather_state_dict(model)
|
||||
yield state_dict
|
||||
if self.use_distributed:
|
||||
dist.barrier()
|
||||
|
||||
def _gather_state_dict(self, model: nn.Module):
|
||||
state_dict = self.unwrap_model(model)
|
||||
if self.use_distributed and get_rank() != 0:
|
||||
return None
|
||||
return state_dict
|
||||
|
||||
@property
|
||||
def use_distributed(self) -> bool:
|
||||
return get_world_size() > 1
|
||||
|
||||
@property
|
||||
def sync_gradients(self) -> bool:
|
||||
return self.gradient_state.sync_gradients
|
||||
|
||||
@property
|
||||
def grad_accum_steps(self) -> int:
|
||||
return self.gradient_state.num_steps
|
||||
|
||||
def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
|
||||
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
|
||||
if isinstance(total_norm, torch.Tensor):
|
||||
return total_norm.item()
|
||||
return total_norm
|
||||
|
||||
|
||||
class ExecutorFactory(BaseFactory[BaseExecutor]):
|
||||
pass
|
||||
|
||||
|
||||
@ExecutorFactory.register("none")
|
||||
class NoneExecutor(BaseExecutor):
|
||||
pass
|
||||
|
||||
|
||||
@ExecutorFactory.register("ddp")
|
||||
class DDPExecutor(BaseExecutor):
|
||||
def __init__(
|
||||
self,
|
||||
grad_accum_steps: int = 1,
|
||||
dim: int = 0,
|
||||
broadcast_buffers: bool = True,
|
||||
init_sync: bool = True,
|
||||
process_group=None,
|
||||
bucket_cap_mb: int = 25,
|
||||
find_unused_parameters: bool = False,
|
||||
check_reduction: bool = False,
|
||||
gradient_as_bucket_view: bool = False,
|
||||
static_graph: bool = False,
|
||||
delay_all_reduce_named_params=None,
|
||||
param_to_hook_all_reduce=None,
|
||||
mixed_precision=None,
|
||||
device_mesh=None,
|
||||
):
|
||||
super().__init__(grad_accum_steps=grad_accum_steps)
|
||||
self._ddp_kwargs = dict(
|
||||
dim=dim,
|
||||
broadcast_buffers=broadcast_buffers,
|
||||
init_sync=init_sync,
|
||||
process_group=process_group,
|
||||
bucket_cap_mb=bucket_cap_mb,
|
||||
find_unused_parameters=find_unused_parameters,
|
||||
check_reduction=check_reduction,
|
||||
gradient_as_bucket_view=gradient_as_bucket_view,
|
||||
static_graph=static_graph,
|
||||
delay_all_reduce_named_params=delay_all_reduce_named_params,
|
||||
param_to_hook_all_reduce=param_to_hook_all_reduce,
|
||||
mixed_precision=mixed_precision,
|
||||
device_mesh=device_mesh,
|
||||
)
|
||||
|
||||
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||
if not self.use_distributed:
|
||||
logger.warning("DDP backend selected but world_size=1, model not wrapped")
|
||||
return model
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", get_rank()))
|
||||
model = DDP(
|
||||
model,
|
||||
device_ids=[local_rank],
|
||||
output_device=local_rank,
|
||||
**self._ddp_kwargs,
|
||||
)
|
||||
logger.info("Model wrapped with DDP (world_size=%d)", get_world_size())
|
||||
return model
|
||||
|
||||
def _no_sync(self, model: nn.Module):
|
||||
if isinstance(model, DDP):
|
||||
return model.no_sync()
|
||||
return contextlib.nullcontext()
|
||||
|
||||
def unwrap_model(self, model: nn.Module):
|
||||
if isinstance(model, DDP):
|
||||
return model.module.state_dict()
|
||||
return model.state_dict()
|
||||
|
||||
|
||||
@ExecutorFactory.register("fsdp")
|
||||
class FSDPExecutor(BaseExecutor):
|
||||
"""FSDP executor using `torch.distributed.fsdp.fully_shard` (per-module API).
|
||||
|
||||
Wraps each child module individually via ``fully_shard``.
|
||||
Skips the root model because ``ABC + Generic[T]`` in the MRO makes
|
||||
``fully_shard``'s dynamic ``__class__`` assignment fail at the CPython level.
|
||||
Original ``Parameter`` objects are preserved (as DTensors) — no
|
||||
``FlatParameter``, no ``use_orig_params=True`` hack.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
grad_accum_steps: int = 1,
|
||||
mesh: Optional[Any] = None,
|
||||
mp_policy: Optional[Any] = None,
|
||||
reshard_after_forward: bool = False,
|
||||
):
|
||||
super().__init__(grad_accum_steps=grad_accum_steps)
|
||||
self._mesh = mesh
|
||||
self._mp_policy = mp_policy
|
||||
self._reshard_after_forward = reshard_after_forward
|
||||
|
||||
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||
if not self.use_distributed:
|
||||
logger.warning("FSDP backend selected but world_size=1, model not wrapped")
|
||||
return model
|
||||
|
||||
kwargs = dict(
|
||||
mesh=self._mesh,
|
||||
mp_policy=self._mp_policy,
|
||||
reshard_after_forward=self._reshard_after_forward,
|
||||
)
|
||||
kwargs = {k: v for k, v in kwargs.items() if v is not None}
|
||||
|
||||
for child in model.children():
|
||||
if isinstance(child, nn.ModuleList):
|
||||
for sub in child:
|
||||
fully_shard(sub, **kwargs)
|
||||
else:
|
||||
fully_shard(child, **kwargs)
|
||||
|
||||
logger.info(
|
||||
"FSDP wrapping applied to %d direct children (root skipped for ABC compat)",
|
||||
len(list(model.children())),
|
||||
)
|
||||
return model
|
||||
|
||||
@contextmanager
|
||||
def _no_sync(self, model: nn.Module):
|
||||
fsdp_modules = [m for m in model.modules() if isinstance(m, FSDPModule)]
|
||||
if fsdp_modules:
|
||||
for m in fsdp_modules:
|
||||
m.set_requires_gradient_sync(False, recurse=True)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
for m in fsdp_modules:
|
||||
m.set_requires_gradient_sync(True, recurse=True)
|
||||
else:
|
||||
yield
|
||||
|
||||
def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
|
||||
if not self.use_distributed:
|
||||
return super().clip_grad_norm(model, max_norm)
|
||||
|
||||
# FSDP params are DTensors (sharded across ranks).
|
||||
# torch.nn.utils.clip_grad_norm_ computes LOCAL norm per rank,
|
||||
# so we must all-reduce to get the global norm before clipping.
|
||||
local_norm = torch.nn.utils.get_total_norm(
|
||||
[p.grad for p in model.parameters() if p.grad is not None],
|
||||
)
|
||||
if isinstance(local_norm, DTensor):
|
||||
local_norm = local_norm.to_local()
|
||||
total_norm_sq = local_norm**2
|
||||
dist.all_reduce(total_norm_sq, op=dist.ReduceOp.SUM)
|
||||
total_norm = total_norm_sq.sqrt()
|
||||
|
||||
clip_coef = max_norm / (total_norm + 1e-6)
|
||||
clip_coef_clamped = torch.clamp(clip_coef, max=1.0)
|
||||
for p in model.parameters():
|
||||
if p.grad is not None:
|
||||
p.grad.mul_(clip_coef_clamped)
|
||||
|
||||
return total_norm.item()
|
||||
|
||||
def unwrap_model(self, model: nn.Module):
|
||||
if not self.use_distributed:
|
||||
return model.state_dict()
|
||||
|
||||
# unshard() and full_tensor() are collective ops — all ranks must
|
||||
# participate. Non-rank-0 ranks still call them but discard results.
|
||||
for module in model.modules():
|
||||
if isinstance(module, FSDPModule):
|
||||
module.unshard()
|
||||
|
||||
state_dict = model.state_dict()
|
||||
result = {}
|
||||
for k, v in state_dict.items():
|
||||
if isinstance(v, DTensor):
|
||||
full = v.full_tensor()
|
||||
if get_rank() == 0:
|
||||
result[k] = full
|
||||
elif get_rank() == 0:
|
||||
result[k] = v
|
||||
|
||||
for module in model.modules():
|
||||
if isinstance(module, FSDPModule):
|
||||
module.reshard()
|
||||
|
||||
if get_rank() != 0:
|
||||
return None
|
||||
|
||||
return result
|
||||
@@ -1,115 +0,0 @@
|
||||
from typing import Dict
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
class ParallelModel(nn.Module):
|
||||
def __init__(self, process_group: dist.ProcessGroup):
|
||||
super().__init__()
|
||||
self.process_group = process_group
|
||||
self.rank = dist.get_rank(self.process_group)
|
||||
self.world_size = dist.get_world_size(self.process_group)
|
||||
|
||||
|
||||
class RowParallelLinear(ParallelModel):
|
||||
def __init__(
|
||||
self,
|
||||
process_group: dist.ProcessGroup,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
bias: bool = True,
|
||||
reduce_results: bool = True,
|
||||
):
|
||||
super().__init__(process_group)
|
||||
|
||||
self.in_features = in_features
|
||||
self.out_features = out_features
|
||||
self.in_features_per_rank = in_features // self.world_size
|
||||
self.reduce_results = reduce_results
|
||||
|
||||
if in_features % self.world_size != 0:
|
||||
raise ValueError(
|
||||
f"in_features must be divisible by world_size. Got {in_features} and {self.world_size}"
|
||||
)
|
||||
|
||||
self.weight = nn.Parameter(torch.empty(out_features, self.in_features_per_rank))
|
||||
self.bias = nn.Parameter(torch.zeros(out_features)) if bias else None
|
||||
|
||||
def forward(self, input: Tensor) -> Tensor:
|
||||
output = F.linear(input, self.weight)
|
||||
|
||||
if self.reduce_results:
|
||||
dist.all_reduce(output, op=dist.ReduceOp.SUM, group=self.process_group)
|
||||
|
||||
if self.bias is not None:
|
||||
output += self.bias
|
||||
|
||||
return output
|
||||
|
||||
def load_state_dict(self, state_dict: Dict[str, Tensor]):
|
||||
full_weight = state_dict.get("weight")
|
||||
full_bias = state_dict.get("bias")
|
||||
|
||||
start_idx = self.rank * self.in_features_per_rank
|
||||
end_idx = start_idx + self.in_features_per_rank
|
||||
weight_slice = full_weight[:, start_idx:end_idx]
|
||||
self.weight.data.copy_(weight_slice)
|
||||
|
||||
if self.bias is not None:
|
||||
self.bias.data.copy_(full_bias)
|
||||
|
||||
|
||||
class ColumnParallelLinear(ParallelModel):
|
||||
def __init__(
|
||||
self,
|
||||
process_group: dist.ProcessGroup,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
bias: bool = True,
|
||||
gather_results: bool = True,
|
||||
):
|
||||
super().__init__(process_group)
|
||||
|
||||
self.in_features = in_features
|
||||
self.out_features = out_features
|
||||
self.out_features_per_rank = out_features // self.world_size
|
||||
self.gather_results = gather_results
|
||||
|
||||
if out_features % self.world_size != 0:
|
||||
raise ValueError(
|
||||
f"out_features must be divisible by world_size. Got {out_features} and {self.world_size}"
|
||||
)
|
||||
|
||||
self.weight = nn.Parameter(
|
||||
torch.empty(self.out_features_per_rank, self.in_features)
|
||||
)
|
||||
self.bias = (
|
||||
nn.Parameter(torch.zeros(self.out_features_per_rank)) if bias else None
|
||||
)
|
||||
|
||||
def forward(self, input: Tensor) -> Tensor:
|
||||
output = F.linear(input, self.weight, self.bias)
|
||||
|
||||
if self.gather_results:
|
||||
output_list = [torch.empty_like(output) for _ in range(self.world_size)]
|
||||
dist.all_gather(output_list, output, group=self.process_group)
|
||||
output = torch.cat(output_list, dim=-1)
|
||||
|
||||
return output
|
||||
|
||||
def load_state_dict(self, state_dict: Dict[str, Tensor]):
|
||||
full_weight = state_dict.get("weight")
|
||||
full_bias = state_dict.get("bias")
|
||||
|
||||
start_idx = self.rank * self.out_features_per_rank
|
||||
end_idx = start_idx + self.out_features_per_rank
|
||||
weight_slice = full_weight[start_idx:end_idx, :]
|
||||
self.weight.data.copy_(weight_slice)
|
||||
|
||||
if self.bias is not None:
|
||||
bias_slice = full_bias[start_idx:end_idx]
|
||||
self.bias.data.copy_(bias_slice)
|
||||
+174
-50
@@ -1,12 +1,27 @@
|
||||
import logging
|
||||
import os
|
||||
import signal
|
||||
import socket
|
||||
import threading
|
||||
from abc import ABC, abstractmethod
|
||||
from contextlib import contextmanager
|
||||
from functools import wraps
|
||||
from typing import Callable
|
||||
from typing import Callable, Optional
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.multiprocessing as mp
|
||||
|
||||
from astrai.signal_handler import install_early_signal_handlers
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def find_free_port() -> str:
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
s.bind(("", 0))
|
||||
return str(s.getsockname()[1])
|
||||
|
||||
|
||||
def get_current_device():
|
||||
return os.environ["LOCAL_DEVICE"]
|
||||
@@ -30,6 +45,7 @@ def get_rank() -> int:
|
||||
def setup_parallel(
|
||||
rank: int,
|
||||
world_size: int,
|
||||
local_rank: int,
|
||||
backend: str = "nccl",
|
||||
master_addr: str = "localhost",
|
||||
master_port: str = "29500",
|
||||
@@ -41,20 +57,26 @@ def setup_parallel(
|
||||
return
|
||||
|
||||
if world_size <= 1:
|
||||
device_id = torch.device(device_type, local_rank)
|
||||
os.environ["LOCAL_RANK"] = str(local_rank)
|
||||
os.environ["WORLD_SIZE"] = "1"
|
||||
os.environ["LOCAL_DEVICE"] = str(device_id)
|
||||
yield None
|
||||
return
|
||||
|
||||
device_id = torch.device(device_type, rank)
|
||||
device_id = torch.device(device_type, local_rank)
|
||||
|
||||
os.environ["MASTER_ADDR"] = master_addr
|
||||
os.environ["MASTER_PORT"] = master_port
|
||||
os.environ["LOCAL_RANK"] = str(rank)
|
||||
os.environ["LOCAL_RANK"] = str(local_rank)
|
||||
os.environ["WORLD_SIZE"] = str(world_size)
|
||||
os.environ["LOCAL_DEVICE"] = str(device_id)
|
||||
|
||||
dist.init_process_group(
|
||||
rank=rank, world_size=world_size, backend=backend, device_id=device_id
|
||||
)
|
||||
pg_kwargs = dict(rank=rank, world_size=world_size, backend=backend)
|
||||
if backend in ("nccl", "ccl"):
|
||||
pg_kwargs["device_id"] = device_id
|
||||
|
||||
dist.init_process_group(**pg_kwargs)
|
||||
|
||||
try:
|
||||
if backend == "nccl" and torch.cuda.is_available():
|
||||
@@ -90,7 +112,7 @@ def only_on_rank(rank, sync=False):
|
||||
return decorator
|
||||
|
||||
|
||||
def wrapper_spawn_func(
|
||||
def _run_single_rank(
|
||||
rank: int,
|
||||
world_size: int,
|
||||
backend: str,
|
||||
@@ -100,20 +122,143 @@ def wrapper_spawn_func(
|
||||
func: Callable,
|
||||
kwargs: dict,
|
||||
):
|
||||
try:
|
||||
install_early_signal_handlers()
|
||||
with setup_parallel(
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
local_rank=rank,
|
||||
backend=backend,
|
||||
master_addr=master_addr,
|
||||
master_port=master_port,
|
||||
device_type=device_type,
|
||||
):
|
||||
func(**kwargs)
|
||||
|
||||
|
||||
class LaunchStrategy(ABC):
|
||||
"""Strategy for launching a function in a distributed context."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
world_size: int,
|
||||
backend: str,
|
||||
master_addr: str,
|
||||
master_port: str,
|
||||
device_type: str,
|
||||
start_method: str,
|
||||
):
|
||||
self.world_size = world_size
|
||||
self.backend = backend
|
||||
self.master_addr = master_addr
|
||||
self.master_port = master_port
|
||||
self.device_type = device_type
|
||||
self.start_method = start_method
|
||||
|
||||
@abstractmethod
|
||||
def launch(self, func: Callable, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class TorchrunStrategy(LaunchStrategy):
|
||||
"""External orchestrator (torchrun, SLURM, K8s) — env vars pre-set."""
|
||||
|
||||
def launch(self, func: Callable, **kwargs):
|
||||
install_early_signal_handlers()
|
||||
rank = int(os.environ["RANK"])
|
||||
world_size = int(os.environ["WORLD_SIZE"])
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", rank))
|
||||
with setup_parallel(
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
backend=backend,
|
||||
master_addr=master_addr,
|
||||
master_port=master_port,
|
||||
device_type=device_type,
|
||||
local_rank=local_rank,
|
||||
backend=self.backend,
|
||||
master_addr=os.environ.get("MASTER_ADDR", self.master_addr),
|
||||
master_port=os.environ.get("MASTER_PORT", self.master_port),
|
||||
device_type=self.device_type,
|
||||
):
|
||||
func(**kwargs)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error in rank {rank}: {e}")
|
||||
raise
|
||||
|
||||
class LocalStrategy(LaunchStrategy):
|
||||
"""Local launcher — single-process or mp.start_processes."""
|
||||
|
||||
def launch(self, func: Callable, **kwargs):
|
||||
args = (
|
||||
self.world_size,
|
||||
self.backend,
|
||||
self.master_addr,
|
||||
self.master_port,
|
||||
self.device_type,
|
||||
func,
|
||||
kwargs,
|
||||
)
|
||||
|
||||
if self.world_size == 1:
|
||||
_run_single_rank(0, *args)
|
||||
return
|
||||
|
||||
install_early_signal_handlers()
|
||||
ctx = mp.start_processes(
|
||||
_run_single_rank,
|
||||
args=args,
|
||||
nprocs=self.world_size,
|
||||
start_method=self.start_method,
|
||||
join=False,
|
||||
)
|
||||
|
||||
parent_stop = threading.Event()
|
||||
original_handlers = {}
|
||||
|
||||
def _parent_handler(signum, frame):
|
||||
sig = signal.Signals(signum)
|
||||
logger.warning(
|
||||
"Parent (pid=%d) received %s, forwarding to children...",
|
||||
os.getpid(),
|
||||
sig.name,
|
||||
)
|
||||
parent_stop.set()
|
||||
for p in ctx.processes:
|
||||
if p.is_alive():
|
||||
p.terminate()
|
||||
|
||||
for sig in (signal.SIGTERM, signal.SIGINT):
|
||||
prev = signal.signal(sig, _parent_handler)
|
||||
if prev not in (signal.SIG_DFL, signal.SIG_IGN, None, _parent_handler):
|
||||
original_handlers[sig] = prev
|
||||
|
||||
try:
|
||||
while not ctx.join() and not parent_stop.is_set():
|
||||
pass
|
||||
except BaseException:
|
||||
logger.warning(
|
||||
"Parent received unexpected exception, terminating children..."
|
||||
)
|
||||
for p in ctx.processes:
|
||||
if p.is_alive():
|
||||
p.terminate()
|
||||
raise
|
||||
finally:
|
||||
for sig, handler in original_handlers.items():
|
||||
signal.signal(sig, handler)
|
||||
|
||||
for p in ctx.processes:
|
||||
p.join()
|
||||
|
||||
ctx.join()
|
||||
|
||||
|
||||
def _detect_launcher() -> str:
|
||||
"""Detect the distributed launcher from environment.
|
||||
|
||||
Returns one of: "torchelastic", "torchrun", "external", "local".
|
||||
"""
|
||||
if dist.is_torchelastic_launched():
|
||||
return "torchelastic"
|
||||
if "LOCAL_WORLD_SIZE" in os.environ:
|
||||
return "torchrun"
|
||||
if "RANK" in os.environ and "WORLD_SIZE" in os.environ:
|
||||
return "external"
|
||||
return "local"
|
||||
|
||||
|
||||
def spawn_parallel_fn(
|
||||
@@ -121,41 +266,20 @@ def spawn_parallel_fn(
|
||||
world_size: int,
|
||||
backend: str = "nccl",
|
||||
master_addr: str = "localhost",
|
||||
master_port: str = "29500",
|
||||
master_port: Optional[str] = None,
|
||||
device_type: str = "cuda",
|
||||
start_method: str = "spawn",
|
||||
**kwargs,
|
||||
):
|
||||
# clear environment variables
|
||||
for key in [
|
||||
"MASTER_ADDR",
|
||||
"MASTER_PORT",
|
||||
"RANK",
|
||||
"WORLD_SIZE",
|
||||
"LOCAL_RANK",
|
||||
"LOCAL_DEVICE",
|
||||
]:
|
||||
if key in os.environ:
|
||||
del os.environ[key]
|
||||
|
||||
if world_size == 1:
|
||||
device_id = torch.device(device_type, 0)
|
||||
os.environ["LOCAL_RANK"] = "0"
|
||||
os.environ["WORLD_SIZE"] = "1"
|
||||
os.environ["LOCAL_DEVICE"] = str(device_id)
|
||||
|
||||
func(**kwargs)
|
||||
return
|
||||
|
||||
wrapper_spawn_func_args = (
|
||||
world_size,
|
||||
backend,
|
||||
master_addr,
|
||||
master_port,
|
||||
device_type,
|
||||
func,
|
||||
kwargs,
|
||||
)
|
||||
|
||||
mp.spawn(
|
||||
wrapper_spawn_func, nprocs=world_size, args=wrapper_spawn_func_args, join=True
|
||||
)
|
||||
if master_port is None:
|
||||
master_port = find_free_port()
|
||||
launcher = _detect_launcher()
|
||||
if launcher in ("torchelastic", "torchrun", "external"):
|
||||
strategy = TorchrunStrategy(
|
||||
world_size, backend, master_addr, master_port, device_type, start_method
|
||||
)
|
||||
else:
|
||||
strategy = LocalStrategy(
|
||||
world_size, backend, master_addr, master_port, device_type, start_method
|
||||
)
|
||||
strategy.launch(func, **kwargs)
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
from astrai.preprocessing.builder import (
|
||||
BaseMaskBuilder,
|
||||
MaskBuilderFactory,
|
||||
MultiOutputMaskBuilder,
|
||||
SectionedMaskBuilder,
|
||||
SingleOutputMaskBuilder,
|
||||
)
|
||||
from astrai.preprocessing.packing import (
|
||||
PackingStrategy,
|
||||
PackingStrategyFactory,
|
||||
plan_bfd,
|
||||
)
|
||||
from astrai.preprocessing.pipeline import Pipeline, filter_by_length
|
||||
from astrai.preprocessing.position_id import (
|
||||
PositionIdStrategy,
|
||||
PositionIdStrategyFactory,
|
||||
)
|
||||
from astrai.preprocessing.transform import TokenizeTransform
|
||||
from astrai.preprocessing.writer import (
|
||||
StoreWriter,
|
||||
StoreWriterFactory,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"BaseMaskBuilder",
|
||||
"MaskBuilderFactory",
|
||||
"MultiOutputMaskBuilder",
|
||||
"PackingStrategy",
|
||||
"PackingStrategyFactory",
|
||||
"Pipeline",
|
||||
"PositionIdStrategy",
|
||||
"PositionIdStrategyFactory",
|
||||
"SectionedMaskBuilder",
|
||||
"SingleOutputMaskBuilder",
|
||||
"StoreWriter",
|
||||
"StoreWriterFactory",
|
||||
"TokenizeTransform",
|
||||
"filter_by_length",
|
||||
"plan_bfd",
|
||||
]
|
||||
@@ -0,0 +1,542 @@
|
||||
"""Mask building for preprocessing pipeline.
|
||||
|
||||
:class:`SectionRenderer` converts section specs into token ids and loss
|
||||
masks (template / text / value extraction). :class:`SingleOutputMaskBuilder`
|
||||
handles single-output (SFT / pretrain), :class:`MultiOutputMaskBuilder`
|
||||
handles multi-output (DPO / GRPO), and :class:`SectionedMaskBuilder`
|
||||
orchestrates both modes as a façade.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
|
||||
def _extract_domain(item: dict, domain_key: Optional[str]) -> str:
|
||||
if not domain_key:
|
||||
return "__default__"
|
||||
val = item.get(domain_key, "__default__")
|
||||
return val if isinstance(val, str) else "__default__"
|
||||
|
||||
|
||||
def _resolve_action(action: str, role: str, config) -> str:
|
||||
if action == "$role":
|
||||
return config.mask.get(role, config.mask_default)
|
||||
return action
|
||||
|
||||
|
||||
class SectionRenderer:
|
||||
"""Render section specs into ``(ids, loss_mask)`` tuples."""
|
||||
|
||||
def process_sections(
|
||||
self,
|
||||
item: dict,
|
||||
sections: list,
|
||||
config,
|
||||
tokenizer,
|
||||
*,
|
||||
is_top_level: bool = False,
|
||||
):
|
||||
all_ids: list[int] = []
|
||||
loss_mask: list[int] = []
|
||||
|
||||
has_template = any(s.get("template") for s in sections)
|
||||
is_text_config = not has_template and all(
|
||||
s["action"] == "train" for s in sections
|
||||
)
|
||||
|
||||
if is_top_level and has_template and tokenizer.bos_token_id is not None:
|
||||
all_ids.append(tokenizer.bos_token_id)
|
||||
loss_mask.append(0)
|
||||
|
||||
first_section = True
|
||||
for sec in sections:
|
||||
field = sec["field"]
|
||||
action = sec["action"]
|
||||
use_template = sec.get("template", False)
|
||||
add_special = sec.get(
|
||||
"add_special_tokens", not use_template and first_section
|
||||
)
|
||||
|
||||
if use_template:
|
||||
success = self._append_template(
|
||||
item, field, action, tokenizer, config, all_ids, loss_mask
|
||||
)
|
||||
if not success:
|
||||
continue
|
||||
else:
|
||||
success = self._append_text(
|
||||
item,
|
||||
field,
|
||||
action,
|
||||
tokenizer,
|
||||
add_special,
|
||||
is_text_config,
|
||||
config,
|
||||
all_ids,
|
||||
loss_mask,
|
||||
)
|
||||
if not success:
|
||||
continue
|
||||
|
||||
first_section = False
|
||||
|
||||
max_len = config.preprocessing.max_seq_len
|
||||
all_ids = all_ids[:max_len]
|
||||
loss_mask = loss_mask[: len(all_ids)]
|
||||
|
||||
if not all_ids:
|
||||
return None, None
|
||||
|
||||
if is_top_level and has_template and len(all_ids) <= 1:
|
||||
return None, None
|
||||
|
||||
return all_ids, loss_mask
|
||||
|
||||
def process_sections_batch(
|
||||
self,
|
||||
items: list[dict],
|
||||
sections: list,
|
||||
config,
|
||||
tokenizer,
|
||||
*,
|
||||
is_top_level=False,
|
||||
filter_text=True,
|
||||
):
|
||||
"""Render and tokenize a group of records with batched Rust tokenization."""
|
||||
has_template = any(s.get("template") for s in sections)
|
||||
is_text_config = not has_template and all(
|
||||
s["action"] == "train" for s in sections
|
||||
)
|
||||
plans: list[list[tuple[str, str, bool]]] = []
|
||||
|
||||
for item in items:
|
||||
plan: list[tuple[str, str, bool]] = []
|
||||
first_section = True
|
||||
for sec in sections:
|
||||
field = sec["field"]
|
||||
action = sec["action"]
|
||||
use_template = sec.get("template", False)
|
||||
add_special = sec.get(
|
||||
"add_special_tokens", not use_template and first_section
|
||||
)
|
||||
|
||||
if use_template:
|
||||
messages = item.get(field)
|
||||
if not isinstance(messages, list) or not messages:
|
||||
continue
|
||||
for msg in messages:
|
||||
role = msg.get("role", "")
|
||||
rendered = tokenizer.apply_chat_template(
|
||||
[msg], tokenize=False, add_generation_prompt=False
|
||||
)
|
||||
plan.append(
|
||||
(rendered, _resolve_action(action, role, config), False)
|
||||
)
|
||||
else:
|
||||
text = str(item.get(field, ""))
|
||||
if not text.strip():
|
||||
continue
|
||||
if is_text_config and filter_text:
|
||||
pp = config.preprocessing
|
||||
if pp.min_chars > 0 and len(text) < pp.min_chars:
|
||||
continue
|
||||
if len(text) > pp.max_chars:
|
||||
continue
|
||||
plan.append((text, action, add_special))
|
||||
|
||||
first_section = False
|
||||
plans.append(plan)
|
||||
|
||||
encoded: dict[tuple[int, int], list[int]] = {}
|
||||
for add_special in (False, True):
|
||||
refs = [
|
||||
(item_idx, unit_idx, text)
|
||||
for item_idx, plan in enumerate(plans)
|
||||
for unit_idx, (text, _, add) in enumerate(plan)
|
||||
if add == add_special
|
||||
]
|
||||
if not refs:
|
||||
continue
|
||||
ids_batch = tokenizer.encode(
|
||||
[text for _, _, text in refs], add_special_tokens=add_special
|
||||
)
|
||||
for (item_idx, unit_idx, _), ids in zip(refs, ids_batch):
|
||||
encoded[(item_idx, unit_idx)] = ids
|
||||
|
||||
outputs = []
|
||||
max_len = config.preprocessing.max_seq_len
|
||||
for item_idx, plan in enumerate(plans):
|
||||
all_ids = []
|
||||
loss_mask = []
|
||||
if is_top_level and has_template and tokenizer.bos_token_id is not None:
|
||||
all_ids.append(tokenizer.bos_token_id)
|
||||
loss_mask.append(0)
|
||||
for unit_idx, (_, action, _) in enumerate(plan):
|
||||
ids = encoded[(item_idx, unit_idx)]
|
||||
all_ids.extend(ids)
|
||||
loss_mask.extend([1 if action == "train" else 0] * len(ids))
|
||||
all_ids = all_ids[:max_len]
|
||||
loss_mask = loss_mask[: len(all_ids)]
|
||||
if not all_ids or (is_top_level and has_template and len(all_ids) <= 1):
|
||||
outputs.append((None, None))
|
||||
else:
|
||||
outputs.append((all_ids, loss_mask))
|
||||
return outputs
|
||||
|
||||
def process_list_field(self, item: dict, sections: list, config, tokenizer):
|
||||
"""Tokenize a list-valued field, preserving per-element boundaries.
|
||||
|
||||
Returns ``(list_of_id_lists, list_of_mask_lists)`` where each
|
||||
inner list corresponds to one element of the source list. This
|
||||
is critical for GRPO where each response must stay a separate
|
||||
sequence so the strategy can form a ``[G, R]`` tensor.
|
||||
"""
|
||||
per_item_ids: list[list[int]] = []
|
||||
per_item_masks: list[list[int]] = []
|
||||
|
||||
for sec in sections:
|
||||
field = sec["field"]
|
||||
action = sec["action"]
|
||||
use_template = sec.get("template", False)
|
||||
|
||||
values = item.get(field)
|
||||
if not isinstance(values, list):
|
||||
continue
|
||||
|
||||
for val in values:
|
||||
ids: list[int] = []
|
||||
mask: list[int] = []
|
||||
if use_template:
|
||||
if isinstance(val, list):
|
||||
wrapper = {field: val}
|
||||
self._append_template(
|
||||
wrapper, field, action, tokenizer, config, ids, mask
|
||||
)
|
||||
else:
|
||||
wrapper = {field: str(val)}
|
||||
self._append_text(
|
||||
wrapper,
|
||||
field,
|
||||
action,
|
||||
tokenizer,
|
||||
False,
|
||||
False,
|
||||
config,
|
||||
ids,
|
||||
mask,
|
||||
)
|
||||
if ids:
|
||||
max_len = config.preprocessing.max_seq_len
|
||||
ids = ids[:max_len]
|
||||
mask = mask[: len(ids)]
|
||||
per_item_ids.append(ids)
|
||||
per_item_masks.append(mask)
|
||||
|
||||
if not per_item_ids:
|
||||
return None, None
|
||||
return per_item_ids, per_item_masks
|
||||
|
||||
def process_list_field_batch(self, items, sections, config, tokenizer):
|
||||
per_item_ids = [[] for _ in items]
|
||||
per_item_masks = [[] for _ in items]
|
||||
|
||||
for sec in sections:
|
||||
wrappers = []
|
||||
owners = []
|
||||
field = sec["field"]
|
||||
for item_idx, item in enumerate(items):
|
||||
values = item.get(field)
|
||||
if not isinstance(values, list):
|
||||
continue
|
||||
for val in values:
|
||||
if sec.get("template", False) and not isinstance(val, list):
|
||||
continue
|
||||
wrappers.append({field: val if isinstance(val, list) else str(val)})
|
||||
owners.append(item_idx)
|
||||
|
||||
rendered = self.process_sections_batch(
|
||||
wrappers,
|
||||
[sec],
|
||||
config,
|
||||
tokenizer,
|
||||
is_top_level=False,
|
||||
filter_text=False,
|
||||
)
|
||||
for owner, (ids, mask) in zip(owners, rendered):
|
||||
if ids:
|
||||
per_item_ids[owner].append(ids)
|
||||
per_item_masks[owner].append(mask)
|
||||
|
||||
return [
|
||||
(ids, masks) if ids else (None, None)
|
||||
for ids, masks in zip(per_item_ids, per_item_masks)
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def is_value_section(sections: list) -> bool:
|
||||
return len(sections) == 1 and sections[0].get("action") == "value"
|
||||
|
||||
@staticmethod
|
||||
def extract_raw_value(item: dict, sections: list):
|
||||
sec = sections[0]
|
||||
field = sec["field"]
|
||||
raw = item.get(field)
|
||||
if raw is None:
|
||||
return None
|
||||
if isinstance(raw, list):
|
||||
return [float(v) for v in raw]
|
||||
return [float(raw)]
|
||||
|
||||
def _append_template(
|
||||
self, item, field, action, tokenizer, config, all_ids, loss_mask
|
||||
):
|
||||
messages = item.get(field)
|
||||
if not isinstance(messages, list) or not messages:
|
||||
return False
|
||||
for msg in messages:
|
||||
role = msg.get("role", "")
|
||||
act = _resolve_action(action, role, config)
|
||||
rendered = tokenizer.apply_chat_template(
|
||||
[msg], tokenize=False, add_generation_prompt=False
|
||||
)
|
||||
ids = tokenizer.encode(rendered, add_special_tokens=False)
|
||||
all_ids.extend(ids)
|
||||
val = 1 if act == "train" else 0
|
||||
loss_mask.extend([val] * len(ids))
|
||||
return True
|
||||
|
||||
def _append_text(
|
||||
self,
|
||||
item,
|
||||
field,
|
||||
action,
|
||||
tokenizer,
|
||||
add_special,
|
||||
is_text_config,
|
||||
config,
|
||||
all_ids,
|
||||
loss_mask,
|
||||
):
|
||||
text = str(item.get(field, ""))
|
||||
if not text.strip():
|
||||
return False
|
||||
if is_text_config:
|
||||
pp = config.preprocessing
|
||||
if pp.min_chars > 0 and len(text) < pp.min_chars:
|
||||
return False
|
||||
if len(text) > pp.max_chars:
|
||||
return False
|
||||
ids = tokenizer.encode(text, add_special_tokens=add_special)
|
||||
all_ids.extend(ids)
|
||||
val = 1 if action == "train" else 0
|
||||
loss_mask.extend([val] * len(ids))
|
||||
return True
|
||||
|
||||
|
||||
class BaseMaskBuilder(ABC):
|
||||
"""Convert a JSONL item into token ids and optional loss_mask."""
|
||||
|
||||
@abstractmethod
|
||||
def build(self, item: dict, config, tokenizer) -> Optional[dict]: ...
|
||||
|
||||
def build_batch(self, items: list[dict], config, tokenizer) -> list[Optional[dict]]:
|
||||
return [self.build(item, config, tokenizer) for item in items]
|
||||
|
||||
|
||||
class MaskBuilderFactory(BaseFactory["BaseMaskBuilder"]):
|
||||
pass
|
||||
|
||||
|
||||
@MaskBuilderFactory.register("single")
|
||||
class SingleOutputMaskBuilder(BaseMaskBuilder):
|
||||
"""Build a single output sequence with optional loss mask.
|
||||
|
||||
Expects ``config.input.sections`` (list of section specs).
|
||||
"""
|
||||
|
||||
def __init__(self, renderer: Optional[SectionRenderer] = None):
|
||||
self.renderer = renderer or SectionRenderer()
|
||||
|
||||
def build(self, item: dict, config, tokenizer) -> Optional[dict]:
|
||||
sections = config.input.sections
|
||||
if not sections:
|
||||
return None
|
||||
|
||||
ids, mask = self.renderer.process_sections(
|
||||
item, sections, config, tokenizer, is_top_level=True
|
||||
)
|
||||
if ids is None:
|
||||
return None
|
||||
|
||||
result: dict = {
|
||||
"sequence": ids,
|
||||
"domain": _extract_domain(item, config.output.domain_key),
|
||||
}
|
||||
if not all(m == 1 for m in mask):
|
||||
result["loss_mask"] = mask
|
||||
return result
|
||||
|
||||
def build_batch(self, items, config, tokenizer):
|
||||
sections = config.input.sections
|
||||
if not sections:
|
||||
return [None] * len(items)
|
||||
rendered = self.renderer.process_sections_batch(
|
||||
items, sections, config, tokenizer, is_top_level=True
|
||||
)
|
||||
results = []
|
||||
for item, (ids, mask) in zip(items, rendered):
|
||||
if ids is None:
|
||||
results.append(None)
|
||||
continue
|
||||
result = {
|
||||
"sequence": ids,
|
||||
"domain": _extract_domain(item, config.output.domain_key),
|
||||
}
|
||||
if not all(m == 1 for m in mask):
|
||||
result["loss_mask"] = mask
|
||||
results.append(result)
|
||||
return results
|
||||
|
||||
|
||||
@MaskBuilderFactory.register("multi")
|
||||
class MultiOutputMaskBuilder(BaseMaskBuilder):
|
||||
"""Build multiple output sequences (DPO / GRPO).
|
||||
|
||||
Expects ``config.input.sources`` (dict of output_key → spec).
|
||||
"""
|
||||
|
||||
def __init__(self, renderer: Optional[SectionRenderer] = None):
|
||||
self.renderer = renderer or SectionRenderer()
|
||||
|
||||
def build(self, item: dict, config, tokenizer) -> Optional[dict]:
|
||||
sources_spec = getattr(config.input, "sources", None)
|
||||
if not sources_spec:
|
||||
return None
|
||||
|
||||
result: dict = {}
|
||||
required_outputs = {
|
||||
output_key
|
||||
for output_key, spec in sources_spec.items()
|
||||
if spec.get("sections")
|
||||
}
|
||||
|
||||
for output_key, spec in sources_spec.items():
|
||||
sections = spec.get("sections", [])
|
||||
if not sections:
|
||||
continue
|
||||
|
||||
if self.renderer.is_value_section(sections):
|
||||
ids = self.renderer.extract_raw_value(item, sections)
|
||||
if ids is None:
|
||||
continue
|
||||
result[output_key] = ids
|
||||
continue
|
||||
|
||||
list_field = spec.get("list_field", False)
|
||||
mask_key = spec.get("mask_key", f"{output_key}_mask")
|
||||
|
||||
if list_field:
|
||||
ids, mask = self.renderer.process_list_field(
|
||||
item, sections, config, tokenizer
|
||||
)
|
||||
if ids is None:
|
||||
continue
|
||||
# ids is List[List[int]] — preserve per-response structure
|
||||
result[output_key] = ids
|
||||
if mask is not None:
|
||||
result[mask_key] = mask
|
||||
continue
|
||||
|
||||
ids, mask = self.renderer.process_sections(
|
||||
item, sections, config, tokenizer, is_top_level=True
|
||||
)
|
||||
|
||||
if ids is None:
|
||||
continue
|
||||
|
||||
result[output_key] = ids
|
||||
if not all(m == 1 for m in mask):
|
||||
result[mask_key] = mask
|
||||
elif "mask_key" in spec:
|
||||
result[mask_key] = mask
|
||||
|
||||
if not required_outputs or not required_outputs.issubset(result):
|
||||
return None
|
||||
|
||||
result["domain"] = _extract_domain(item, config.output.domain_key)
|
||||
return result
|
||||
|
||||
def build_batch(self, items, config, tokenizer):
|
||||
sources_spec = getattr(config.input, "sources", None)
|
||||
if not sources_spec:
|
||||
return [None] * len(items)
|
||||
|
||||
results = [{} for _ in items]
|
||||
required_outputs = {
|
||||
output_key
|
||||
for output_key, spec in sources_spec.items()
|
||||
if spec.get("sections")
|
||||
}
|
||||
for output_key, spec in sources_spec.items():
|
||||
sections = spec.get("sections", [])
|
||||
if not sections:
|
||||
continue
|
||||
if self.renderer.is_value_section(sections):
|
||||
for item, result in zip(items, results):
|
||||
value = self.renderer.extract_raw_value(item, sections)
|
||||
if value is not None:
|
||||
result[output_key] = value
|
||||
continue
|
||||
|
||||
mask_key = spec.get("mask_key", f"{output_key}_mask")
|
||||
if spec.get("list_field", False):
|
||||
rendered = self.renderer.process_list_field_batch(
|
||||
items, sections, config, tokenizer
|
||||
)
|
||||
else:
|
||||
rendered = self.renderer.process_sections_batch(
|
||||
items, sections, config, tokenizer, is_top_level=True
|
||||
)
|
||||
|
||||
for result, (ids, mask) in zip(results, rendered):
|
||||
if ids is None:
|
||||
continue
|
||||
result[output_key] = ids
|
||||
if spec.get("list_field", False) or not all(m == 1 for m in mask):
|
||||
result[mask_key] = mask
|
||||
elif "mask_key" in spec:
|
||||
result[mask_key] = mask
|
||||
|
||||
return [
|
||||
({**result, "domain": _extract_domain(item, config.output.domain_key)})
|
||||
if required_outputs and required_outputs.issubset(result)
|
||||
else None
|
||||
for item, result in zip(items, results)
|
||||
]
|
||||
|
||||
|
||||
@MaskBuilderFactory.register("sectioned")
|
||||
class SectionedMaskBuilder(BaseMaskBuilder):
|
||||
"""Façade that dispatches to SingleOutputMaskBuilder or MultiOutputMaskBuilder.
|
||||
|
||||
Preserves backward compatibility for existing configs and code that rely
|
||||
on the ``"sectioned"`` factory name.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._single = SingleOutputMaskBuilder()
|
||||
self._multi = MultiOutputMaskBuilder()
|
||||
|
||||
def build(self, item: dict, config, tokenizer) -> Optional[dict]:
|
||||
sources_spec = getattr(config.input, "sources", None)
|
||||
if sources_spec:
|
||||
return self._multi.build(item, config, tokenizer)
|
||||
return self._single.build(item, config, tokenizer)
|
||||
|
||||
def build_batch(self, items, config, tokenizer):
|
||||
sources_spec = getattr(config.input, "sources", None)
|
||||
if sources_spec:
|
||||
return self._multi.build_batch(items, config, tokenizer)
|
||||
return self._single.build_batch(items, config, tokenizer)
|
||||
@@ -0,0 +1,124 @@
|
||||
"""Shared preprocessing kernel used by both :class:`Pipeline` and
|
||||
:class:`TokenizeTransform`.
|
||||
|
||||
The two entry points previously duplicated ~60 % of their logic:
|
||||
record iteration, mask-builder invocation, primary-id extraction,
|
||||
per-key accumulation, dtype inference and position-id generation.
|
||||
This module factors out the common core as pure functions so that
|
||||
the online (``TokenizeTransform``) and offline (``Pipeline``) paths
|
||||
stay in lockstep.
|
||||
"""
|
||||
|
||||
from itertools import chain
|
||||
from typing import Dict, Iterator, List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
from astrai.preprocessing.builder import MaskBuilderFactory
|
||||
from astrai.preprocessing.position_id import PositionIdStrategyFactory
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
|
||||
def build_preprocessing_components(config: PipelineConfig, tokenizer_path: str):
|
||||
"""Load tokenizer, mask builder and position-id strategy together.
|
||||
|
||||
Both ``Pipeline`` and ``TokenizeTransform`` need the same triple;
|
||||
centralising the construction avoids drift (e.g. one path forgetting
|
||||
to create the position-id strategy).
|
||||
"""
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
||||
mask_builder = MaskBuilderFactory.create("sectioned")
|
||||
position_strategy = PositionIdStrategyFactory.create(
|
||||
config.output.position_ids_mode
|
||||
)
|
||||
return tokenizer, mask_builder, position_strategy
|
||||
|
||||
|
||||
def primary_ids(result: dict) -> List[int]:
|
||||
"""Return the first flat int-list value in *result*.
|
||||
|
||||
Used for token counting and position-id generation when the
|
||||
primary key name is not known (DPO uses ``chosen``, GRPO uses
|
||||
``prompts``, SFT uses ``sequence``).
|
||||
"""
|
||||
for val in result.values():
|
||||
if isinstance(val, list) and val and isinstance(val[0], int):
|
||||
return val
|
||||
return []
|
||||
|
||||
|
||||
def infer_dtype(ids: List) -> torch.dtype:
|
||||
"""Float values become float32, everything else int32."""
|
||||
if ids and isinstance(ids[0], float):
|
||||
return torch.float32
|
||||
return torch.int32
|
||||
|
||||
|
||||
def iter_raw_records(
|
||||
records: List[dict],
|
||||
mask_builder,
|
||||
config: PipelineConfig,
|
||||
tokenizer,
|
||||
) -> Iterator[dict]:
|
||||
"""Yield mask-builder output dicts for each record, skipping failures.
|
||||
|
||||
Drops ``domain`` from the result (callers that need it should read
|
||||
it before calling this). Each yielded dict maps a key
|
||||
(``sequence``, ``chosen``, ``responses``…) to either a flat
|
||||
``List[int]`` or a nested ``List[List[int]]`` (GRPO responses/masks).
|
||||
"""
|
||||
for item in records:
|
||||
result = mask_builder.build(item, config, tokenizer)
|
||||
if result is None:
|
||||
continue
|
||||
result.pop("domain", None)
|
||||
if not primary_ids(result):
|
||||
continue
|
||||
yield result
|
||||
|
||||
|
||||
def to_per_record_tensors(
|
||||
raw: Dict[str, list],
|
||||
) -> Dict[str, List[torch.Tensor]]:
|
||||
"""Convert an accumulated ``{key: [per-record ids]}`` dict to tensors.
|
||||
|
||||
Handles three shapes transparently:
|
||||
|
||||
- ``List[int]`` per record (``sequence``, ``chosen``…) → one tensor per record.
|
||||
- ``List[List[int]]`` per record (GRPO ``responses``/``masks``) → one
|
||||
``List[Tensor]`` per record (nested), preserving the per-response
|
||||
boundary so downstream code can index responses individually.
|
||||
- ``List[int]`` for the whole shard (pre-packed keys) → single tensor.
|
||||
|
||||
The detection mirrors the previous inline logic in
|
||||
``Pipeline._flush`` and ``TokenizeTransform.apply``.
|
||||
"""
|
||||
tensors: Dict[str, List[torch.Tensor]] = {}
|
||||
for key, ids_list in raw.items():
|
||||
if ids_list and isinstance(ids_list[0], list):
|
||||
tensors[key] = [
|
||||
[torch.tensor(sub, dtype=infer_dtype(sub)) for sub in ids]
|
||||
if ids and isinstance(ids[0], list)
|
||||
else torch.tensor(ids, dtype=infer_dtype(ids))
|
||||
for ids in ids_list
|
||||
]
|
||||
else:
|
||||
tensors[key] = [
|
||||
torch.tensor(list(chain.from_iterable(ids_list)), dtype=torch.int32)
|
||||
]
|
||||
return tensors
|
||||
|
||||
|
||||
def build_position_ids(
|
||||
sequences: List[List[int]],
|
||||
strategy,
|
||||
) -> Optional[List[int]]:
|
||||
"""Generate position ids for *sequences* using *strategy*.
|
||||
|
||||
Returns ``None`` when the strategy produces no ids (e.g. ``none``
|
||||
mode), so callers can skip attaching the key instead of storing
|
||||
an empty list.
|
||||
"""
|
||||
pos_ids = strategy.generate(sequences)
|
||||
return pos_ids or None
|
||||
@@ -0,0 +1,176 @@
|
||||
"""Sequence packing strategies for shard-level reordering and truncation.
|
||||
|
||||
Each strategy receives the accumulated ``{key: [list of token lists]}``
|
||||
dict for a shard and returns a reordered / truncated version. The
|
||||
pipeline later flattens the result into contiguous tensors.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Dict, List
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
|
||||
def _truncate(seq: List[int], max_len: int, mode: str) -> List[int]:
|
||||
if len(seq) <= max_len:
|
||||
return seq
|
||||
if mode == "keep_end":
|
||||
return seq[-max_len:]
|
||||
return seq[:max_len]
|
||||
|
||||
|
||||
def plan_bfd(
|
||||
sequences: List[List[int]], max_packed_len: int, truncation_mode: str = "keep_start"
|
||||
) -> List[List[int]]:
|
||||
"""Best-Fit Decreasing bin packing of *sequences* into bins.
|
||||
|
||||
Returns a list of bins, each bin a list of original indices into
|
||||
*sequences*. Bin capacities are respected on the *truncated*
|
||||
length of each sequence (so a sequence longer than
|
||||
*max_packed_len* counts at *max_packed_len*).
|
||||
|
||||
Pure index-based so callers can apply the same plan to any
|
||||
aligned key (``loss_mask``, ``position_ids``…).
|
||||
"""
|
||||
n = len(sequences)
|
||||
order = sorted(range(n), key=lambda i: len(sequences[i]), reverse=True)
|
||||
bins: List[List[int]] = []
|
||||
bin_lengths: List[int] = []
|
||||
|
||||
for orig_idx in order:
|
||||
seq_len = len(_truncate(sequences[orig_idx], max_packed_len, truncation_mode))
|
||||
best_bin = None
|
||||
best_remain = max_packed_len + 1
|
||||
for i, bl in enumerate(bin_lengths):
|
||||
remain = max_packed_len - bl
|
||||
if seq_len <= remain < best_remain:
|
||||
best_remain = remain
|
||||
best_bin = i
|
||||
if best_bin is not None:
|
||||
bins[best_bin].append(orig_idx)
|
||||
bin_lengths[best_bin] += seq_len
|
||||
else:
|
||||
bins.append([orig_idx])
|
||||
bin_lengths.append(seq_len)
|
||||
|
||||
return bins
|
||||
|
||||
|
||||
class PackingStrategy(ABC):
|
||||
"""Reorder and truncate sequences within a shard."""
|
||||
|
||||
@abstractmethod
|
||||
def apply(
|
||||
self,
|
||||
keys: Dict[str, List[List[int]]],
|
||||
max_packed_len: int,
|
||||
truncation_mode: str,
|
||||
) -> Dict[str, List[List[int]]]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class PackingStrategyFactory(BaseFactory["PackingStrategy"]):
|
||||
pass
|
||||
|
||||
|
||||
@PackingStrategyFactory.register("simple")
|
||||
class SimplePacking(PackingStrategy):
|
||||
def apply(
|
||||
self,
|
||||
keys: Dict[str, List[List[int]]],
|
||||
max_packed_len: int,
|
||||
truncation_mode: str,
|
||||
) -> Dict[str, List[List[int]]]:
|
||||
return {
|
||||
k: [_truncate(v, max_packed_len, truncation_mode) for v in vals]
|
||||
for k, vals in keys.items()
|
||||
}
|
||||
|
||||
|
||||
@PackingStrategyFactory.register("bfd")
|
||||
class BFDPacking(PackingStrategy):
|
||||
"""Best-Fit Decreasing bin packing.
|
||||
|
||||
Assigns sequences to bins using a best-fit heuristic (sorted by
|
||||
decreasing length) and concatenates sequences within each bin into
|
||||
a single packed sequence. Packed sequences are truncated to
|
||||
*max_packed_len* so that each packed bin fits within one context
|
||||
window during training.
|
||||
"""
|
||||
|
||||
def apply(
|
||||
self,
|
||||
keys: Dict[str, List[List[int]]],
|
||||
max_packed_len: int,
|
||||
truncation_mode: str,
|
||||
) -> Dict[str, List[List[int]]]:
|
||||
sequences = keys.get("sequence", [])
|
||||
if not sequences:
|
||||
return keys
|
||||
bins = plan_bfd(sequences, max_packed_len, truncation_mode)
|
||||
|
||||
packed: Dict[str, List[List[int]]] = {}
|
||||
for k, vals in keys.items():
|
||||
packed[k] = [
|
||||
_truncate(
|
||||
self._concat_bin(vals, bin_indices),
|
||||
max_packed_len,
|
||||
truncation_mode,
|
||||
)
|
||||
for bin_indices in bins
|
||||
]
|
||||
return packed
|
||||
|
||||
@staticmethod
|
||||
def _concat_bin(vals: List[List[int]], indices: List[int]) -> List[int]:
|
||||
result: List[int] = []
|
||||
for i in indices:
|
||||
result.extend(vals[i])
|
||||
return result
|
||||
|
||||
|
||||
@PackingStrategyFactory.register("bfd_split")
|
||||
class BFDSplitPacking(BFDPacking):
|
||||
"""BFD packing with over-length sequences split into chunks.
|
||||
|
||||
Sequences longer than *max_packed_len* are split into consecutive
|
||||
chunks of at most *max_packed_len* tokens instead of being
|
||||
truncated. Each chunk becomes an independent sequence that enters
|
||||
BFD planning. All keys (``loss_mask``, ``position_ids``, …) are
|
||||
split in lockstep so per-token alignment is preserved.
|
||||
|
||||
Note: because each chunk is treated as a separate document, the
|
||||
second chunk of a split sequence loses the preceding context.
|
||||
"""
|
||||
|
||||
def apply(
|
||||
self,
|
||||
keys: Dict[str, List[List[int]]],
|
||||
max_packed_len: int,
|
||||
truncation_mode: str,
|
||||
) -> Dict[str, List[List[int]]]:
|
||||
sequences = keys.get("sequence", [])
|
||||
if not sequences:
|
||||
return keys
|
||||
if max_packed_len <= 0:
|
||||
return super().apply(keys, max_packed_len, truncation_mode)
|
||||
|
||||
split_keys = self._split_all(keys, max_packed_len)
|
||||
return super().apply(split_keys, max_packed_len, truncation_mode)
|
||||
|
||||
@staticmethod
|
||||
def _split_all(
|
||||
keys: Dict[str, List[List[int]]], max_packed_len: int
|
||||
) -> Dict[str, List[List[int]]]:
|
||||
"""Split every sequence exceeding *max_packed_len* into chunks,
|
||||
applying the same chunk boundaries to all keys."""
|
||||
sequences = keys["sequence"]
|
||||
chunk_bounds = [list(range(0, len(s), max_packed_len)) for s in sequences]
|
||||
result: Dict[str, List[List[int]]] = {}
|
||||
for key, vals in keys.items():
|
||||
split_vals: List[List[int]] = []
|
||||
for val, starts in zip(vals, chunk_bounds):
|
||||
for start in starts:
|
||||
split_vals.append(val[start : start + max_packed_len])
|
||||
result[key] = split_vals
|
||||
return result
|
||||
@@ -0,0 +1,277 @@
|
||||
"""Config-driven JSONL preprocessing pipeline.
|
||||
|
||||
Composes a :class:`BaseMaskBuilder` (selected by ``input.type``) with
|
||||
sharding and flush to ``.bin`` storage. Packing, position-id
|
||||
generation and storage writing are each delegated to pluggable strategies,
|
||||
dispatched by configuration keys.
|
||||
|
||||
Record iteration, mask building, primary-id extraction and per-key
|
||||
accumulation are shared with :class:`TokenizeTransform` via the
|
||||
:mod:`astrai.preprocessing.core` helpers.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from collections import defaultdict
|
||||
from itertools import chain
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
import torch
|
||||
import tqdm
|
||||
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
from astrai.preprocessing.core import (
|
||||
build_preprocessing_components,
|
||||
primary_ids,
|
||||
)
|
||||
from astrai.preprocessing.packing import PackingStrategyFactory
|
||||
from astrai.preprocessing.writer import StoreWriterFactory
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_STR_TO_DTYPE: dict[str, torch.dtype] = {
|
||||
"bool": torch.bool,
|
||||
"uint8": torch.uint8,
|
||||
"int8": torch.int8,
|
||||
"int16": torch.int16,
|
||||
"int32": torch.int32,
|
||||
"int64": torch.int64,
|
||||
"float16": torch.float16,
|
||||
"float32": torch.float32,
|
||||
"float64": torch.float64,
|
||||
}
|
||||
|
||||
|
||||
def filter_by_length(text: str, min_len: int = 50, max_len: int = 2_000_000) -> bool:
|
||||
return min_len <= len(text) <= max_len
|
||||
|
||||
|
||||
class Pipeline:
|
||||
"""Tokenization pipeline driven by a declarative :class:`PipelineConfig`.
|
||||
|
||||
Usage::
|
||||
|
||||
config = PipelineConfig.from_file("sft_pipeline.json")
|
||||
Pipeline(config, ["data.jsonl"], output_dir="out", tokenizer_path="params").run()
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: PipelineConfig,
|
||||
input_paths: list[str],
|
||||
output_dir: str,
|
||||
tokenizer_path: str,
|
||||
):
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
self.config = config
|
||||
self.paths = input_paths
|
||||
self.output_dir = output_dir
|
||||
self.tokenizer_path = tokenizer_path
|
||||
|
||||
self.tokenizer, self.mask_builder, self._position_id = (
|
||||
build_preprocessing_components(config, tokenizer_path)
|
||||
)
|
||||
self._packer = PackingStrategyFactory.create(
|
||||
config.preprocessing.packing_strategy
|
||||
)
|
||||
self._writer = StoreWriterFactory.create(config.output.storage_format)
|
||||
|
||||
def transform(self, item: dict) -> Optional[dict]:
|
||||
return self.mask_builder.build(item, self.config, self.tokenizer)
|
||||
|
||||
def transform_batch(self, items: list[dict]) -> list[Optional[dict]]:
|
||||
return self.mask_builder.build_batch(items, self.config, self.tokenizer)
|
||||
|
||||
def run(self):
|
||||
domains: dict = defaultdict(lambda: defaultdict(list))
|
||||
total_tokens = 0
|
||||
shard_idx: dict[str, int] = defaultdict(int)
|
||||
count = 0
|
||||
|
||||
pp = self.config.preprocessing
|
||||
|
||||
progress = tqdm.tqdm(desc="Tokenizing", unit="docs", mininterval=0.5)
|
||||
stop = False
|
||||
for items in self._iter_batches(pp.batch_size):
|
||||
progress.update(len(items))
|
||||
try:
|
||||
results = self.transform_batch(items)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to process batch, retrying records individually",
|
||||
exc_info=True,
|
||||
)
|
||||
results = []
|
||||
for item in items:
|
||||
try:
|
||||
results.append(self.transform(item))
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to process item, skipping", exc_info=True
|
||||
)
|
||||
results.append(None)
|
||||
|
||||
for result in results:
|
||||
if pp.max_items and count >= pp.max_items:
|
||||
stop = True
|
||||
break
|
||||
if result is None:
|
||||
continue
|
||||
|
||||
domain = result.pop("domain", "__default__")
|
||||
ids = primary_ids(result)
|
||||
if not ids:
|
||||
continue
|
||||
|
||||
bucket = domains[domain]
|
||||
self._align_bucket(bucket, result, ids)
|
||||
for key, val in result.items():
|
||||
bucket[key].append(val)
|
||||
|
||||
count += 1
|
||||
total_tokens += len(ids)
|
||||
|
||||
if total_tokens >= self.config.output.max_tokens_per_shard:
|
||||
self._flush(domains, shard_idx)
|
||||
domains.clear()
|
||||
total_tokens = 0
|
||||
if stop:
|
||||
break
|
||||
|
||||
progress.close()
|
||||
|
||||
if total_tokens > 0:
|
||||
self._flush(domains, shard_idx)
|
||||
|
||||
@staticmethod
|
||||
def _align_bucket(bucket: dict, result: dict, ids: list):
|
||||
"""Pad previously-accumulated keys that are missing from *result*."""
|
||||
for key in list(bucket.keys()):
|
||||
if key in result:
|
||||
continue
|
||||
bucket[key].append([0] * len(ids))
|
||||
|
||||
def _iter_items(self):
|
||||
for path in self.paths:
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
if path.endswith(".json"):
|
||||
data = json.load(f)
|
||||
if isinstance(data, dict):
|
||||
yield data
|
||||
elif isinstance(data, list):
|
||||
yield from data
|
||||
else:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
yield json.loads(line)
|
||||
|
||||
def _iter_batches(self, batch_size: int):
|
||||
batch_size = max(1, batch_size)
|
||||
batch = []
|
||||
for item in self._iter_items():
|
||||
batch.append(item)
|
||||
if len(batch) >= batch_size:
|
||||
yield batch
|
||||
batch = []
|
||||
if batch:
|
||||
yield batch
|
||||
|
||||
def _flush(self, domains, shard_idx):
|
||||
for domain, keys in domains.items():
|
||||
idx = shard_idx[domain]
|
||||
|
||||
pp = self.config.preprocessing
|
||||
original_sequences = keys.get("sequence", [])
|
||||
mode = self.config.output.position_ids_mode
|
||||
|
||||
keys = self._inject_doc_reset_position_ids(keys, mode, original_sequences)
|
||||
keys = self._packer.apply(dict(keys), pp.max_packed_len, pp.truncation_mode)
|
||||
tensors = self._to_tensors(keys)
|
||||
tensors = self._inject_continuous_position_ids(
|
||||
tensors, mode, keys.get("sequence", [])
|
||||
)
|
||||
|
||||
self._writer.save(self.output_dir, domain, idx, tensors)
|
||||
shard_idx[domain] = idx + 1
|
||||
|
||||
first_key = "sequence" if "sequence" in tensors else next(iter(tensors))
|
||||
tqdm.tqdm.write(
|
||||
f" saved {domain}/shard_{idx:04d} "
|
||||
f"({tensors[first_key][0].numel():,} tokens)"
|
||||
)
|
||||
|
||||
def _inject_doc_reset_position_ids(
|
||||
self,
|
||||
keys: Dict[str, list],
|
||||
mode: str,
|
||||
original_sequences: List[List[int]],
|
||||
) -> Dict[str, list]:
|
||||
"""Attach per-document position_ids before packing (``doc_reset``).
|
||||
|
||||
``doc_reset`` position ids must enter the packer so that each
|
||||
packed bin concatenates the per-doc ranges in bin order. The
|
||||
per-record structure ``[range(len(s)) for s in seqs]`` is required
|
||||
by the packer (it concatenates per-record lists per bin); the
|
||||
``PositionIdStrategy.generate`` flattens, so it cannot be used
|
||||
directly here — it is only consulted for the ``continuous``
|
||||
post-packing path.
|
||||
"""
|
||||
if mode != "doc_reset" or not original_sequences:
|
||||
return keys
|
||||
keys["position_ids"] = [list(range(len(s))) for s in original_sequences]
|
||||
return keys
|
||||
|
||||
def _inject_continuous_position_ids(
|
||||
self,
|
||||
tensors: Dict[str, List[torch.Tensor]],
|
||||
mode: str,
|
||||
packed_sequences: List[List[int]],
|
||||
) -> Dict[str, List[torch.Tensor]]:
|
||||
"""Attach a single continuous position_ids tensor after packing.
|
||||
|
||||
``continuous`` mode spans the whole shard (post-packing), so it
|
||||
cannot participate in bin packing — it is computed from the
|
||||
packed sequences and appended directly to the tensor dict.
|
||||
"""
|
||||
if mode != "continuous" or not packed_sequences:
|
||||
return tensors
|
||||
pos_ids = self._position_id.generate(packed_sequences)
|
||||
if pos_ids:
|
||||
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
||||
return tensors
|
||||
|
||||
def _to_tensors(self, keys: Dict[str, list]) -> Dict[str, List[torch.Tensor]]:
|
||||
"""Convert packed per-key id lists to tensors.
|
||||
|
||||
Honours ``config.output.dtype`` overrides per key; falls back to
|
||||
``int32``. Handles three shapes (see
|
||||
:func:`astrai.preprocessing.core.to_per_record_tensors` for the
|
||||
equivalent online-path helper):
|
||||
- ``List[int]`` per record → one tensor per record.
|
||||
- ``List[List[int]]`` per record (GRPO responses/masks) → one tensor
|
||||
per record, inner lists flattened.
|
||||
- ``List[int]`` for the whole shard (pre-packed keys) → single tensor.
|
||||
"""
|
||||
tensors: Dict[str, List[torch.Tensor]] = {}
|
||||
for key, ids_list in keys.items():
|
||||
dt = _STR_TO_DTYPE.get(
|
||||
self.config.output.dtype.get(key, "int32"), torch.int32
|
||||
)
|
||||
if ids_list and isinstance(ids_list[0], list):
|
||||
tensors[key] = [
|
||||
torch.tensor(
|
||||
list(chain.from_iterable(ids))
|
||||
if ids and isinstance(ids[0], list)
|
||||
else ids,
|
||||
dtype=dt,
|
||||
)
|
||||
for ids in ids_list
|
||||
]
|
||||
else:
|
||||
tensors[key] = [
|
||||
torch.tensor(list(chain.from_iterable(ids_list)), dtype=dt)
|
||||
]
|
||||
return tensors
|
||||
@@ -0,0 +1,46 @@
|
||||
"""Position-id generation strategies for packed sequences.
|
||||
|
||||
Each strategy takes the list of per-document token sequences after packing
|
||||
and returns a flat list of position ids (same total length as all
|
||||
sequences combined). The pipeline wraps the result into a tensor and
|
||||
attaches it as ``position_ids``.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
|
||||
class PositionIdStrategy(ABC):
|
||||
"""Generate ``position_ids`` for packed sequences."""
|
||||
|
||||
@abstractmethod
|
||||
def generate(self, sequences: List[List[int]]) -> List[int]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class PositionIdStrategyFactory(BaseFactory["PositionIdStrategy"]):
|
||||
pass
|
||||
|
||||
|
||||
@PositionIdStrategyFactory.register("none")
|
||||
class NoPositionId(PositionIdStrategy):
|
||||
def generate(self, sequences: List[List[int]]) -> List[int]:
|
||||
return []
|
||||
|
||||
|
||||
@PositionIdStrategyFactory.register("doc_reset")
|
||||
class DocResetPositionId(PositionIdStrategy):
|
||||
def generate(self, sequences: List[List[int]]) -> List[int]:
|
||||
pos_ids = []
|
||||
for seq in sequences:
|
||||
pos_ids.extend(range(len(seq)))
|
||||
return pos_ids
|
||||
|
||||
|
||||
@PositionIdStrategyFactory.register("continuous")
|
||||
class ContinuousPositionId(PositionIdStrategy):
|
||||
def generate(self, sequences: List[List[int]]) -> List[int]:
|
||||
total = sum(len(seq) for seq in sequences)
|
||||
return list(range(total))
|
||||
@@ -0,0 +1,92 @@
|
||||
"""Tokenization transform for JSONL record streams.
|
||||
|
||||
Bridges the Reader layer (``JsonlStore`` reads raw JSON records) and the
|
||||
Dataset layer (expects per-record tensors). Holds the tokenizer,
|
||||
mask-builder and position-id strategy together so that I/O code stays
|
||||
free of model dependencies.
|
||||
|
||||
The record-processing core (mask building, primary-id extraction,
|
||||
per-key tensorisation, position-id generation) is shared with
|
||||
:class:`astrai.preprocessing.pipeline.Pipeline` via the
|
||||
:mod:`astrai.preprocessing.core` helpers.
|
||||
"""
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Dict, List
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
from astrai.preprocessing.core import (
|
||||
build_position_ids,
|
||||
build_preprocessing_components,
|
||||
iter_raw_records,
|
||||
to_per_record_tensors,
|
||||
)
|
||||
|
||||
|
||||
class TokenizeTransform:
|
||||
"""Tokenize raw JSONL record dicts into per-key tensor lists.
|
||||
|
||||
Owns the three preprocessing concerns that were previously inlined in
|
||||
``JsonlStore``: tokenization, loss-mask construction and position-id
|
||||
generation. Constructing it loads the tokenizer, so it is intentionally
|
||||
cheap to pass around once built.
|
||||
|
||||
Args:
|
||||
config: Pipeline config describing sections / masks / position mode.
|
||||
tokenizer_path: Path passed to ``AutoTokenizer.from_pretrained``.
|
||||
"""
|
||||
|
||||
def __init__(self, config: PipelineConfig, tokenizer_path: str):
|
||||
self.config = config
|
||||
self.tokenizer, self.mask_builder, self.position_strategy = (
|
||||
build_preprocessing_components(config, tokenizer_path)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_config_file(cls, config_path: str) -> "TokenizeTransform":
|
||||
"""Build from a ``dataset_config.json`` file path.
|
||||
|
||||
The config file follows :class:`PipelineConfig` schema with an
|
||||
extra ``tokenizer_path`` field. When omitted, the config's
|
||||
parent directory is used as the tokenizer path.
|
||||
"""
|
||||
root = Path(config_path).parent
|
||||
with open(config_path, "r", encoding="utf-8") as f:
|
||||
raw_config = json.load(f)
|
||||
tokenizer_path = raw_config.pop("tokenizer_path", None) or str(root)
|
||||
config = PipelineConfig.from_dict(raw_config)
|
||||
return cls(config, tokenizer_path)
|
||||
|
||||
def apply(self, records: List[dict]) -> Dict[str, list]:
|
||||
"""Tokenize a list of raw record dicts.
|
||||
|
||||
Returns a dict mapping key (``sequence``, ``chosen``, ``responses``,
|
||||
…) to a list of per-record tensors (or nested tensor lists for
|
||||
multi-response keys such as GRPO ``responses``).
|
||||
"""
|
||||
raw: Dict[str, list] = {}
|
||||
doc_sequences: List[List[int]] = []
|
||||
|
||||
for result in iter_raw_records(
|
||||
records, self.mask_builder, self.config, self.tokenizer
|
||||
):
|
||||
primary = None
|
||||
for val in result.values():
|
||||
if isinstance(val, list) and val and isinstance(val[0], int):
|
||||
primary = val
|
||||
break
|
||||
if primary is not None:
|
||||
doc_sequences.append(primary)
|
||||
for key, ids in result.items():
|
||||
raw.setdefault(key, []).append(ids)
|
||||
|
||||
tensors = to_per_record_tensors(raw)
|
||||
|
||||
pos_ids = build_position_ids(doc_sequences, self.position_strategy)
|
||||
if pos_ids is not None:
|
||||
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
||||
|
||||
return tensors
|
||||
@@ -0,0 +1,56 @@
|
||||
"""Storage writer strategies for pipeline output.
|
||||
|
||||
The :class:`StoreWriter` abstraction decouples the pipeline from the
|
||||
concrete storage format (bin). The pipeline builds a ``{key:
|
||||
List[Tensor]}`` dict and delegates the write to the writer selected
|
||||
by ``output.storage_format``.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Dict, List
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.serialization import save_bin
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class StoreWriter(ABC):
|
||||
"""Write pre-tokenized tensors to disk in a format-specific way."""
|
||||
|
||||
@abstractmethod
|
||||
def save(
|
||||
self,
|
||||
output_dir: str,
|
||||
domain: str,
|
||||
shard_idx: int,
|
||||
tensors: Dict[str, List[torch.Tensor]],
|
||||
) -> None: ...
|
||||
|
||||
|
||||
class StoreWriterFactory(BaseFactory["StoreWriter"]):
|
||||
pass
|
||||
|
||||
|
||||
@StoreWriterFactory.register("bin")
|
||||
class BinWriter(StoreWriter):
|
||||
def save(self, output_dir, domain, shard_idx, tensors):
|
||||
shard_path = os.path.join(output_dir, domain, f"shard_{shard_idx:04d}")
|
||||
try:
|
||||
save_bin(shard_path, tensors)
|
||||
except Exception:
|
||||
if os.path.exists(shard_path):
|
||||
shutil.rmtree(shard_path, ignore_errors=True)
|
||||
logger.error(
|
||||
"Failed to write shard %s/%s_%04d, cleaned up partial output",
|
||||
domain,
|
||||
"shard",
|
||||
shard_idx,
|
||||
exc_info=True,
|
||||
)
|
||||
raise
|
||||
@@ -0,0 +1,21 @@
|
||||
"""Training component protocols — structural subtyping for optimizer/scheduler wrappers."""
|
||||
|
||||
from typing import Any, Protocol, runtime_checkable
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class OptimizerProtocol(Protocol):
|
||||
def step(self, closure=None): ...
|
||||
def zero_grad(self): ...
|
||||
@property
|
||||
def param_groups(self) -> Any: ...
|
||||
def state_dict(self) -> dict: ...
|
||||
def load_state_dict(self, d: dict): ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class SchedulerProtocol(Protocol):
|
||||
def step(self): ...
|
||||
def state_dict(self) -> dict: ...
|
||||
def load_state_dict(self, d: dict): ...
|
||||
def get_last_lr(self): ...
|
||||
@@ -1,116 +0,0 @@
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import h5py
|
||||
import safetensors.torch as st
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.parallel.setup import get_rank
|
||||
|
||||
|
||||
def save_h5(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
|
||||
os.makedirs(file_path, exist_ok=True)
|
||||
full_file_path = os.path.join(file_path, f"{file_name}.h5")
|
||||
with h5py.File(full_file_path, "w") as f:
|
||||
for key, tensors in tensor_group.items():
|
||||
grp = f.create_group(key)
|
||||
for idx, tensor in enumerate(tensors):
|
||||
arr = tensor.cpu().numpy()
|
||||
grp.create_dataset(f"data_{idx}", data=arr)
|
||||
|
||||
|
||||
def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
|
||||
tensor_group: Dict[str, List[Tensor]] = {}
|
||||
|
||||
root_path = Path(file_path)
|
||||
h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5"))
|
||||
|
||||
for h5_file in h5_files:
|
||||
with h5py.File(h5_file, "r") as f:
|
||||
for key in f.keys():
|
||||
grp = f[key]
|
||||
dsets = []
|
||||
for dset_name in grp.keys():
|
||||
dset = grp[dset_name]
|
||||
tensor = torch.from_numpy(dset[:])
|
||||
if share_memory:
|
||||
tensor = tensor.share_memory_()
|
||||
dsets.append(tensor)
|
||||
|
||||
if tensor_group.get(key) is None:
|
||||
tensor_group[key] = []
|
||||
tensor_group[key].extend(dsets)
|
||||
|
||||
return tensor_group
|
||||
|
||||
|
||||
class Checkpoint:
|
||||
def __init__(
|
||||
self,
|
||||
state_dict: Dict[str, Any],
|
||||
epoch: int = 0,
|
||||
iteration: int = 0,
|
||||
extra: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
self.state_dict = state_dict
|
||||
self.epoch = epoch
|
||||
self.iteration = iteration
|
||||
self.extra = extra or {}
|
||||
|
||||
def save(
|
||||
self,
|
||||
save_dir: str,
|
||||
) -> None:
|
||||
|
||||
save_path = Path(save_dir)
|
||||
save_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
rank = get_rank()
|
||||
if rank == 0:
|
||||
meta = {
|
||||
"epoch": self.epoch,
|
||||
"iteration": self.iteration,
|
||||
}
|
||||
with open(save_path / "meta.json", "w") as f:
|
||||
json.dump(meta, f, indent=2)
|
||||
|
||||
st.save_file(self.state_dict, save_path / "state_dict.safetensors")
|
||||
if self.extra:
|
||||
torch.save(self.extra, save_path / "extra.pt")
|
||||
|
||||
@classmethod
|
||||
def load(
|
||||
cls,
|
||||
save_dir: str,
|
||||
) -> "Checkpoint":
|
||||
|
||||
rank = get_rank()
|
||||
save_path = Path(save_dir)
|
||||
|
||||
meta = {}
|
||||
if rank == 0:
|
||||
with open(Path(save_dir) / "meta.json", "r") as f:
|
||||
meta = json.load(f)
|
||||
|
||||
if dist.is_initialized():
|
||||
meta_list = [meta]
|
||||
dist.broadcast_object_list(meta_list, src=0)
|
||||
meta = meta_list[0]
|
||||
|
||||
state_dict = st.load_file(save_path / "state_dict.safetensors")
|
||||
|
||||
extra = None
|
||||
extra_path = save_path / "extra.pt"
|
||||
if extra_path.exists():
|
||||
extra = torch.load(extra_path, map_location="cpu", weights_only=False)
|
||||
|
||||
return cls(
|
||||
state_dict=state_dict,
|
||||
epoch=meta["epoch"],
|
||||
iteration=meta["iteration"],
|
||||
extra=extra,
|
||||
)
|
||||
@@ -0,0 +1,41 @@
|
||||
"""Serialization utilities for models and datasets.
|
||||
|
||||
This package re-exports checkpoint helpers and dataset storage helpers so
|
||||
that existing imports from ``astrai.serialization`` continue to work.
|
||||
"""
|
||||
|
||||
from astrai.serialization.checkpoint import (
|
||||
Checkpoint,
|
||||
load_json,
|
||||
load_model_config,
|
||||
load_model_weights,
|
||||
load_safetensors,
|
||||
load_state_dict,
|
||||
load_torch,
|
||||
save_json,
|
||||
save_model,
|
||||
save_safetensors,
|
||||
save_torch,
|
||||
)
|
||||
from astrai.serialization.dataset import (
|
||||
load_bin,
|
||||
load_bin_offsets,
|
||||
save_bin,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Checkpoint",
|
||||
"load_json",
|
||||
"load_model_config",
|
||||
"load_model_weights",
|
||||
"load_safetensors",
|
||||
"load_state_dict",
|
||||
"load_torch",
|
||||
"save_json",
|
||||
"save_model",
|
||||
"save_safetensors",
|
||||
"save_torch",
|
||||
"load_bin",
|
||||
"load_bin_offsets",
|
||||
"save_bin",
|
||||
]
|
||||
@@ -0,0 +1,201 @@
|
||||
"""Model checkpoint serialization helpers."""
|
||||
|
||||
import io
|
||||
import json
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Optional, Union
|
||||
|
||||
import safetensors.torch as st
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from astrai.parallel.setup import get_rank
|
||||
|
||||
_META_FILE = "meta.json"
|
||||
_CONFIG_FILE = "config.json"
|
||||
_WEIGHTS_FILE = "model.safetensors"
|
||||
|
||||
|
||||
def save_safetensors(state_dict: dict, path: Union[str, Path]):
|
||||
st.save_file(state_dict, str(path))
|
||||
|
||||
|
||||
def load_safetensors(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||
if not broadcast or not dist.is_initialized():
|
||||
return st.load_file(str(path))
|
||||
|
||||
rank = get_rank()
|
||||
if rank == 0:
|
||||
state_dict = st.load_file(str(path))
|
||||
else:
|
||||
state_dict = {}
|
||||
tmp = [state_dict]
|
||||
dist.broadcast_object_list(tmp, src=0)
|
||||
return tmp[0]
|
||||
|
||||
|
||||
def save_json(data: dict, path: Union[str, Path]):
|
||||
with open(str(path), "w") as f:
|
||||
json.dump(data, f, indent=2)
|
||||
|
||||
|
||||
def load_json(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||
if not broadcast or not dist.is_initialized():
|
||||
with open(str(path), "r") as f:
|
||||
return json.load(f)
|
||||
|
||||
rank = get_rank()
|
||||
if rank == 0:
|
||||
with open(str(path), "r") as f:
|
||||
data = json.load(f)
|
||||
else:
|
||||
data = {}
|
||||
tmp = [data]
|
||||
dist.broadcast_object_list(tmp, src=0)
|
||||
return tmp[0]
|
||||
|
||||
|
||||
def save_torch(obj: Any, path: Union[str, Path]):
|
||||
torch.save(obj, str(path))
|
||||
|
||||
|
||||
def load_torch(path: Union[str, Path], broadcast: bool = False) -> Any:
|
||||
if not broadcast or not dist.is_initialized():
|
||||
return torch.load(str(path), map_location="cpu", weights_only=False)
|
||||
|
||||
path = Path(path)
|
||||
rank = get_rank()
|
||||
|
||||
if rank == 0:
|
||||
with open(path, "rb") as f:
|
||||
raw = f.read()
|
||||
data_tensor = torch.frombuffer(bytearray(raw), dtype=torch.uint8)
|
||||
num_bytes = torch.tensor([len(raw)], dtype=torch.long)
|
||||
else:
|
||||
num_bytes = torch.tensor([0], dtype=torch.long)
|
||||
|
||||
dist.broadcast(num_bytes, src=0)
|
||||
|
||||
if rank != 0:
|
||||
data_tensor = torch.empty(num_bytes.item(), dtype=torch.uint8)
|
||||
|
||||
dist.broadcast(data_tensor, src=0)
|
||||
|
||||
buf = io.BytesIO(data_tensor.numpy().tobytes())
|
||||
return torch.load(buf, map_location="cpu", weights_only=False)
|
||||
|
||||
|
||||
def save_model(config: dict, state_dict: dict, save_directory: str):
|
||||
save_path = Path(save_directory)
|
||||
save_path.mkdir(parents=True, exist_ok=True)
|
||||
save_json(config, save_path / _CONFIG_FILE)
|
||||
save_safetensors(state_dict, save_path / _WEIGHTS_FILE)
|
||||
|
||||
|
||||
def load_model_config(save_directory: str) -> dict:
|
||||
return load_json(Path(save_directory) / _CONFIG_FILE)
|
||||
|
||||
|
||||
def load_model_weights(save_directory: str) -> dict:
|
||||
return load_state_dict(Path(save_directory) / _WEIGHTS_FILE)
|
||||
|
||||
|
||||
def load_state_dict(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||
path = Path(path)
|
||||
if not broadcast or not dist.is_initialized():
|
||||
return load_safetensors(path)
|
||||
|
||||
rank = get_rank()
|
||||
if rank == 0:
|
||||
state_dict = load_safetensors(path)
|
||||
specs = [
|
||||
(k, list(state_dict[k].shape), str(state_dict[k].dtype).split(".")[-1])
|
||||
for k in sorted(state_dict)
|
||||
]
|
||||
else:
|
||||
state_dict = {}
|
||||
specs = []
|
||||
|
||||
specs_list = [specs]
|
||||
dist.broadcast_object_list(specs_list, src=0)
|
||||
specs = specs_list[0]
|
||||
|
||||
for key, shape, dtype_name in specs:
|
||||
dtype = getattr(torch, dtype_name)
|
||||
if rank != 0:
|
||||
tensor = torch.empty(shape, dtype=dtype, device="cpu")
|
||||
else:
|
||||
tensor = state_dict[key].contiguous().cpu()
|
||||
dist.broadcast(tensor, src=0)
|
||||
if rank != 0:
|
||||
state_dict[key] = tensor
|
||||
return state_dict
|
||||
|
||||
|
||||
@dataclass
|
||||
class Checkpoint:
|
||||
state_dict: Dict[str, Any] = field(default_factory=dict)
|
||||
epoch: int = 0
|
||||
consumed_samples: int = 0
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
meta: Dict[str, Any] = field(default_factory=dict)
|
||||
config: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def save(self, save_dir: str):
|
||||
save_path = Path(save_dir)
|
||||
save_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
meta = {
|
||||
"epoch": self.epoch,
|
||||
"consumed_samples": self.consumed_samples,
|
||||
"timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"),
|
||||
**self.meta,
|
||||
}
|
||||
save_json(meta, save_path / _META_FILE)
|
||||
save_json(self.config, save_path / _CONFIG_FILE)
|
||||
save_safetensors(self.state_dict, save_path / _WEIGHTS_FILE)
|
||||
for key, value in self.extra.items():
|
||||
save_torch(value, save_path / f"{key}.pt")
|
||||
|
||||
@classmethod
|
||||
def load(cls, save_dir: str, broadcast: bool = False) -> "Checkpoint":
|
||||
save_path = Path(save_dir)
|
||||
|
||||
meta = load_json(save_path / _META_FILE, broadcast)
|
||||
config = load_json(save_path / _CONFIG_FILE, broadcast)
|
||||
state_dict = load_state_dict(save_path / _WEIGHTS_FILE, broadcast=broadcast)
|
||||
|
||||
extra = {}
|
||||
for f in sorted(save_path.iterdir()):
|
||||
if f.suffix == ".pt":
|
||||
extra[f.stem] = load_torch(f, broadcast=broadcast)
|
||||
|
||||
return cls(
|
||||
state_dict=state_dict,
|
||||
epoch=meta.get("epoch", 0),
|
||||
consumed_samples=meta.get("consumed_samples", 0),
|
||||
extra=extra,
|
||||
meta=meta,
|
||||
config=config,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def load_any(cls, save_dir: str, broadcast: bool = False) -> Optional["Checkpoint"]:
|
||||
save_path = Path(save_dir)
|
||||
meta_path = save_path / _META_FILE
|
||||
weights_path = save_path / _WEIGHTS_FILE
|
||||
|
||||
if meta_path.exists():
|
||||
return cls.load(save_dir, broadcast=broadcast)
|
||||
|
||||
if weights_path.exists():
|
||||
state_dict = load_state_dict(weights_path, broadcast=broadcast)
|
||||
config = {}
|
||||
config_path = save_path / _CONFIG_FILE
|
||||
if config_path.exists():
|
||||
config = load_json(config_path, broadcast)
|
||||
return cls(state_dict=state_dict, config=config)
|
||||
|
||||
return None
|
||||
@@ -0,0 +1,82 @@
|
||||
"""Dataset storage serialization helpers (memory-mapped binary)."""
|
||||
|
||||
import json
|
||||
import os
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
def save_bin(
|
||||
file_path: str,
|
||||
tensor_group: Dict[str, List[Tensor]],
|
||||
record_keys: Optional[List[str]] = None,
|
||||
):
|
||||
"""Save tensors as memory-mapped binary files.
|
||||
|
||||
When *record_keys* is provided, those keys are written with per-record
|
||||
cumulative offsets in ``meta.json`` so that ``MmapStore.fetch_record``
|
||||
can slice individual records from the concatenated binary without
|
||||
cross-record concatenation. Keys not in *record_keys* (e.g. SEQ
|
||||
``sequence``) are written as a single contiguous stream without
|
||||
offsets, preserving backward compatibility.
|
||||
|
||||
Nested keys (``List[List[Tensor]]`` such as GRPO ``responses``) are
|
||||
not supported in bin format — use JSONL for those.
|
||||
"""
|
||||
os.makedirs(file_path, exist_ok=True)
|
||||
record_keys = set(record_keys or [])
|
||||
meta = {}
|
||||
for key, tensors in tensor_group.items():
|
||||
if tensors and isinstance(tensors[0], list):
|
||||
raise ValueError(
|
||||
f"Nested key '{key}' (List[List[Tensor]]) is not supported "
|
||||
f"in bin format. Use JSONL storage instead."
|
||||
)
|
||||
cat = torch.cat(tensors, dim=0)
|
||||
entry: Dict[str, Any] = {
|
||||
"shape": list(cat.shape),
|
||||
"dtype": str(cat.dtype).split(".")[-1],
|
||||
}
|
||||
if key in record_keys:
|
||||
offsets = [0]
|
||||
for t in tensors:
|
||||
offsets.append(offsets[-1] + t.shape[0])
|
||||
entry["offsets"] = offsets
|
||||
meta[key] = entry
|
||||
np.asarray(cat.cpu().numpy()).tofile(os.path.join(file_path, f"{key}.bin"))
|
||||
with open(os.path.join(file_path, "meta.json"), "w") as f:
|
||||
json.dump(meta, f)
|
||||
|
||||
|
||||
def load_bin(file_path: str) -> Dict[str, List[Tensor]]:
|
||||
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
||||
meta = json.load(f)
|
||||
segments: Dict[str, List[Tensor]] = {}
|
||||
for key, info in meta.items():
|
||||
arr = np.memmap(
|
||||
os.path.join(file_path, f"{key}.bin"),
|
||||
dtype=info["dtype"],
|
||||
mode="c",
|
||||
shape=tuple(info["shape"]),
|
||||
)
|
||||
segments[key] = [torch.from_numpy(arr)]
|
||||
return segments
|
||||
|
||||
|
||||
def load_bin_offsets(file_path: str) -> Dict[str, List[int]]:
|
||||
"""Read per-record cumulative offsets from ``meta.json``.
|
||||
|
||||
Returns an empty dict when no key has offsets (legacy bin files),
|
||||
in which case record-mode access falls back to per-record segment
|
||||
indexing (JSONL layout).
|
||||
"""
|
||||
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
||||
meta = json.load(f)
|
||||
offsets: Dict[str, List[int]] = {}
|
||||
for key, info in meta.items():
|
||||
if "offsets" in info:
|
||||
offsets[key] = info["offsets"]
|
||||
return offsets
|
||||
@@ -0,0 +1,53 @@
|
||||
import logging
|
||||
import os
|
||||
import signal
|
||||
import threading
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_early_stop = threading.Event()
|
||||
_active_context = None
|
||||
|
||||
|
||||
def _early_handler(signum: int, frame):
|
||||
sig = signal.Signals(signum)
|
||||
logger.warning(
|
||||
"Received %s (pid=%d), requesting graceful training stop...",
|
||||
sig.name,
|
||||
os.getpid(),
|
||||
)
|
||||
_early_stop.set()
|
||||
if _active_context is not None:
|
||||
_active_context.request_stop()
|
||||
|
||||
|
||||
def install_early_signal_handlers():
|
||||
for sig in (signal.SIGTERM, signal.SIGINT):
|
||||
signal.signal(sig, _early_handler)
|
||||
_unblock_signals()
|
||||
|
||||
|
||||
def _unblock_signals():
|
||||
try:
|
||||
mask = signal.pthread_sigmask(signal.SIG_BLOCK, set())
|
||||
blocked = {signal.SIGTERM, signal.SIGINT} & mask
|
||||
if blocked:
|
||||
signal.pthread_sigmask(signal.SIG_UNBLOCK, blocked)
|
||||
except (AttributeError, OSError):
|
||||
pass
|
||||
|
||||
|
||||
def register_signal_handlers(context):
|
||||
global _active_context
|
||||
_active_context = context
|
||||
for sig in (signal.SIGTERM, signal.SIGINT):
|
||||
signal.signal(sig, _early_handler)
|
||||
if _early_stop.is_set():
|
||||
context.request_stop()
|
||||
logger.warning("Signal was received during initialization, stopping...")
|
||||
|
||||
|
||||
def unregister_signal_handlers():
|
||||
global _active_context
|
||||
_active_context = None
|
||||
_early_stop.clear()
|
||||
@@ -1,8 +1,10 @@
|
||||
from astrai.tokenize.chat_template import ChatTemplate, MessageType
|
||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||
from astrai.tokenize.tokenizer import AutoTokenizer, Message, Messages
|
||||
|
||||
__all__ = [
|
||||
"AutoTokenizer",
|
||||
"ChatTemplate",
|
||||
"MessageType",
|
||||
"Message",
|
||||
"Messages",
|
||||
]
|
||||
|
||||
@@ -1,13 +1,11 @@
|
||||
from dataclasses import dataclass
|
||||
from functools import cached_property
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from jinja2 import Template
|
||||
|
||||
# Message type for chat messages
|
||||
type MessageType = Dict[str, Any]
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChatTemplate:
|
||||
"""A chat template with Jinja2 rendering support.
|
||||
|
||||
@@ -15,23 +13,51 @@ class ChatTemplate:
|
||||
name: Unique identifier for the template.
|
||||
template_str: Jinja2 template string.
|
||||
description: Optional description.
|
||||
default_variables: Optional dictionary of default variable values
|
||||
that will be passed to the template if not overridden during rendering.
|
||||
default_variables: Optional dictionary of default variable values.
|
||||
special_tokens: Optional dictionary mapping token names to their string values.
|
||||
These tokens are automatically added to the template variables.
|
||||
"""
|
||||
|
||||
name: str
|
||||
template_str: str
|
||||
description: str = ""
|
||||
default_variables: Dict[str, Any] = None
|
||||
special_tokens: Dict[str, str] = None
|
||||
def __init__(
|
||||
self,
|
||||
name: str = "",
|
||||
template_str: str = "",
|
||||
description: str = "",
|
||||
default_variables: Optional[Dict[str, Any]] = None,
|
||||
special_tokens: Optional[Dict[str, str]] = None,
|
||||
):
|
||||
self.name = name
|
||||
self.template_str = template_str
|
||||
self.description = description
|
||||
self.default_variables = default_variables or {}
|
||||
self.special_tokens = special_tokens or {}
|
||||
|
||||
def __post_init__(self):
|
||||
if self.default_variables is None:
|
||||
self.default_variables = {}
|
||||
if self.special_tokens is None:
|
||||
self.special_tokens = {}
|
||||
@cached_property
|
||||
def _compiled(self) -> Template:
|
||||
"""Lazy-compiled Jinja2 template, cached on first access.
|
||||
|
||||
The compiled :class:`~jinja2.Template` holds a dynamically-generated
|
||||
``root`` render function whose ``__module__`` is ``None``; under
|
||||
``pickle`` it falls back to ``__main__`` and breaks ``spawn``-based
|
||||
multiprocessing. :meth:`__getstate__` drops the cached template so
|
||||
that pickle serialises only ``template_str``; each worker rebuilds
|
||||
the cache on first render.
|
||||
"""
|
||||
return Template(self.template_str)
|
||||
|
||||
def __getstate__(self) -> Dict[str, Any]:
|
||||
"""Exclude the cached Jinja2 template from pickling.
|
||||
|
||||
``Template.root_render_func`` is a dynamically generated closure
|
||||
that cannot be pickled by reference. Dropping ``_compiled`` here
|
||||
lets :class:`cached_property` rebuild it on first access after
|
||||
unpickle.
|
||||
"""
|
||||
state = self.__dict__.copy()
|
||||
state.pop("_compiled", None)
|
||||
return state
|
||||
|
||||
def __setstate__(self, state: Dict[str, Any]) -> None:
|
||||
self.__dict__.update(state)
|
||||
|
||||
@classmethod
|
||||
def from_string(
|
||||
@@ -43,7 +69,7 @@ class ChatTemplate:
|
||||
) -> "ChatTemplate":
|
||||
"""Create a ChatTemplate instance directly from a template string."""
|
||||
return cls(
|
||||
name="", # empty name for ad‑hoc templates
|
||||
name="",
|
||||
template_str=template_str,
|
||||
description=description,
|
||||
default_variables=default_variables,
|
||||
@@ -73,5 +99,4 @@ class ChatTemplate:
|
||||
if system_prompt is not None:
|
||||
variables["system_prompt"] = system_prompt
|
||||
|
||||
jinja_template = Template(self.template_str)
|
||||
return jinja_template.render(**variables)
|
||||
return self._compiled.render(**variables)
|
||||
|
||||
@@ -10,12 +10,16 @@ from tokenizers import Tokenizer
|
||||
|
||||
from astrai.tokenize.chat_template import ChatTemplate
|
||||
|
||||
Message = Dict[str, str]
|
||||
"""Single chat message with ``role`` and ``content`` keys."""
|
||||
|
||||
Messages = List[Message]
|
||||
"""Single conversation — a list of messages."""
|
||||
|
||||
|
||||
class AutoTokenizer:
|
||||
"""Base tokenizer class with automatic loading support"""
|
||||
|
||||
TOKENIZER_CLASSES = {} # Registry for auto-loading
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path: Optional[Union[str, Path]] = None,
|
||||
@@ -51,9 +55,26 @@ class AutoTokenizer:
|
||||
self.set_chat_template(config["chat_template"])
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, path: Union[str, Path], **kwargs) -> "AutoTokenizer":
|
||||
"""Load tokenizer from pretrained directory."""
|
||||
def from_pretrained(cls, path: Union[str, Path]) -> "AutoTokenizer":
|
||||
"""Load tokenizer from pretrained directory.
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If tokenizer.json is missing.
|
||||
RuntimeError: If tokenizer failed to initialize.
|
||||
"""
|
||||
path = Path(path)
|
||||
tokenizer_file = path / "tokenizer.json"
|
||||
if not tokenizer_file.exists():
|
||||
raise FileNotFoundError(
|
||||
f"Tokenizer file not found: {tokenizer_file}. "
|
||||
"A valid tokenizer.json is required."
|
||||
)
|
||||
instance = cls(path)
|
||||
if instance._tokenizer is None:
|
||||
raise RuntimeError(
|
||||
f"Failed to load tokenizer from {path}. "
|
||||
"The tokenizer.json may be corrupted or incompatible."
|
||||
)
|
||||
return instance
|
||||
|
||||
def save_pretrained(self, save_path: str):
|
||||
@@ -85,17 +106,6 @@ class AutoTokenizer:
|
||||
with open(save_path / "tokenizer_config.json", "w", encoding="utf-8") as f:
|
||||
json.dump(config, f, ensure_ascii=False, indent=2)
|
||||
|
||||
@classmethod
|
||||
def register_tokenizer(cls, name: str, tokenizer_class: type):
|
||||
"""
|
||||
Register a new tokenizer class.
|
||||
|
||||
Args:
|
||||
name: Name to register the tokenizer class under
|
||||
tokenizer_class: The tokenizer class to register
|
||||
"""
|
||||
cls.TOKENIZER_CLASSES[name] = tokenizer_class
|
||||
|
||||
def encode(
|
||||
self,
|
||||
tokens: Union[str, List[str]],
|
||||
@@ -103,7 +113,16 @@ class AutoTokenizer:
|
||||
is_pretokenized: bool = False,
|
||||
add_special_tokens: bool = True,
|
||||
) -> List:
|
||||
"""Encode text to tokens or token IDs."""
|
||||
"""Encode text to token IDs.
|
||||
|
||||
Accepts both single strings and batches:
|
||||
|
||||
- ``encode("hello")`` → ``[123, 456]``
|
||||
- ``encode(["hello", "world"])`` → ``[[123, 456], [789]]``
|
||||
|
||||
Batches are tokenised in parallel via the Rust backend's
|
||||
``encode_batch`` (uses all available CPU cores).
|
||||
"""
|
||||
if self._tokenizer is None:
|
||||
raise RuntimeError(
|
||||
"Tokenizer not initialized. Load or create a tokenizer first."
|
||||
@@ -116,15 +135,13 @@ class AutoTokenizer:
|
||||
add_special_tokens=add_special_tokens,
|
||||
)
|
||||
return encoded.ids if out_ids else encoded.tokens
|
||||
else:
|
||||
encoded_list = self._tokenizer.encode_batch(
|
||||
tokens,
|
||||
is_pretokenized=is_pretokenized,
|
||||
add_special_tokens=add_special_tokens,
|
||||
)
|
||||
return [
|
||||
encoded.ids if out_ids else encoded.tokens for encoded in encoded_list
|
||||
]
|
||||
|
||||
encoded_list = self._tokenizer.encode_batch(
|
||||
tokens,
|
||||
is_pretokenized=is_pretokenized,
|
||||
add_special_tokens=add_special_tokens,
|
||||
)
|
||||
return [encoded.ids if out_ids else encoded.tokens for encoded in encoded_list]
|
||||
|
||||
def decode(self, tokens: List[int], skip_special_tokens: bool = True) -> str:
|
||||
"""Decode token IDs to text."""
|
||||
@@ -147,7 +164,14 @@ class AutoTokenizer:
|
||||
- tokenizer.bos_token → returns string
|
||||
- tokenizer.bos_token_id → returns corresponding integer ID
|
||||
- tokenizer.stop_ids → returns list of corresponding integer IDs for all special tokens
|
||||
|
||||
Internal/private attrs are not intercepted: during unpickle
|
||||
``__dict__`` is empty, so probing ``self._special_token_map``
|
||||
would recurse infinitely.
|
||||
"""
|
||||
if key.startswith("_"):
|
||||
raise AttributeError(key)
|
||||
|
||||
# Handle stop_ids - return IDs for all special tokens
|
||||
if key == "stop_ids":
|
||||
stop_ids = []
|
||||
@@ -203,45 +227,63 @@ class AutoTokenizer:
|
||||
|
||||
def apply_chat_template(
|
||||
self,
|
||||
messages: List[Dict[str, str]],
|
||||
messages: Union[Messages, List[Messages]],
|
||||
system_prompt: Optional[str] = None,
|
||||
tokenize: bool = True,
|
||||
add_generation_prompt: bool = True,
|
||||
**kwargs,
|
||||
) -> Union[str, List[int]]:
|
||||
"""
|
||||
Apply the chat template to messages and optionally tokenize the result.
|
||||
) -> Union[str, List[int], List[str], List[List[int]]]:
|
||||
"""Apply the chat template and optionally tokenize.
|
||||
|
||||
Accepts both single conversations and batches:
|
||||
|
||||
- ``apply_chat_template([msg1, msg2])`` → ``"..."`` or ``[ids]``
|
||||
- ``apply_chat_template([[msg1, msg2], [msg3]])`` → ``["..", ".."]``
|
||||
or ``[[ids], [ids]]``
|
||||
|
||||
Batches render each conversation list and tokenise all at once via
|
||||
:meth:`encode` (``List[str]`` → Rust parallel ``encode_batch``).
|
||||
|
||||
Args:
|
||||
messages: List of message dicts with 'role' and 'content'.
|
||||
system_prompt: Optional system prompt string (auto-converted to first message).
|
||||
messages: Single conversation (``Messages``) or batch of
|
||||
conversations (``BatchMessages``).
|
||||
system_prompt: Optional system prompt prepended (single mode only).
|
||||
tokenize: Whether to return token IDs (True) or raw string (False).
|
||||
add_generation_prompt: Whether to add the generation prompt (default: True).
|
||||
**kwargs: Additional variables to pass to the template.
|
||||
add_generation_prompt: Whether to add the generation prompt.
|
||||
**kwargs: Additional template variables.
|
||||
|
||||
Returns:
|
||||
Either the rendered string or list of token IDs.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If chat template is not set.
|
||||
Single mode: ``str`` or ``List[int]``.
|
||||
Batch mode: ``List[str]`` or ``List[List[int]]``.
|
||||
"""
|
||||
if self._chat_template is None:
|
||||
raise RuntimeError(
|
||||
"Chat template not set. Use set_chat_template() to set a template first."
|
||||
)
|
||||
|
||||
# Auto-convert system_prompt to first message if provided
|
||||
is_batch = bool(messages) and isinstance(messages[0], list)
|
||||
|
||||
if is_batch:
|
||||
rendered = [
|
||||
self._chat_template.render(
|
||||
messages=msgs,
|
||||
add_generation_prompt=add_generation_prompt,
|
||||
**kwargs,
|
||||
)
|
||||
for msgs in messages
|
||||
]
|
||||
if tokenize:
|
||||
return self.encode(rendered) # List[str] → batch encode
|
||||
return rendered
|
||||
|
||||
# Single conversation
|
||||
if system_prompt:
|
||||
messages = [{"role": "system", "content": system_prompt}] + list(messages)
|
||||
|
||||
# Render the template
|
||||
rendered = self._chat_template.render(
|
||||
messages=messages,
|
||||
add_generation_prompt=add_generation_prompt,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
if tokenize:
|
||||
return self.encode(rendered)
|
||||
|
||||
return rendered
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user