Compare commits
604
Commits
v1.3.2
..
2eeac02d70
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2eeac02d70 | ||
|
|
5e76fbd1bf | ||
|
|
4dc5e923e0 | ||
|
|
998b443aa3 | ||
|
|
cebdd45d3a | ||
|
|
7da1439c9e | ||
|
|
29e5f571af | ||
|
|
74e694921c | ||
|
|
d5067af064 | ||
|
|
f6db546578 | ||
|
|
31ca357c61 | ||
|
|
34471252ab | ||
|
|
aa08479285 | ||
|
|
4b10d3ca37 | ||
|
|
2bc4d2b8a8 | ||
|
|
4244df2785 | ||
|
|
a29bdfae46 | ||
|
|
10fec8dca1 | ||
|
|
75304d084d | ||
|
|
16a55bb474 | ||
|
|
cb21af38ba | ||
|
|
dcc96de12a | ||
|
|
7d27f3e078 | ||
|
|
84753d3e08 | ||
|
|
53a7149577 | ||
|
|
c79d34eee1 | ||
|
|
398e8a3ea3 | ||
|
|
f252af495c | ||
|
|
00c2c80c8f | ||
|
|
a6c6a54ace | ||
|
|
3d3ea47d37 | ||
|
|
f7f14d0e5f | ||
|
|
7580d80d45 | ||
|
|
cb51a3587b | ||
|
|
1bcd8f53ab | ||
|
|
0d0dc64884 | ||
|
|
6ac3b51496 | ||
|
|
3406157431 | ||
|
|
0dd9a417b7 | ||
|
|
a01c1fd427 | ||
|
|
f8d9ab344d | ||
|
|
3fb4b8ab13 | ||
|
|
b5afe3d7a4 | ||
|
|
69f35c46e0 | ||
|
|
0378e62e17 | ||
|
|
5244f1a8fc | ||
|
|
5104638447 | ||
|
|
a711d9f478 | ||
|
|
15862d4b56 | ||
|
|
f9efb705b8 | ||
|
|
c6a82a5029 | ||
|
|
a5b238dd86 | ||
|
|
da6d94492d | ||
|
|
71b6e3aaaf | ||
|
|
f95722a277 | ||
|
|
9f48cb8928 | ||
|
|
9b58fef222 | ||
|
|
c5fba9c238 | ||
|
|
cd31f1f62f | ||
|
|
a5a3cc1fc2 | ||
|
|
d565d44c43 | ||
|
|
596c35fd71 | ||
|
|
47b3ed4e44 | ||
|
|
c1d05ae11d | ||
|
|
cf4f5ab9f6 | ||
|
|
3416f98c58 | ||
|
|
d28552f878 | ||
|
|
be90dfe2bd | ||
|
|
a33ca04f60 | ||
|
|
7f0e8bb8c2 | ||
|
|
0c1b7664c1 | ||
|
|
3fa7e66676 | ||
|
|
ca50fe4721 | ||
|
|
d9240ab149 | ||
|
|
d7cd69fef5 | ||
|
|
9bff61fb91 | ||
|
|
0b661bae85 | ||
|
|
ae9fd546ef | ||
|
|
e3ea850dc9 | ||
|
|
6e5088cc7d | ||
|
|
cbc584470d | ||
|
|
cb60713a72 | ||
|
|
c52a2487ae | ||
|
|
49aaa9a714 | ||
|
|
056c1382ff | ||
|
|
f163520fff | ||
|
|
1b1f1a0707 | ||
|
|
184fbbce5c | ||
|
|
02469887f5 | ||
|
|
05739629fc | ||
|
|
e0f7fa8e13 | ||
|
|
af25833fab | ||
|
|
6572be4f98 | ||
|
|
81788faef4 | ||
|
|
0e7fe57d96 | ||
|
|
55ee258e95 | ||
|
|
ef1bb6f401 | ||
|
|
6f49738991 | ||
|
|
a59ae8f32e | ||
|
|
6054b8dbd4 | ||
|
|
6f67ba8942 | ||
|
|
d0c5debbab | ||
|
|
4f2e03880b | ||
|
|
5c180cfa90 | ||
|
|
6f09b1d2ee | ||
|
|
b2230fefd8 | ||
|
|
654e6eb0d1 | ||
|
|
a317a4756b | ||
|
|
9b7e6c205f | ||
|
|
602b5ce216 | ||
|
|
8152760b5f | ||
|
|
8c052c99ee | ||
|
|
2667b8116d | ||
|
|
6dffb0305a | ||
|
|
49a9c6b3d2 | ||
|
|
cdf9145ecf | ||
|
|
85f0461b3b | ||
|
|
9f0e9195f7 | ||
|
|
88751d0b08 | ||
|
|
d0e5d910de | ||
|
|
a03504a280 | ||
|
|
d033b2ef0f | ||
|
|
8447f88f61 | ||
|
|
b1b65a657e | ||
|
|
3439e3104e | ||
|
|
288ba20db1 | ||
|
|
020e2eff4e | ||
|
|
1c7369f293 | ||
|
|
0fc1b1bd46 | ||
|
|
d7db37a70f | ||
|
|
6d98bb4f9f | ||
|
|
925cbedc93 | ||
|
|
fda82ee232 | ||
|
|
4b25664c79 | ||
|
|
a27c8a819d | ||
|
|
91acaf4b0b | ||
|
|
41dcf0feb9 | ||
|
|
9960f79920 | ||
|
|
7feeb0b93e | ||
|
|
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 | ||
|
|
523eacf5fe | ||
|
|
cffedaad5e | ||
|
|
3583c46b66 | ||
|
|
ca4e6b907c | ||
|
|
db99d8b254 | ||
|
|
b98c9cefdc | ||
|
|
283bcaf2ff | ||
|
|
bc7c82977e | ||
|
|
34a511e36e | ||
|
|
d73f52a2f8 | ||
|
|
9d96b0431d | ||
|
|
f81e2b4a73 | ||
|
|
4e324d8f26 | ||
|
|
6ed0506491 | ||
|
|
30cc2d67a4 | ||
|
|
7ddebf2cd9 | ||
|
|
78dc2bd41c | ||
|
|
44d7a4e959 | ||
|
|
c4401512f2 | ||
|
|
a6f5ff3b37 | ||
|
|
ffff05b2c6 | ||
|
|
b89f8436ea | ||
|
|
123f25e339 | ||
|
|
520de3ebe8 | ||
|
|
466c34d7a8 | ||
|
|
6831a15424 | ||
|
|
0f9e5c5049 | ||
|
|
cb0e7f2a80 | ||
|
|
296db909aa | ||
|
|
a2ae742988 | ||
|
|
29beb174a5 | ||
|
|
bbeaff4c60 | ||
|
|
ab5e207f42 | ||
|
|
b0eff02446 | ||
|
|
408f0cb513 | ||
|
|
64b78ecce3 | ||
|
|
f2ffdf60d0 | ||
|
|
ace8f6ee68 | ||
|
|
a57a16430d | ||
|
|
3fee87897d | ||
|
|
3f67e53088 | ||
|
|
bf7adb35b3 | ||
|
|
feaa3fca36 | ||
|
|
39766aa1dc | ||
|
|
9b22b1651e | ||
|
|
e58dbd7c57 | ||
|
|
d2fe8afbd1 | ||
|
|
23ce4bc3ae | ||
|
|
d2b36cc85d | ||
|
|
fc278d17ab | ||
|
|
ff43a2fab8 | ||
|
|
2b26f03bd3 | ||
|
|
861d33b1a1 | ||
|
|
99b821ebf5 | ||
|
|
c94a246c71 | ||
|
|
2dc9545d7f | ||
|
|
9c31d78a22 | ||
|
|
bd9741dc5f | ||
|
|
b531232a9b | ||
|
|
3346c75584 | ||
|
|
aa5e03d7f6 | ||
|
|
073baf105c | ||
|
|
e97536758f | ||
|
|
7861af12e4 | ||
|
|
7f0552013a | ||
|
|
3535de5cc4 | ||
|
|
26989e54aa | ||
|
|
70d52935f0 | ||
|
|
c0e0e6afd9 | ||
|
|
0852b852f8 | ||
|
|
3a7d98a950 | ||
|
|
c5560740b6 | ||
|
|
94c6a015c8 | ||
|
|
8b6509b305 | ||
|
|
912d7c7f54 | ||
|
|
475de51c7d | ||
|
|
9f1561afe7 | ||
|
|
80c0b20877 | ||
|
|
e7721eafc6 | ||
|
|
4ead0a20cf | ||
|
|
b1527d9575 | ||
|
|
2e009cf59a | ||
|
|
780b9e1855 | ||
|
|
aef7615abd | ||
|
|
50488bd659 | ||
|
|
eb57e55fca | ||
|
|
426af2d75f | ||
|
|
345fd2f091 | ||
|
|
e1f9901384 | ||
|
|
0e7fc623b4 | ||
|
|
3e33c14376 | ||
|
|
60f4df95bd | ||
|
|
c01791ff54 | ||
|
|
980299cd54 | ||
|
|
3e8f2eba81 | ||
|
|
361cdeb296 | ||
|
|
50f76cd7c7 | ||
|
|
0f518473af | ||
|
|
a5574f92e2 | ||
|
|
abcedf892e | ||
|
|
abc3a06266 | ||
|
|
62fba9a298 | ||
|
|
e23a5ca426 | ||
|
|
e55b57d771 | ||
|
|
c4feab96fe | ||
|
|
e35cb0d84a | ||
|
|
6d6ef6dbb6 | ||
|
|
493fe4e84b |
@@ -0,0 +1,11 @@
|
|||||||
|
# Ignore everything
|
||||||
|
*
|
||||||
|
|
||||||
|
# Allow necessary files
|
||||||
|
!astrai/
|
||||||
|
!scripts/
|
||||||
|
!docs/
|
||||||
|
!csrc/
|
||||||
|
!setup.py
|
||||||
|
!pyproject.toml
|
||||||
|
!README.md
|
||||||
@@ -0,0 +1,19 @@
|
|||||||
|
# Auto detect text files
|
||||||
|
* text=auto
|
||||||
|
|
||||||
|
# Files that MUST use LF (Unix/Linux execution)
|
||||||
|
*.sh text eol=lf
|
||||||
|
*.py text eol=lf
|
||||||
|
*.md text eol=lf
|
||||||
|
*.yml text eol=lf
|
||||||
|
|
||||||
|
Dockerfile text eol=lf
|
||||||
|
.dockerignore text eol=lf
|
||||||
|
|
||||||
|
.gitignore text eol=lf
|
||||||
|
.gitattributes text eol=lf
|
||||||
|
|
||||||
|
# Windows scripts - use CRLF
|
||||||
|
*.bat text eol=crlf
|
||||||
|
*.cmd text eol=crlf
|
||||||
|
*.ps1 text eol=crlf
|
||||||
@@ -0,0 +1,27 @@
|
|||||||
|
---
|
||||||
|
name: Bug report
|
||||||
|
about: Create a report to help us improve
|
||||||
|
title: "[BUG]"
|
||||||
|
labels: bug
|
||||||
|
assignees: ''
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Description
|
||||||
|
A clear and concise description of what the bug is.
|
||||||
|
## Steps to Reproduce
|
||||||
|
1. ...
|
||||||
|
2. ...
|
||||||
|
3. ...
|
||||||
|
## Expected Behavior
|
||||||
|
What you expected to happen.
|
||||||
|
## Actual Behavior
|
||||||
|
What actually happened.
|
||||||
|
## Environment
|
||||||
|
- Python version:
|
||||||
|
- AstrAI version (or commit hash):
|
||||||
|
- Operating System:
|
||||||
|
- GPU (if applicable):
|
||||||
|
- CUDA/cuDNN version (if applicable):
|
||||||
|
## Additional Context
|
||||||
|
Add any other context, screenshots, or logs here.
|
||||||
@@ -0,0 +1,10 @@
|
|||||||
|
---
|
||||||
|
name: Custom issue template
|
||||||
|
about: Describe this issue template's purpose here.
|
||||||
|
title: ''
|
||||||
|
labels: ''
|
||||||
|
assignees: ''
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
|
||||||
@@ -0,0 +1,19 @@
|
|||||||
|
---
|
||||||
|
name: Feature request
|
||||||
|
about: Suggest an idea for this project
|
||||||
|
title: "[FEAT]"
|
||||||
|
labels: ''
|
||||||
|
assignees: ''
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Description
|
||||||
|
A clear and concise description of the feature you'd like to see.
|
||||||
|
## Problem Statement
|
||||||
|
What problem does this feature solve? Why is it needed?
|
||||||
|
## Proposed Solution
|
||||||
|
Describe the solution you'd like. Include any design ideas, API changes, or implementation details.
|
||||||
|
## Alternatives Considered
|
||||||
|
Describe any alternative solutions or features you've considered.
|
||||||
|
## Additional Context
|
||||||
|
Add any other context, screenshots, or references here.
|
||||||
@@ -0,0 +1,26 @@
|
|||||||
|
## Description
|
||||||
|
Please include a summary of the change and which issue is fixed. Please also include relevant motivation and context.
|
||||||
|
|
||||||
|
Fixes # (issue number)
|
||||||
|
|
||||||
|
## Type of Change
|
||||||
|
Please delete options that are not relevant.
|
||||||
|
|
||||||
|
- [ ] Bug fix (non-breaking change which fixes an issue)
|
||||||
|
- [ ] New feature (non-breaking change which adds functionality)
|
||||||
|
- [ ] Breaking change (fix or feature that would cause existing functionality to not work as expected)
|
||||||
|
- [ ] Documentation update
|
||||||
|
- [ ] Other (please describe):
|
||||||
|
|
||||||
|
## How Has This Been Tested?
|
||||||
|
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 . --select I`)
|
||||||
|
- [ ] I have performed a self-review of my own code
|
||||||
|
- [ ] 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
|
||||||
|
- [ ] New and existing unit tests pass locally with my changes
|
||||||
|
- [ ] Any dependent changes have been merged and published in downstream modules
|
||||||
@@ -0,0 +1,50 @@
|
|||||||
|
name: Build and Push Docker Image
|
||||||
|
|
||||||
|
on:
|
||||||
|
push:
|
||||||
|
tags:
|
||||||
|
- 'v*'
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
build:
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
permissions:
|
||||||
|
contents: read
|
||||||
|
packages: write
|
||||||
|
|
||||||
|
steps:
|
||||||
|
- name: Checkout
|
||||||
|
uses: actions/checkout@v4
|
||||||
|
|
||||||
|
- name: Set up QEMU
|
||||||
|
uses: docker/setup-qemu-action@v3
|
||||||
|
|
||||||
|
- name: Set up Docker Buildx
|
||||||
|
uses: docker/setup-buildx-action@v3
|
||||||
|
|
||||||
|
- name: Login to GitHub Container Registry
|
||||||
|
uses: docker/login-action@v3
|
||||||
|
with:
|
||||||
|
registry: ghcr.io
|
||||||
|
username: ${{ github.actor }}
|
||||||
|
password: ${{ secrets.GITHUB_TOKEN }}
|
||||||
|
|
||||||
|
- name: Extract metadata
|
||||||
|
id: meta
|
||||||
|
uses: docker/metadata-action@v5
|
||||||
|
with:
|
||||||
|
images: ghcr.io/${{ github.repository }}
|
||||||
|
tags: |
|
||||||
|
type=ref,event=tag
|
||||||
|
type=raw,value=latest
|
||||||
|
|
||||||
|
- name: Build and push
|
||||||
|
uses: docker/build-push-action@v5
|
||||||
|
with:
|
||||||
|
context: .
|
||||||
|
platforms: linux/amd64
|
||||||
|
push: true
|
||||||
|
tags: ${{ steps.meta.outputs.tags }}
|
||||||
|
labels: ${{ steps.meta.outputs.labels }}
|
||||||
|
cache-from: type=gha
|
||||||
|
cache-to: type=gha,mode=max
|
||||||
@@ -0,0 +1,31 @@
|
|||||||
|
name: Lint
|
||||||
|
|
||||||
|
on:
|
||||||
|
push:
|
||||||
|
branches: [main]
|
||||||
|
pull_request:
|
||||||
|
branches: [main]
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
lint:
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
|
||||||
|
- name: Set up Python 3.12
|
||||||
|
uses: actions/setup-python@v5
|
||||||
|
with:
|
||||||
|
python-version: "3.12"
|
||||||
|
|
||||||
|
- name: Install dependencies
|
||||||
|
run: |
|
||||||
|
pip install --upgrade pip
|
||||||
|
pip install .[dev]
|
||||||
|
|
||||||
|
- name: Check formatting with ruff
|
||||||
|
run: |
|
||||||
|
ruff format --check .
|
||||||
|
|
||||||
|
- name: Check import sorting
|
||||||
|
run: |
|
||||||
|
ruff check . --select I
|
||||||
@@ -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
|
||||||
@@ -1,17 +0,0 @@
|
|||||||
name: Spell Check
|
|
||||||
on: [push, pull_request]
|
|
||||||
|
|
||||||
permissions:
|
|
||||||
contents: read
|
|
||||||
|
|
||||||
jobs:
|
|
||||||
spellcheck:
|
|
||||||
runs-on: ubuntu-latest
|
|
||||||
steps:
|
|
||||||
- uses: actions/checkout@v4
|
|
||||||
- name: Check spelling in specific files
|
|
||||||
uses: codespell-project/actions-codespell@v2
|
|
||||||
with:
|
|
||||||
check_filenames: true
|
|
||||||
only_warn: false
|
|
||||||
path: "**/*.{md, py}"
|
|
||||||
@@ -0,0 +1,31 @@
|
|||||||
|
name: Tests
|
||||||
|
|
||||||
|
on:
|
||||||
|
push:
|
||||||
|
branches: [main]
|
||||||
|
pull_request:
|
||||||
|
branches: [main]
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
test:
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
strategy:
|
||||||
|
matrix:
|
||||||
|
python-version: ["3.12"]
|
||||||
|
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
|
||||||
|
- name: Set up Python ${{ matrix.python-version }}
|
||||||
|
uses: actions/setup-python@v5
|
||||||
|
with:
|
||||||
|
python-version: ${{ matrix.python-version }}
|
||||||
|
|
||||||
|
- name: Install dependencies
|
||||||
|
run: |
|
||||||
|
pip install --upgrade pip
|
||||||
|
pip install .[dev]
|
||||||
|
|
||||||
|
- name: Run tests with pytest
|
||||||
|
run: |
|
||||||
|
python -m pytest tests/ -v
|
||||||
+30
-5
@@ -5,8 +5,33 @@
|
|||||||
!*/
|
!*/
|
||||||
|
|
||||||
# Allow specific file types and root files
|
# Allow specific file types and root files
|
||||||
!*.py
|
!astrai/**/*.py
|
||||||
!*.md
|
!scripts/**/*.py
|
||||||
!*.png
|
!tests/**/*.py
|
||||||
!LICENSE
|
!csrc/**/*.py
|
||||||
!pyproject.toml
|
!csrc/CMakeLists.txt
|
||||||
|
|
||||||
|
!csrc/**/*.cu
|
||||||
|
!csrc/**/*.h
|
||||||
|
!csrc/**/*.cuh
|
||||||
|
|
||||||
|
!scripts/**/*.sh
|
||||||
|
|
||||||
|
# Allow GitHub files
|
||||||
|
!/.github/**
|
||||||
|
|
||||||
|
# Allow root files
|
||||||
|
!/.gitattributes
|
||||||
|
!/.dockerignore
|
||||||
|
!/Dockerfile
|
||||||
|
!/docker-compose.yml
|
||||||
|
!/docs/**
|
||||||
|
!/CONTRIBUTING.md
|
||||||
|
!/LICENSE
|
||||||
|
!/pyproject.toml
|
||||||
|
!/README.md
|
||||||
|
# Allow extension modules (only source .py)
|
||||||
|
!/astrai/extension/**/*.py
|
||||||
|
|
||||||
|
# Allow build files
|
||||||
|
!/setup.py
|
||||||
|
|||||||
+103
@@ -0,0 +1,103 @@
|
|||||||
|
# Contributing to AstrAI
|
||||||
|
|
||||||
|
Thank you for your interest in contributing! This document provides step-by-step guidelines.
|
||||||
|
|
||||||
|
## Quick Start
|
||||||
|
|
||||||
|
```bash
|
||||||
|
git clone https://github.com/ViperEkura/AstrAI.git
|
||||||
|
cd AstrAI
|
||||||
|
pip install -e ".[dev]" # install with dev dependencies (pytest, ruff)
|
||||||
|
```
|
||||||
|
|
||||||
|
## Before You Commit
|
||||||
|
|
||||||
|
Run the following checks **in order** — CI will reject if any fail.
|
||||||
|
|
||||||
|
### 1. Format
|
||||||
|
|
||||||
|
```bash
|
||||||
|
ruff format .
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2. Import sorting
|
||||||
|
|
||||||
|
```bash
|
||||||
|
ruff check . --select I
|
||||||
|
```
|
||||||
|
|
||||||
|
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 the system temp directory
|
||||||
|
> (`$TMPDIR` on Linux/macOS, `%TEMP%` on Windows). Clean them manually if needed.
|
||||||
|
|
||||||
|
### 4. (Optional) Full pre-commit check script
|
||||||
|
|
||||||
|
If you have `bash` available (Git Bash on Windows works too):
|
||||||
|
|
||||||
|
```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 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 [Apache-2.0 License](LICENSE).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
Questions? Ask in [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions) or open an issue.
|
||||||
+70
@@ -0,0 +1,70 @@
|
|||||||
|
# AstrAI Dockerfile - Multi-stage Build (Optimized)
|
||||||
|
#
|
||||||
|
# CUDA version selection:
|
||||||
|
# docker build -t astrai .
|
||||||
|
# docker build -t astrai --build-arg CUDA_TAG=cu128 .
|
||||||
|
# docker build -t astrai --build-arg CUDA_TAG=cu130 .
|
||||||
|
# Default: cu128
|
||||||
|
|
||||||
|
# Build stage - use base image with minimal build tools
|
||||||
|
FROM ubuntu:24.04 AS builder
|
||||||
|
|
||||||
|
ARG CUDA_TAG=cu128
|
||||||
|
|
||||||
|
WORKDIR /app
|
||||||
|
|
||||||
|
# Install Python 3.12 and minimal build dependencies
|
||||||
|
RUN apt-get update && DEBIAN_FRONTEND=noninteractive apt-get install -y --no-install-recommends \
|
||||||
|
python3.12 \
|
||||||
|
python3.12-dev \
|
||||||
|
python3.12-venv \
|
||||||
|
gcc \
|
||||||
|
g++ \
|
||||||
|
&& rm -rf /var/lib/apt/lists/*
|
||||||
|
|
||||||
|
# Create isolated virtual environment
|
||||||
|
RUN python3.12 -m venv --copies /opt/venv
|
||||||
|
ENV PATH="/opt/venv/bin:$PATH"
|
||||||
|
|
||||||
|
# Copy source code and install (deps read from pyproject.toml)
|
||||||
|
COPY astrai/ ./astrai/
|
||||||
|
COPY csrc/ ./csrc/
|
||||||
|
COPY setup.py .
|
||||||
|
COPY pyproject.toml .
|
||||||
|
RUN pip install --no-cache-dir --upgrade pip \
|
||||||
|
&& pip install --no-cache-dir . \
|
||||||
|
--extra-index-url "https://download.pytorch.org/whl/${CUDA_TAG}"
|
||||||
|
|
||||||
|
# Production stage
|
||||||
|
FROM ubuntu:24.04 AS production
|
||||||
|
|
||||||
|
WORKDIR /app
|
||||||
|
|
||||||
|
# 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
|
||||||
|
COPY --from=builder /opt/venv /opt/venv
|
||||||
|
ENV PATH="/opt/venv/bin:$PATH"
|
||||||
|
|
||||||
|
# Copy application code
|
||||||
|
COPY astrai/ ./astrai/
|
||||||
|
COPY scripts/ ./scripts/
|
||||||
|
COPY docs/ ./docs/
|
||||||
|
COPY pyproject.toml .
|
||||||
|
COPY README.md .
|
||||||
|
|
||||||
|
# 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 \
|
||||||
|
PYTHONDONTWRITEBYTECODE=1
|
||||||
@@ -1,674 +1,201 @@
|
|||||||
GNU GENERAL PUBLIC LICENSE
|
Apache License
|
||||||
Version 3, 29 June 2007
|
Version 2.0, January 2004
|
||||||
|
http://www.apache.org/licenses/
|
||||||
Copyright (C) 2007 Free Software Foundation, Inc. <https://fsf.org/>
|
|
||||||
Everyone is permitted to copy and distribute verbatim copies
|
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||||
of this license document, but changing it is not allowed.
|
|
||||||
|
1. Definitions.
|
||||||
Preamble
|
|
||||||
|
"License" shall mean the terms and conditions for use, reproduction,
|
||||||
The GNU General Public License is a free, copyleft license for
|
and distribution as defined by Sections 1 through 9 of this document.
|
||||||
software and other kinds of works.
|
|
||||||
|
"Licensor" shall mean the copyright owner or entity authorized by
|
||||||
The licenses for most software and other practical works are designed
|
the copyright owner that is granting the License.
|
||||||
to take away your freedom to share and change the works. By contrast,
|
|
||||||
the GNU General Public License is intended to guarantee your freedom to
|
"Legal Entity" shall mean the union of the acting entity and all
|
||||||
share and change all versions of a program--to make sure it remains free
|
other entities that control, are controlled by, or are under common
|
||||||
software for all its users. We, the Free Software Foundation, use the
|
control with that entity. For the purposes of this definition,
|
||||||
GNU General Public License for most of our software; it applies also to
|
"control" means (i) the power, direct or indirect, to cause the
|
||||||
any other work released this way by its authors. You can apply it to
|
direction or management of such entity, whether by contract or
|
||||||
your programs, too.
|
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||||
|
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||||
When we speak of free software, we are referring to freedom, not
|
|
||||||
price. Our General Public Licenses are designed to make sure that you
|
"You" (or "Your") shall mean an individual or Legal Entity
|
||||||
have the freedom to distribute copies of free software (and charge for
|
exercising permissions granted by this License.
|
||||||
them if you wish), that you receive source code or can get it if you
|
|
||||||
want it, that you can change the software or use pieces of it in new
|
"Source" form shall mean the preferred form for making modifications,
|
||||||
free programs, and that you know you can do these things.
|
including but not limited to software source code, documentation
|
||||||
|
source, and configuration files.
|
||||||
To protect your rights, we need to prevent others from denying you
|
|
||||||
these rights or asking you to surrender the rights. Therefore, you have
|
"Object" form shall mean any form resulting from mechanical
|
||||||
certain responsibilities if you distribute copies of the software, or if
|
transformation or translation of a Source form, including but
|
||||||
you modify it: responsibilities to respect the freedom of others.
|
not limited to compiled object code, generated documentation,
|
||||||
|
and conversions to other media types.
|
||||||
For example, if you distribute copies of such a program, whether
|
|
||||||
gratis or for a fee, you must pass on to the recipients the same
|
"Work" shall mean the work of authorship, whether in Source or
|
||||||
freedoms that you received. You must make sure that they, too, receive
|
Object form, made available under the License, as indicated by a
|
||||||
or can get the source code. And you must show them these terms so they
|
copyright notice that is included in or attached to the work
|
||||||
know their rights.
|
(an example is provided in the Appendix below).
|
||||||
|
|
||||||
Developers that use the GNU GPL protect your rights with two steps:
|
"Derivative Works" shall mean any work, whether in Source or Object
|
||||||
(1) assert copyright on the software, and (2) offer you this License
|
form, that is based on (or derived from) the Work and for which the
|
||||||
giving you legal permission to copy, distribute and/or modify it.
|
editorial revisions, annotations, elaborations, or other modifications
|
||||||
|
represent, as a whole, an original work of authorship. For the purposes
|
||||||
For the developers' and authors' protection, the GPL clearly explains
|
of this License, Derivative Works shall not include works that remain
|
||||||
that there is no warranty for this free software. For both users' and
|
separable from, or merely link (or bind by name) to the interfaces of,
|
||||||
authors' sake, the GPL requires that modified versions be marked as
|
the Work and Derivative Works thereof.
|
||||||
changed, so that their problems will not be attributed erroneously to
|
|
||||||
authors of previous versions.
|
"Contribution" shall mean any work of authorship, including
|
||||||
|
the original version of the Work and any modifications or additions
|
||||||
Some devices are designed to deny users access to install or run
|
to that Work or Derivative Works thereof, that is intentionally
|
||||||
modified versions of the software inside them, although the manufacturer
|
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||||
can do so. This is fundamentally incompatible with the aim of
|
or by an individual or Legal Entity authorized to submit on behalf of
|
||||||
protecting users' freedom to change the software. The systematic
|
the copyright owner. For the purposes of this definition, "submitted"
|
||||||
pattern of such abuse occurs in the area of products for individuals to
|
means any form of electronic, verbal, or written communication sent
|
||||||
use, which is precisely where it is most unacceptable. Therefore, we
|
to the Licensor or its representatives, including but not limited to
|
||||||
have designed this version of the GPL to prohibit the practice for those
|
communication on electronic mailing lists, source code control systems,
|
||||||
products. If such problems arise substantially in other domains, we
|
and issue tracking systems that are managed by, or on behalf of, the
|
||||||
stand ready to extend this provision to those domains in future versions
|
Licensor for the purpose of discussing and improving the Work, but
|
||||||
of the GPL, as needed to protect the freedom of users.
|
excluding communication that is conspicuously marked or otherwise
|
||||||
|
designated in writing by the copyright owner as "Not a Contribution."
|
||||||
Finally, every program is threatened constantly by software patents.
|
|
||||||
States should not allow patents to restrict development and use of
|
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||||
software on general-purpose computers, but in those that do, we wish to
|
on behalf of whom a Contribution has been received by Licensor and
|
||||||
avoid the special danger that patents applied to a free program could
|
subsequently incorporated within the Work.
|
||||||
make it effectively proprietary. To prevent this, the GPL assures that
|
|
||||||
patents cannot be used to render the program non-free.
|
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||||
|
this License, each Contributor hereby grants to You a perpetual,
|
||||||
The precise terms and conditions for copying, distribution and
|
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||||
modification follow.
|
copyright license to reproduce, prepare Derivative Works of,
|
||||||
|
publicly display, publicly perform, sublicense, and distribute the
|
||||||
TERMS AND CONDITIONS
|
Work and such Derivative Works in Source or Object form.
|
||||||
|
|
||||||
0. Definitions.
|
3. Grant of Patent License. Subject to the terms and conditions of
|
||||||
|
this License, each Contributor hereby grants to You a perpetual,
|
||||||
"This License" refers to version 3 of the GNU General Public License.
|
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||||
|
(except as stated in this section) patent license to make, have made,
|
||||||
"Copyright" also means copyright-like laws that apply to other kinds of
|
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||||
works, such as semiconductor masks.
|
where such license applies only to those patent claims licensable
|
||||||
|
by such Contributor that are necessarily infringed by their
|
||||||
"The Program" refers to any copyrightable work licensed under this
|
Contribution(s) alone or by combination of their Contribution(s)
|
||||||
License. Each licensee is addressed as "you". "Licensees" and
|
with the Work to which such Contribution(s) was submitted. If You
|
||||||
"recipients" may be individuals or organizations.
|
institute patent litigation against any entity (including a
|
||||||
|
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||||
To "modify" a work means to copy from or adapt all or part of the work
|
or a Contribution incorporated within the Work constitutes direct
|
||||||
in a fashion requiring copyright permission, other than the making of an
|
or contributory patent infringement, then any patent licenses
|
||||||
exact copy. The resulting work is called a "modified version" of the
|
granted to You under this License for that Work shall terminate
|
||||||
earlier work or a work "based on" the earlier work.
|
as of the date such litigation is filed.
|
||||||
|
|
||||||
A "covered work" means either the unmodified Program or a work based
|
4. Redistribution. You may reproduce and distribute copies of the
|
||||||
on the Program.
|
Work or Derivative Works thereof in any medium, with or without
|
||||||
|
modifications, and in Source or Object form, provided that You
|
||||||
To "propagate" a work means to do anything with it that, without
|
meet the following conditions:
|
||||||
permission, would make you directly or secondarily liable for
|
|
||||||
infringement under applicable copyright law, except executing it on a
|
(a) You must give any other recipients of the Work or
|
||||||
computer or modifying a private copy. Propagation includes copying,
|
Derivative Works a copy of this License; and
|
||||||
distribution (with or without modification), making available to the
|
|
||||||
public, and in some countries other activities as well.
|
(b) You must cause any modified files to carry prominent notices
|
||||||
|
stating that You changed the files; and
|
||||||
To "convey" a work means any kind of propagation that enables other
|
|
||||||
parties to make or receive copies. Mere interaction with a user through
|
(c) You must retain, in the Source form of any Derivative Works
|
||||||
a computer network, with no transfer of a copy, is not conveying.
|
that You distribute, all copyright, patent, trademark, and
|
||||||
|
attribution notices from the Source form of the Work,
|
||||||
An interactive user interface displays "Appropriate Legal Notices"
|
excluding those notices that do not pertain to any part of
|
||||||
to the extent that it includes a convenient and prominently visible
|
the Derivative Works; and
|
||||||
feature that (1) displays an appropriate copyright notice, and (2)
|
|
||||||
tells the user that there is no warranty for the work (except to the
|
(d) If the Work includes a "NOTICE" text file as part of its
|
||||||
extent that warranties are provided), that licensees may convey the
|
distribution, then any Derivative Works that You distribute must
|
||||||
work under this License, and how to view a copy of this License. If
|
include a readable copy of the attribution notices contained
|
||||||
the interface presents a list of user commands or options, such as a
|
within such NOTICE file, excluding those notices that do not
|
||||||
menu, a prominent item in the list meets this criterion.
|
pertain to any part of the Derivative Works, in at least one
|
||||||
|
of the following places: within a NOTICE text file distributed
|
||||||
1. Source Code.
|
as part of the Derivative Works; within the Source form or
|
||||||
|
documentation, if provided along with the Derivative Works; or,
|
||||||
The "source code" for a work means the preferred form of the work
|
within a display generated by the Derivative Works, if and
|
||||||
for making modifications to it. "Object code" means any non-source
|
wherever such third-party notices normally appear. The contents
|
||||||
form of a work.
|
of the NOTICE file are for informational purposes only and
|
||||||
|
do not modify the License. You may add Your own attribution
|
||||||
A "Standard Interface" means an interface that either is an official
|
notices within Derivative Works that You distribute, alongside
|
||||||
standard defined by a recognized standards body, or, in the case of
|
or as an addendum to the NOTICE text from the Work, provided
|
||||||
interfaces specified for a particular programming language, one that
|
that such additional attribution notices cannot be construed
|
||||||
is widely used among developers working in that language.
|
as modifying the License.
|
||||||
|
|
||||||
The "System Libraries" of an executable work include anything, other
|
You may add Your own copyright statement to Your modifications and
|
||||||
than the work as a whole, that (a) is included in the normal form of
|
may provide additional or different license terms and conditions
|
||||||
packaging a Major Component, but which is not part of that Major
|
for use, reproduction, or distribution of Your modifications, or
|
||||||
Component, and (b) serves only to enable use of the work with that
|
for any such Derivative Works as a whole, provided Your use,
|
||||||
Major Component, or to implement a Standard Interface for which an
|
reproduction, and distribution of the Work otherwise complies with
|
||||||
implementation is available to the public in source code form. A
|
the conditions stated in this License.
|
||||||
"Major Component", in this context, means a major essential component
|
|
||||||
(kernel, window system, and so on) of the specific operating system
|
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||||
(if any) on which the executable work runs, or a compiler used to
|
any Contribution intentionally submitted for inclusion in the Work
|
||||||
produce the work, or an object code interpreter used to run it.
|
by You to the Licensor shall be under the terms and conditions of
|
||||||
|
this License, without any additional terms or conditions.
|
||||||
The "Corresponding Source" for a work in object code form means all
|
Notwithstanding the above, nothing herein shall supersede or modify
|
||||||
the source code needed to generate, install, and (for an executable
|
the terms of any separate license agreement you may have executed
|
||||||
work) run the object code and to modify the work, including scripts to
|
with Licensor regarding such Contributions.
|
||||||
control those activities. However, it does not include the work's
|
|
||||||
System Libraries, or general-purpose tools or generally available free
|
6. Trademarks. This License does not grant permission to use the trade
|
||||||
programs which are used unmodified in performing those activities but
|
names, trademarks, service marks, or product names of the Licensor,
|
||||||
which are not part of the work. For example, Corresponding Source
|
except as required for reasonable and customary use in describing the
|
||||||
includes interface definition files associated with source files for
|
origin of the Work and reproducing the content of the NOTICE file.
|
||||||
the work, and the source code for shared libraries and dynamically
|
|
||||||
linked subprograms that the work is specifically designed to require,
|
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||||
such as by intimate data communication or control flow between those
|
agreed to in writing, Licensor provides the Work (and each
|
||||||
subprograms and other parts of the work.
|
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||||
|
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||||
The Corresponding Source need not include anything that users
|
implied, including, without limitation, any warranties or conditions
|
||||||
can regenerate automatically from other parts of the Corresponding
|
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||||
Source.
|
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||||
|
appropriateness of using or redistributing the Work and assume any
|
||||||
The Corresponding Source for a work in source code form is that
|
risks associated with Your exercise of permissions under this License.
|
||||||
same work.
|
|
||||||
|
8. Limitation of Liability. In no event and under no legal theory,
|
||||||
2. Basic Permissions.
|
whether in tort (including negligence), contract, or otherwise,
|
||||||
|
unless required by applicable law (such as deliberate and grossly
|
||||||
All rights granted under this License are granted for the term of
|
negligent acts) or agreed to in writing, shall any Contributor be
|
||||||
copyright on the Program, and are irrevocable provided the stated
|
liable to You for damages, including any direct, indirect, special,
|
||||||
conditions are met. This License explicitly affirms your unlimited
|
incidental, or consequential damages of any character arising as a
|
||||||
permission to run the unmodified Program. The output from running a
|
result of this License or out of the use or inability to use the
|
||||||
covered work is covered by this License only if the output, given its
|
Work (including but not limited to damages for loss of goodwill,
|
||||||
content, constitutes a covered work. This License acknowledges your
|
work stoppage, computer failure or malfunction, or any and all
|
||||||
rights of fair use or other equivalent, as provided by copyright law.
|
other commercial damages or losses), even if such Contributor
|
||||||
|
has been advised of the possibility of such damages.
|
||||||
You may make, run and propagate covered works that you do not
|
|
||||||
convey, without conditions so long as your license otherwise remains
|
9. Accepting Warranty or Additional Liability. While redistributing
|
||||||
in force. You may convey covered works to others for the sole purpose
|
the Work or Derivative Works thereof, You may choose to offer,
|
||||||
of having them make modifications exclusively for you, or provide you
|
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||||
with facilities for running those works, provided that you comply with
|
or other liability obligations and/or rights consistent with this
|
||||||
the terms of this License in conveying all material for which you do
|
License. However, in accepting such obligations, You may act only
|
||||||
not control copyright. Those thus making or running the covered works
|
on Your own behalf and on Your sole responsibility, not on behalf
|
||||||
for you must do so exclusively on your behalf, under your direction
|
of any other Contributor, and only if You agree to indemnify,
|
||||||
and control, on terms that prohibit them from making any copies of
|
defend, and hold each Contributor harmless for any liability
|
||||||
your copyrighted material outside their relationship with you.
|
incurred by, or claims asserted against, such Contributor by reason
|
||||||
|
of your accepting any such warranty or additional liability.
|
||||||
Conveying under any other circumstances is permitted solely under
|
|
||||||
the conditions stated below. Sublicensing is not allowed; section 10
|
END OF TERMS AND CONDITIONS
|
||||||
makes it unnecessary.
|
|
||||||
|
APPENDIX: How to apply the Apache License to your work.
|
||||||
3. Protecting Users' Legal Rights From Anti-Circumvention Law.
|
|
||||||
|
To apply the Apache License to your work, attach the following
|
||||||
No covered work shall be deemed part of an effective technological
|
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||||
measure under any applicable law fulfilling obligations under article
|
replaced with your own identifying information. (Don't include
|
||||||
11 of the WIPO copyright treaty adopted on 20 December 1996, or
|
the brackets!) The text should be enclosed in the appropriate
|
||||||
similar laws prohibiting or restricting circumvention of such
|
comment syntax for the file format. We also recommend that a
|
||||||
measures.
|
file or class name and description of purpose be included on the
|
||||||
|
same "printed page" as the copyright notice for easier
|
||||||
When you convey a covered work, you waive any legal power to forbid
|
identification within third-party archives.
|
||||||
circumvention of technological measures to the extent such circumvention
|
|
||||||
is effected by exercising rights under this License with respect to
|
Copyright [yyyy] [name of copyright owner]
|
||||||
the covered work, and you disclaim any intention to limit operation or
|
|
||||||
modification of the work as a means of enforcing, against the work's
|
Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
users, your or third parties' legal rights to forbid circumvention of
|
you may not use this file except in compliance with the License.
|
||||||
technological measures.
|
You may obtain a copy of the License at
|
||||||
|
|
||||||
4. Conveying Verbatim Copies.
|
http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
|
||||||
You may convey verbatim copies of the Program's source code as you
|
Unless required by applicable law or agreed to in writing, software
|
||||||
receive it, in any medium, provided that you conspicuously and
|
distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
appropriately publish on each copy an appropriate copyright notice;
|
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
keep intact all notices stating that this License and any
|
See the License for the specific language governing permissions and
|
||||||
non-permissive terms added in accord with section 7 apply to the code;
|
limitations under the License.
|
||||||
keep intact all notices of the absence of any warranty; and give all
|
|
||||||
recipients a copy of this License along with the Program.
|
|
||||||
|
|
||||||
You may charge any price or no price for each copy that you convey,
|
|
||||||
and you may offer support or warranty protection for a fee.
|
|
||||||
|
|
||||||
5. Conveying Modified Source Versions.
|
|
||||||
|
|
||||||
You may convey a work based on the Program, or the modifications to
|
|
||||||
produce it from the Program, in the form of source code under the
|
|
||||||
terms of section 4, provided that you also meet all of these conditions:
|
|
||||||
|
|
||||||
a) The work must carry prominent notices stating that you modified
|
|
||||||
it, and giving a relevant date.
|
|
||||||
|
|
||||||
b) The work must carry prominent notices stating that it is
|
|
||||||
released under this License and any conditions added under section
|
|
||||||
7. This requirement modifies the requirement in section 4 to
|
|
||||||
"keep intact all notices".
|
|
||||||
|
|
||||||
c) You must license the entire work, as a whole, under this
|
|
||||||
License to anyone who comes into possession of a copy. This
|
|
||||||
License will therefore apply, along with any applicable section 7
|
|
||||||
additional terms, to the whole of the work, and all its parts,
|
|
||||||
regardless of how they are packaged. This License gives no
|
|
||||||
permission to license the work in any other way, but it does not
|
|
||||||
invalidate such permission if you have separately received it.
|
|
||||||
|
|
||||||
d) If the work has interactive user interfaces, each must display
|
|
||||||
Appropriate Legal Notices; however, if the Program has interactive
|
|
||||||
interfaces that do not display Appropriate Legal Notices, your
|
|
||||||
work need not make them do so.
|
|
||||||
|
|
||||||
A compilation of a covered work with other separate and independent
|
|
||||||
works, which are not by their nature extensions of the covered work,
|
|
||||||
and which are not combined with it such as to form a larger program,
|
|
||||||
in or on a volume of a storage or distribution medium, is called an
|
|
||||||
"aggregate" if the compilation and its resulting copyright are not
|
|
||||||
used to limit the access or legal rights of the compilation's users
|
|
||||||
beyond what the individual works permit. Inclusion of a covered work
|
|
||||||
in an aggregate does not cause this License to apply to the other
|
|
||||||
parts of the aggregate.
|
|
||||||
|
|
||||||
6. Conveying Non-Source Forms.
|
|
||||||
|
|
||||||
You may convey a covered work in object code form under the terms
|
|
||||||
of sections 4 and 5, provided that you also convey the
|
|
||||||
machine-readable Corresponding Source under the terms of this License,
|
|
||||||
in one of these ways:
|
|
||||||
|
|
||||||
a) Convey the object code in, or embodied in, a physical product
|
|
||||||
(including a physical distribution medium), accompanied by the
|
|
||||||
Corresponding Source fixed on a durable physical medium
|
|
||||||
customarily used for software interchange.
|
|
||||||
|
|
||||||
b) Convey the object code in, or embodied in, a physical product
|
|
||||||
(including a physical distribution medium), accompanied by a
|
|
||||||
written offer, valid for at least three years and valid for as
|
|
||||||
long as you offer spare parts or customer support for that product
|
|
||||||
model, to give anyone who possesses the object code either (1) a
|
|
||||||
copy of the Corresponding Source for all the software in the
|
|
||||||
product that is covered by this License, on a durable physical
|
|
||||||
medium customarily used for software interchange, for a price no
|
|
||||||
more than your reasonable cost of physically performing this
|
|
||||||
conveying of source, or (2) access to copy the
|
|
||||||
Corresponding Source from a network server at no charge.
|
|
||||||
|
|
||||||
c) Convey individual copies of the object code with a copy of the
|
|
||||||
written offer to provide the Corresponding Source. This
|
|
||||||
alternative is allowed only occasionally and noncommercially, and
|
|
||||||
only if you received the object code with such an offer, in accord
|
|
||||||
with subsection 6b.
|
|
||||||
|
|
||||||
d) Convey the object code by offering access from a designated
|
|
||||||
place (gratis or for a charge), and offer equivalent access to the
|
|
||||||
Corresponding Source in the same way through the same place at no
|
|
||||||
further charge. You need not require recipients to copy the
|
|
||||||
Corresponding Source along with the object code. If the place to
|
|
||||||
copy the object code is a network server, the Corresponding Source
|
|
||||||
may be on a different server (operated by you or a third party)
|
|
||||||
that supports equivalent copying facilities, provided you maintain
|
|
||||||
clear directions next to the object code saying where to find the
|
|
||||||
Corresponding Source. Regardless of what server hosts the
|
|
||||||
Corresponding Source, you remain obligated to ensure that it is
|
|
||||||
available for as long as needed to satisfy these requirements.
|
|
||||||
|
|
||||||
e) Convey the object code using peer-to-peer transmission, provided
|
|
||||||
you inform other peers where the object code and Corresponding
|
|
||||||
Source of the work are being offered to the general public at no
|
|
||||||
charge under subsection 6d.
|
|
||||||
|
|
||||||
A separable portion of the object code, whose source code is excluded
|
|
||||||
from the Corresponding Source as a System Library, need not be
|
|
||||||
included in conveying the object code work.
|
|
||||||
|
|
||||||
A "User Product" is either (1) a "consumer product", which means any
|
|
||||||
tangible personal property which is normally used for personal, family,
|
|
||||||
or household purposes, or (2) anything designed or sold for incorporation
|
|
||||||
into a dwelling. In determining whether a product is a consumer product,
|
|
||||||
doubtful cases shall be resolved in favor of coverage. For a particular
|
|
||||||
product received by a particular user, "normally used" refers to a
|
|
||||||
typical or common use of that class of product, regardless of the status
|
|
||||||
of the particular user or of the way in which the particular user
|
|
||||||
actually uses, or expects or is expected to use, the product. A product
|
|
||||||
is a consumer product regardless of whether the product has substantial
|
|
||||||
commercial, industrial or non-consumer uses, unless such uses represent
|
|
||||||
the only significant mode of use of the product.
|
|
||||||
|
|
||||||
"Installation Information" for a User Product means any methods,
|
|
||||||
procedures, authorization keys, or other information required to install
|
|
||||||
and execute modified versions of a covered work in that User Product from
|
|
||||||
a modified version of its Corresponding Source. The information must
|
|
||||||
suffice to ensure that the continued functioning of the modified object
|
|
||||||
code is in no case prevented or interfered with solely because
|
|
||||||
modification has been made.
|
|
||||||
|
|
||||||
If you convey an object code work under this section in, or with, or
|
|
||||||
specifically for use in, a User Product, and the conveying occurs as
|
|
||||||
part of a transaction in which the right of possession and use of the
|
|
||||||
User Product is transferred to the recipient in perpetuity or for a
|
|
||||||
fixed term (regardless of how the transaction is characterized), the
|
|
||||||
Corresponding Source conveyed under this section must be accompanied
|
|
||||||
by the Installation Information. But this requirement does not apply
|
|
||||||
if neither you nor any third party retains the ability to install
|
|
||||||
modified object code on the User Product (for example, the work has
|
|
||||||
been installed in ROM).
|
|
||||||
|
|
||||||
The requirement to provide Installation Information does not include a
|
|
||||||
requirement to continue to provide support service, warranty, or updates
|
|
||||||
for a work that has been modified or installed by the recipient, or for
|
|
||||||
the User Product in which it has been modified or installed. Access to a
|
|
||||||
network may be denied when the modification itself materially and
|
|
||||||
adversely affects the operation of the network or violates the rules and
|
|
||||||
protocols for communication across the network.
|
|
||||||
|
|
||||||
Corresponding Source conveyed, and Installation Information provided,
|
|
||||||
in accord with this section must be in a format that is publicly
|
|
||||||
documented (and with an implementation available to the public in
|
|
||||||
source code form), and must require no special password or key for
|
|
||||||
unpacking, reading or copying.
|
|
||||||
|
|
||||||
7. Additional Terms.
|
|
||||||
|
|
||||||
"Additional permissions" are terms that supplement the terms of this
|
|
||||||
License by making exceptions from one or more of its conditions.
|
|
||||||
Additional permissions that are applicable to the entire Program shall
|
|
||||||
be treated as though they were included in this License, to the extent
|
|
||||||
that they are valid under applicable law. If additional permissions
|
|
||||||
apply only to part of the Program, that part may be used separately
|
|
||||||
under those permissions, but the entire Program remains governed by
|
|
||||||
this License without regard to the additional permissions.
|
|
||||||
|
|
||||||
When you convey a copy of a covered work, you may at your option
|
|
||||||
remove any additional permissions from that copy, or from any part of
|
|
||||||
it. (Additional permissions may be written to require their own
|
|
||||||
removal in certain cases when you modify the work.) You may place
|
|
||||||
additional permissions on material, added by you to a covered work,
|
|
||||||
for which you have or can give appropriate copyright permission.
|
|
||||||
|
|
||||||
Notwithstanding any other provision of this License, for material you
|
|
||||||
add to a covered work, you may (if authorized by the copyright holders of
|
|
||||||
that material) supplement the terms of this License with terms:
|
|
||||||
|
|
||||||
a) Disclaiming warranty or limiting liability differently from the
|
|
||||||
terms of sections 15 and 16 of this License; or
|
|
||||||
|
|
||||||
b) Requiring preservation of specified reasonable legal notices or
|
|
||||||
author attributions in that material or in the Appropriate Legal
|
|
||||||
Notices displayed by works containing it; or
|
|
||||||
|
|
||||||
c) Prohibiting misrepresentation of the origin of that material, or
|
|
||||||
requiring that modified versions of such material be marked in
|
|
||||||
reasonable ways as different from the original version; or
|
|
||||||
|
|
||||||
d) Limiting the use for publicity purposes of names of licensors or
|
|
||||||
authors of the material; or
|
|
||||||
|
|
||||||
e) Declining to grant rights under trademark law for use of some
|
|
||||||
trade names, trademarks, or service marks; or
|
|
||||||
|
|
||||||
f) Requiring indemnification of licensors and authors of that
|
|
||||||
material by anyone who conveys the material (or modified versions of
|
|
||||||
it) with contractual assumptions of liability to the recipient, for
|
|
||||||
any liability that these contractual assumptions directly impose on
|
|
||||||
those licensors and authors.
|
|
||||||
|
|
||||||
All other non-permissive additional terms are considered "further
|
|
||||||
restrictions" within the meaning of section 10. If the Program as you
|
|
||||||
received it, or any part of it, contains a notice stating that it is
|
|
||||||
governed by this License along with a term that is a further
|
|
||||||
restriction, you may remove that term. If a license document contains
|
|
||||||
a further restriction but permits relicensing or conveying under this
|
|
||||||
License, you may add to a covered work material governed by the terms
|
|
||||||
of that license document, provided that the further restriction does
|
|
||||||
not survive such relicensing or conveying.
|
|
||||||
|
|
||||||
If you add terms to a covered work in accord with this section, you
|
|
||||||
must place, in the relevant source files, a statement of the
|
|
||||||
additional terms that apply to those files, or a notice indicating
|
|
||||||
where to find the applicable terms.
|
|
||||||
|
|
||||||
Additional terms, permissive or non-permissive, may be stated in the
|
|
||||||
form of a separately written license, or stated as exceptions;
|
|
||||||
the above requirements apply either way.
|
|
||||||
|
|
||||||
8. Termination.
|
|
||||||
|
|
||||||
You may not propagate or modify a covered work except as expressly
|
|
||||||
provided under this License. Any attempt otherwise to propagate or
|
|
||||||
modify it is void, and will automatically terminate your rights under
|
|
||||||
this License (including any patent licenses granted under the third
|
|
||||||
paragraph of section 11).
|
|
||||||
|
|
||||||
However, if you cease all violation of this License, then your
|
|
||||||
license from a particular copyright holder is reinstated (a)
|
|
||||||
provisionally, unless and until the copyright holder explicitly and
|
|
||||||
finally terminates your license, and (b) permanently, if the copyright
|
|
||||||
holder fails to notify you of the violation by some reasonable means
|
|
||||||
prior to 60 days after the cessation.
|
|
||||||
|
|
||||||
Moreover, your license from a particular copyright holder is
|
|
||||||
reinstated permanently if the copyright holder notifies you of the
|
|
||||||
violation by some reasonable means, this is the first time you have
|
|
||||||
received notice of violation of this License (for any work) from that
|
|
||||||
copyright holder, and you cure the violation prior to 30 days after
|
|
||||||
your receipt of the notice.
|
|
||||||
|
|
||||||
Termination of your rights under this section does not terminate the
|
|
||||||
licenses of parties who have received copies or rights from you under
|
|
||||||
this License. If your rights have been terminated and not permanently
|
|
||||||
reinstated, you do not qualify to receive new licenses for the same
|
|
||||||
material under section 10.
|
|
||||||
|
|
||||||
9. Acceptance Not Required for Having Copies.
|
|
||||||
|
|
||||||
You are not required to accept this License in order to receive or
|
|
||||||
run a copy of the Program. Ancillary propagation of a covered work
|
|
||||||
occurring solely as a consequence of using peer-to-peer transmission
|
|
||||||
to receive a copy likewise does not require acceptance. However,
|
|
||||||
nothing other than this License grants you permission to propagate or
|
|
||||||
modify any covered work. These actions infringe copyright if you do
|
|
||||||
not accept this License. Therefore, by modifying or propagating a
|
|
||||||
covered work, you indicate your acceptance of this License to do so.
|
|
||||||
|
|
||||||
10. Automatic Licensing of Downstream Recipients.
|
|
||||||
|
|
||||||
Each time you convey a covered work, the recipient automatically
|
|
||||||
receives a license from the original licensors, to run, modify and
|
|
||||||
propagate that work, subject to this License. You are not responsible
|
|
||||||
for enforcing compliance by third parties with this License.
|
|
||||||
|
|
||||||
An "entity transaction" is a transaction transferring control of an
|
|
||||||
organization, or substantially all assets of one, or subdividing an
|
|
||||||
organization, or merging organizations. If propagation of a covered
|
|
||||||
work results from an entity transaction, each party to that
|
|
||||||
transaction who receives a copy of the work also receives whatever
|
|
||||||
licenses to the work the party's predecessor in interest had or could
|
|
||||||
give under the previous paragraph, plus a right to possession of the
|
|
||||||
Corresponding Source of the work from the predecessor in interest, if
|
|
||||||
the predecessor has it or can get it with reasonable efforts.
|
|
||||||
|
|
||||||
You may not impose any further restrictions on the exercise of the
|
|
||||||
rights granted or affirmed under this License. For example, you may
|
|
||||||
not impose a license fee, royalty, or other charge for exercise of
|
|
||||||
rights granted under this License, and you may not initiate litigation
|
|
||||||
(including a cross-claim or counterclaim in a lawsuit) alleging that
|
|
||||||
any patent claim is infringed by making, using, selling, offering for
|
|
||||||
sale, or importing the Program or any portion of it.
|
|
||||||
|
|
||||||
11. Patents.
|
|
||||||
|
|
||||||
A "contributor" is a copyright holder who authorizes use under this
|
|
||||||
License of the Program or a work on which the Program is based. The
|
|
||||||
work thus licensed is called the contributor's "contributor version".
|
|
||||||
|
|
||||||
A contributor's "essential patent claims" are all patent claims
|
|
||||||
owned or controlled by the contributor, whether already acquired or
|
|
||||||
hereafter acquired, that would be infringed by some manner, permitted
|
|
||||||
by this License, of making, using, or selling its contributor version,
|
|
||||||
but do not include claims that would be infringed only as a
|
|
||||||
consequence of further modification of the contributor version. For
|
|
||||||
purposes of this definition, "control" includes the right to grant
|
|
||||||
patent sublicenses in a manner consistent with the requirements of
|
|
||||||
this License.
|
|
||||||
|
|
||||||
Each contributor grants you a non-exclusive, worldwide, royalty-free
|
|
||||||
patent license under the contributor's essential patent claims, to
|
|
||||||
make, use, sell, offer for sale, import and otherwise run, modify and
|
|
||||||
propagate the contents of its contributor version.
|
|
||||||
|
|
||||||
In the following three paragraphs, a "patent license" is any express
|
|
||||||
agreement or commitment, however denominated, not to enforce a patent
|
|
||||||
(such as an express permission to practice a patent or covenant not to
|
|
||||||
sue for patent infringement). To "grant" such a patent license to a
|
|
||||||
party means to make such an agreement or commitment not to enforce a
|
|
||||||
patent against the party.
|
|
||||||
|
|
||||||
If you convey a covered work, knowingly relying on a patent license,
|
|
||||||
and the Corresponding Source of the work is not available for anyone
|
|
||||||
to copy, free of charge and under the terms of this License, through a
|
|
||||||
publicly available network server or other readily accessible means,
|
|
||||||
then you must either (1) cause the Corresponding Source to be so
|
|
||||||
available, or (2) arrange to deprive yourself of the benefit of the
|
|
||||||
patent license for this particular work, or (3) arrange, in a manner
|
|
||||||
consistent with the requirements of this License, to extend the patent
|
|
||||||
license to downstream recipients. "Knowingly relying" means you have
|
|
||||||
actual knowledge that, but for the patent license, your conveying the
|
|
||||||
covered work in a country, or your recipient's use of the covered work
|
|
||||||
in a country, would infringe one or more identifiable patents in that
|
|
||||||
country that you have reason to believe are valid.
|
|
||||||
|
|
||||||
If, pursuant to or in connection with a single transaction or
|
|
||||||
arrangement, you convey, or propagate by procuring conveyance of, a
|
|
||||||
covered work, and grant a patent license to some of the parties
|
|
||||||
receiving the covered work authorizing them to use, propagate, modify
|
|
||||||
or convey a specific copy of the covered work, then the patent license
|
|
||||||
you grant is automatically extended to all recipients of the covered
|
|
||||||
work and works based on it.
|
|
||||||
|
|
||||||
A patent license is "discriminatory" if it does not include within
|
|
||||||
the scope of its coverage, prohibits the exercise of, or is
|
|
||||||
conditioned on the non-exercise of one or more of the rights that are
|
|
||||||
specifically granted under this License. You may not convey a covered
|
|
||||||
work if you are a party to an arrangement with a third party that is
|
|
||||||
in the business of distributing software, under which you make payment
|
|
||||||
to the third party based on the extent of your activity of conveying
|
|
||||||
the work, and under which the third party grants, to any of the
|
|
||||||
parties who would receive the covered work from you, a discriminatory
|
|
||||||
patent license (a) in connection with copies of the covered work
|
|
||||||
conveyed by you (or copies made from those copies), or (b) primarily
|
|
||||||
for and in connection with specific products or compilations that
|
|
||||||
contain the covered work, unless you entered into that arrangement,
|
|
||||||
or that patent license was granted, prior to 28 March 2007.
|
|
||||||
|
|
||||||
Nothing in this License shall be construed as excluding or limiting
|
|
||||||
any implied license or other defenses to infringement that may
|
|
||||||
otherwise be available to you under applicable patent law.
|
|
||||||
|
|
||||||
12. No Surrender of Others' Freedom.
|
|
||||||
|
|
||||||
If conditions are imposed on you (whether by court order, agreement or
|
|
||||||
otherwise) that contradict the conditions of this License, they do not
|
|
||||||
excuse you from the conditions of this License. If you cannot convey a
|
|
||||||
covered work so as to satisfy simultaneously your obligations under this
|
|
||||||
License and any other pertinent obligations, then as a consequence you may
|
|
||||||
not convey it at all. For example, if you agree to terms that obligate you
|
|
||||||
to collect a royalty for further conveying from those to whom you convey
|
|
||||||
the Program, the only way you could satisfy both those terms and this
|
|
||||||
License would be to refrain entirely from conveying the Program.
|
|
||||||
|
|
||||||
13. Use with the GNU Affero General Public License.
|
|
||||||
|
|
||||||
Notwithstanding any other provision of this License, you have
|
|
||||||
permission to link or combine any covered work with a work licensed
|
|
||||||
under version 3 of the GNU Affero General Public License into a single
|
|
||||||
combined work, and to convey the resulting work. The terms of this
|
|
||||||
License will continue to apply to the part which is the covered work,
|
|
||||||
but the special requirements of the GNU Affero General Public License,
|
|
||||||
section 13, concerning interaction through a network will apply to the
|
|
||||||
combination as such.
|
|
||||||
|
|
||||||
14. Revised Versions of this License.
|
|
||||||
|
|
||||||
The Free Software Foundation may publish revised and/or new versions of
|
|
||||||
the GNU General Public License from time to time. Such new versions will
|
|
||||||
be similar in spirit to the present version, but may differ in detail to
|
|
||||||
address new problems or concerns.
|
|
||||||
|
|
||||||
Each version is given a distinguishing version number. If the
|
|
||||||
Program specifies that a certain numbered version of the GNU General
|
|
||||||
Public License "or any later version" applies to it, you have the
|
|
||||||
option of following the terms and conditions either of that numbered
|
|
||||||
version or of any later version published by the Free Software
|
|
||||||
Foundation. If the Program does not specify a version number of the
|
|
||||||
GNU General Public License, you may choose any version ever published
|
|
||||||
by the Free Software Foundation.
|
|
||||||
|
|
||||||
If the Program specifies that a proxy can decide which future
|
|
||||||
versions of the GNU General Public License can be used, that proxy's
|
|
||||||
public statement of acceptance of a version permanently authorizes you
|
|
||||||
to choose that version for the Program.
|
|
||||||
|
|
||||||
Later license versions may give you additional or different
|
|
||||||
permissions. However, no additional obligations are imposed on any
|
|
||||||
author or copyright holder as a result of your choosing to follow a
|
|
||||||
later version.
|
|
||||||
|
|
||||||
15. Disclaimer of Warranty.
|
|
||||||
|
|
||||||
THERE IS NO WARRANTY FOR THE PROGRAM, TO THE EXTENT PERMITTED BY
|
|
||||||
APPLICABLE LAW. EXCEPT WHEN OTHERWISE STATED IN WRITING THE COPYRIGHT
|
|
||||||
HOLDERS AND/OR OTHER PARTIES PROVIDE THE PROGRAM "AS IS" WITHOUT WARRANTY
|
|
||||||
OF ANY KIND, EITHER EXPRESSED OR IMPLIED, INCLUDING, BUT NOT LIMITED TO,
|
|
||||||
THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR
|
|
||||||
PURPOSE. THE ENTIRE RISK AS TO THE QUALITY AND PERFORMANCE OF THE PROGRAM
|
|
||||||
IS WITH YOU. SHOULD THE PROGRAM PROVE DEFECTIVE, YOU ASSUME THE COST OF
|
|
||||||
ALL NECESSARY SERVICING, REPAIR OR CORRECTION.
|
|
||||||
|
|
||||||
16. Limitation of Liability.
|
|
||||||
|
|
||||||
IN NO EVENT UNLESS REQUIRED BY APPLICABLE LAW OR AGREED TO IN WRITING
|
|
||||||
WILL ANY COPYRIGHT HOLDER, OR ANY OTHER PARTY WHO MODIFIES AND/OR CONVEYS
|
|
||||||
THE PROGRAM AS PERMITTED ABOVE, BE LIABLE TO YOU FOR DAMAGES, INCLUDING ANY
|
|
||||||
GENERAL, SPECIAL, INCIDENTAL OR CONSEQUENTIAL DAMAGES ARISING OUT OF THE
|
|
||||||
USE OR INABILITY TO USE THE PROGRAM (INCLUDING BUT NOT LIMITED TO LOSS OF
|
|
||||||
DATA OR DATA BEING RENDERED INACCURATE OR LOSSES SUSTAINED BY YOU OR THIRD
|
|
||||||
PARTIES OR A FAILURE OF THE PROGRAM TO OPERATE WITH ANY OTHER PROGRAMS),
|
|
||||||
EVEN IF SUCH HOLDER OR OTHER PARTY HAS BEEN ADVISED OF THE POSSIBILITY OF
|
|
||||||
SUCH DAMAGES.
|
|
||||||
|
|
||||||
17. Interpretation of Sections 15 and 16.
|
|
||||||
|
|
||||||
If the disclaimer of warranty and limitation of liability provided
|
|
||||||
above cannot be given local legal effect according to their terms,
|
|
||||||
reviewing courts shall apply local law that most closely approximates
|
|
||||||
an absolute waiver of all civil liability in connection with the
|
|
||||||
Program, unless a warranty or assumption of liability accompanies a
|
|
||||||
copy of the Program in return for a fee.
|
|
||||||
|
|
||||||
END OF TERMS AND CONDITIONS
|
|
||||||
|
|
||||||
How to Apply These Terms to Your New Programs
|
|
||||||
|
|
||||||
If you develop a new program, and you want it to be of the greatest
|
|
||||||
possible use to the public, the best way to achieve this is to make it
|
|
||||||
free software which everyone can redistribute and change under these terms.
|
|
||||||
|
|
||||||
To do so, attach the following notices to the program. It is safest
|
|
||||||
to attach them to the start of each source file to most effectively
|
|
||||||
state the exclusion of warranty; and each file should have at least
|
|
||||||
the "copyright" line and a pointer to where the full notice is found.
|
|
||||||
|
|
||||||
<one line to give the program's name and a brief idea of what it does.>
|
|
||||||
Copyright (C) <year> <name of author>
|
|
||||||
|
|
||||||
This program is free software: you can redistribute it and/or modify
|
|
||||||
it under the terms of the GNU General Public License as published by
|
|
||||||
the Free Software Foundation, either version 3 of the License, or
|
|
||||||
(at your option) any later version.
|
|
||||||
|
|
||||||
This program is distributed in the hope that it will be useful,
|
|
||||||
but WITHOUT ANY WARRANTY; without even the implied warranty of
|
|
||||||
MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
|
|
||||||
GNU General Public License for more details.
|
|
||||||
|
|
||||||
You should have received a copy of the GNU General Public License
|
|
||||||
along with this program. If not, see <https://www.gnu.org/licenses/>.
|
|
||||||
|
|
||||||
Also add information on how to contact you by electronic and paper mail.
|
|
||||||
|
|
||||||
If the program does terminal interaction, make it output a short
|
|
||||||
notice like this when it starts in an interactive mode:
|
|
||||||
|
|
||||||
<program> Copyright (C) <year> <name of author>
|
|
||||||
This program comes with ABSOLUTELY NO WARRANTY; for details type `show w'.
|
|
||||||
This is free software, and you are welcome to redistribute it
|
|
||||||
under certain conditions; type `show c' for details.
|
|
||||||
|
|
||||||
The hypothetical commands `show w' and `show c' should show the appropriate
|
|
||||||
parts of the General Public License. Of course, your program's commands
|
|
||||||
might be different; for a GUI interface, you would use an "about box".
|
|
||||||
|
|
||||||
You should also get your employer (if you work as a programmer) or school,
|
|
||||||
if any, to sign a "copyright disclaimer" for the program, if necessary.
|
|
||||||
For more information on this, and how to apply and follow the GNU GPL, see
|
|
||||||
<https://www.gnu.org/licenses/>.
|
|
||||||
|
|
||||||
The GNU General Public License does not permit incorporating your program
|
|
||||||
into proprietary programs. If your program is a subroutine library, you
|
|
||||||
may consider it more useful to permit linking proprietary applications with
|
|
||||||
the library. If this is what you want to do, use the GNU Lesser General
|
|
||||||
Public License instead of this License. But first, please read
|
|
||||||
<https://www.gnu.org/licenses/why-not-lgpl.html>.
|
|
||||||
|
|||||||
@@ -1,286 +1,271 @@
|
|||||||

|
<div align="center">
|
||||||
|
|
||||||
<div style="display: flex; flex-direction: column; align-items: center; justify-content: center; text-align: center; font-size: 16px; font-weight: bold; margin-top: 50px;">
|
<img src="docs/images/logo.png" width="auto" alt="Logo">
|
||||||
|
<p>
|
||||||
<div>
|
<strong>A lightweight Transformer training & inference framework</strong>
|
||||||
<a href="#english" style="text-decoration: none; margin: 0 10px; color: blue;">English</a> |
|
</p>
|
||||||
<a href="#chinese" style="text-decoration: none; margin: 0 10px; color: blue;">中文</a>
|
|
||||||
</div>
|
|
||||||
|
|
||||||
<h1 style="margin: 20px 0 0 0; font-size: 2.5em; font-weight: bold;">KHAOSZ </h1>
|
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<h2 id="english">English Version</h2>
|
<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-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>
|
||||||
|
|
||||||
A training and inference framework for autoregressive Transformer language models.
|
<div align="center">
|
||||||
|
<a href="#english">English</a> •
|
||||||
|
<a href="docs/README-zh-CN.md">中文</a> •
|
||||||
|
<a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> •
|
||||||
|
<a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> •
|
||||||
|
<a href="https://huggingface.co/ViperEkura">HuggingFace</a>
|
||||||
|
</div>
|
||||||
|
|
||||||
**Model Download Options (choose one):**
|
<br>
|
||||||
|
|
||||||
1. Visit [HuggingFace](https://huggingface.co/ViperEk/KHAOSZ) and check **Files and versions**
|
## 📖 Table of Contents
|
||||||
2. Run `scripts/download.py` to download model parameters
|
|
||||||
|
|
||||||
**Demo Video:** [bilibili](https://www.bilibili.com/video/BV1z5RPYHEkd)
|
- [Overview](#overview)
|
||||||
|
- [Getting Started](#getting-started)
|
||||||
|
- [Demo](#demo)
|
||||||
|
- [Documentation](#documentation)
|
||||||
|
- [Contributing](#contributing)
|
||||||
|
- [Community](#community)
|
||||||
|
- [License](#license)
|
||||||
|
|
||||||
For training data sources, please refer to the **Model Card** section on the HuggingFace download page.
|
---
|
||||||
|
|
||||||
**License:** The code follows the GPL-3.0 license. Please provide attribution when using it.
|
<a id="english"></a>
|
||||||
|
## English
|
||||||
|
|
||||||
- **📊 Device Selection:** Uses CUDA for training by default
|
### Overview
|
||||||
- **🌐 Performance Optimization:** Enable `dtype=torch.bfloat16` to accelerate training and reduce memory usage. Ensure your hardware supports this feature
|
|
||||||
- **🤖 Language Support:** The model supports training in Chinese and English. Since the BBPE tokenizer hasn't been trained on multilingual text, OOV (Out-of-Vocabulary) issues are minimal for Chinese and English, but may exist for other languages
|
|
||||||
|
|
||||||
|
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.
|
||||||
|
|
||||||
### 📌 Training Guide
|
| 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, ROUGE, and weight-analysis evaluation tools |
|
||||||
|
| **Extensibility** | Factory and registry architecture for models, datasets, training strategies, callbacks, kernels, and protocol components |
|
||||||
|
|
||||||
To train this Transformer model, follow these steps:
|
### Getting Started
|
||||||
|
|
||||||
**(1). Prepare the Dataset:**
|
End-to-end walkthrough in 5 steps:
|
||||||
|
|
||||||
Place the dataset in the specified root directory. This system uses the BBPE tokenizer for tokenization and requires training with pre-tokenized segments (stored as *.h5 format files).
|
**1. Install**
|
||||||
|
|
||||||
**(2). Install Dependencies:**
|
AstrAI requires Python 3.12+ and pins PyTorch exactly to `2.11.0`. Training, `scripts/tools/generate.py`, generation evaluations, and the generation demos require CUDA; CPU support is limited to components with an explicit CPU device path, such as the HTTP server and direct-scoring evaluations.
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
pip install -e .
|
git clone https://github.com/ViperEkura/AstrAI.git
|
||||||
|
cd AstrAI
|
||||||
|
pip install -e . # kernels auto-build when nvcc + CUDA are detected
|
||||||
|
# CSRC_KERNELS=false pip install -e . # skip kernels (pure PyTorch)
|
||||||
|
# CSRC_KERNELS=true pip install -e . --no-build-isolation # force the fused CUDA kernel build
|
||||||
|
# pip install -e ".[dev]" # dev dependencies (pytest, ruff)
|
||||||
```
|
```
|
||||||
|
|
||||||
**(3). Run the Training Script:**
|
**2. Download model**
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python train.py \
|
python scripts/demo/download.py # downloads 1B checkpoint to params/
|
||||||
--train_type=train_type[seq, sft, dpo] \
|
|
||||||
--data_root_path=/path/to/dataset \
|
|
||||||
--param_path=/path/to/param_path \
|
|
||||||
--n_epoch=5 \
|
|
||||||
--batch_size=8 \
|
|
||||||
--max_lr=2e-4 \
|
|
||||||
--checkpoint_interval=10000 \
|
|
||||||
--checkpoint_dir=checkpoints
|
|
||||||
```
|
```
|
||||||
|
|
||||||
**Parameter Explanation:**
|
**3. Preprocess data**
|
||||||
- `--train_type`: Training type (seq, sft, dpo)
|
|
||||||
- `--data_root_path`: Dataset root directory
|
|
||||||
- `--param_path`: Path to model training parameters
|
|
||||||
- `--n_epoch`: Total number of training epochs
|
|
||||||
- `--batch_size`: Batch size
|
|
||||||
- `--accumulation_steps`: Number of batches per training step
|
|
||||||
- `--warmup_steps`: Warmup steps
|
|
||||||
- `--max_lr`: Maximum learning rate (using warmup + cosine decay)
|
|
||||||
- `--checkpoint_interval`: Checkpoint saving interval
|
|
||||||
- `--checkpoint_dir`: Checkpoint saving directory
|
|
||||||
- `--resume_dir`: Resume training from specified path
|
|
||||||
|
|
||||||
|
Create `pretrain.json` (preprocessing config for `seq` strategy):
|
||||||
|
|
||||||
|
```json
|
||||||
### 👉 Usage Guide
|
{
|
||||||
|
"version": 1,
|
||||||
**(1). Chat with the Model:**
|
"input": {"sections": [{"field": "text", "action": "train"}]},
|
||||||
|
"preprocessing": {"max_seq_len": 2048},
|
||||||
Open `chat.py` or use the streaming/non-streaming interfaces:
|
"output": {"storage_format": "bin"}
|
||||||
|
}
|
||||||
**Streaming Output:**
|
|
||||||
```python
|
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
model_dir = "your_model_parameter_dir"
|
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
history = []
|
|
||||||
|
|
||||||
while True:
|
|
||||||
query = input(">> ")
|
|
||||||
if query == "!exit":
|
|
||||||
break
|
|
||||||
|
|
||||||
response_size = 0
|
|
||||||
for response, history in model.stream_generate(
|
|
||||||
query=query,
|
|
||||||
history=history,
|
|
||||||
temperature=0.85,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50
|
|
||||||
):
|
|
||||||
print(response[response_size:], end="")
|
|
||||||
response_size = len(response)
|
|
||||||
```
|
```
|
||||||
|
|
||||||
**Non-streaming Output:**
|
|
||||||
```python
|
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
model_dir = "your_model_parameter_dir"
|
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
history = []
|
|
||||||
|
|
||||||
while True:
|
|
||||||
query = input(">> ")
|
|
||||||
if query == "!exit":
|
|
||||||
break
|
|
||||||
|
|
||||||
response = model.generate(
|
|
||||||
query=query,
|
|
||||||
history=history,
|
|
||||||
temperature=0.85,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50
|
|
||||||
)
|
|
||||||
print(response)
|
|
||||||
```
|
|
||||||
|
|
||||||
**(2). Retrieval-Augmented Generation (RAG):**
|
|
||||||
|
|
||||||
```python
|
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
model_dir = "your_model_parameter_dir"
|
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
|
|
||||||
retrieved_content = model.retrieve_generate(
|
|
||||||
query=query,
|
|
||||||
retrieve_top_k=5,
|
|
||||||
temperature=0.6,
|
|
||||||
top_k=30,
|
|
||||||
top_p=0.95
|
|
||||||
)
|
|
||||||
print(retrieved_content)
|
|
||||||
```
|
|
||||||
|
|
||||||
<h2 id="chinese">中文版本</h2>
|
|
||||||
这是一个支持基于自回归模式的 Transfomer 语言模型训练以及推理框架
|
|
||||||
|
|
||||||
**模型下载选项(任选其一):**
|
|
||||||
|
|
||||||
1. 访问 [HuggingFace](https://huggingface.co/ViperEk/KHAOSZ) 查看 **Files and versions**
|
|
||||||
2. 运行 `scripts/download.py` 下载模型参数
|
|
||||||
|
|
||||||
**演示视频:** [bilibili](https://www.bilibili.com/video/BV1z5RPYHEkd)
|
|
||||||
|
|
||||||
训练数据来源请参见 HuggingFace 下载页面中的 **Model Card** 部分。
|
|
||||||
|
|
||||||
**许可证:** 代码遵循 GPL-3.0 协议,使用时请注明出处。
|
|
||||||
|
|
||||||
- **📊 设备选择:** 默认使用 CUDA 进行训练
|
|
||||||
- **🌐 性能优化:** 启用 `dtype=torch.bfloat16` 以加速训练并减少内存占用,请确保硬件支持该特性
|
|
||||||
- **🤖 语言支持:** 模型支持中文和英文训练。由于 BBPE 分词器未使用多语言文本训练,因此中英文的 OOV(未登录词)问题较少,其他语言可能存在 OOV 问题
|
|
||||||
|
|
||||||
|
|
||||||
### 📌 训练指南
|
|
||||||
|
|
||||||
要训练该 Transformer 模型,请按照以下步骤操作:
|
|
||||||
|
|
||||||
**(1). 准备数据集:**
|
|
||||||
|
|
||||||
将数据集放置在指定的根目录下, 本系统采用 BBPE 分词器进行分词,并且要求使用已经经过分词的 token 分段训练(分段存储为 *.h5 格式)
|
|
||||||
|
|
||||||
**(2). 安装依赖:**
|
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
pip install -e .
|
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c pretrain.json
|
||||||
```
|
```
|
||||||
|
|
||||||
**(3). 运行训练脚本:**
|
**4. Train**
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python train.py \
|
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||||
--train_type=train_type[seq, sft, dpo] \
|
|
||||||
--data_root_path=/path/to/dataset \
|
nohup python scripts/tools/train.py \
|
||||||
--param_path=/path/to/param_path \
|
--nprocs=4 \
|
||||||
--n_epoch=5 \
|
--parallel_mode=ddp \
|
||||||
--batch_size=8 \
|
--train_type=seq \
|
||||||
--max_lr=2e-4 \
|
--data_root_path=/path/to/dataset \
|
||||||
--checkpoint_interval=10000 \
|
--param_path=/path/to/model \
|
||||||
--checkpoint_dir=checkpoints
|
--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**
|
||||||
- `--train_type`: 训练类型(seq, sft, dpo)
|
|
||||||
- `--data_root_path`: 数据集根目录
|
|
||||||
- `--param_path`: 模型训练参数路径
|
|
||||||
- `--n_epoch`: 总训练轮数
|
|
||||||
- `--batch_size`: 批量大小
|
|
||||||
- `--accumulation_steps`: 每个训练步骤的 batch 数量
|
|
||||||
- `--warmup_steps`: 预热步数(warmup steps)
|
|
||||||
- `--max_lr`: 最大学习率(使用预热 + 余弦衰减)
|
|
||||||
- `--checkpoint_interval`: 检查点保存间隔
|
|
||||||
- `--checkpoint_dir`: 检查点保存目录
|
|
||||||
- `--resume_dir`: 从指定路径恢复训练
|
|
||||||
|
|
||||||
|
```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" \
|
||||||
**(1). 与模型对话:**
|
-d '{"messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||||
|
|
||||||
打开 `chat.py` 或使用流式/非流式接口:
|
|
||||||
|
|
||||||
**流式输出:**
|
|
||||||
```python
|
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
model_dir = "your_model_parameter_dir"
|
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
history = []
|
|
||||||
|
|
||||||
while True:
|
|
||||||
query = input(">> ")
|
|
||||||
if query == "!exit":
|
|
||||||
break
|
|
||||||
|
|
||||||
response_size = 0
|
|
||||||
for response, history in model.stream_generate(
|
|
||||||
query=query,
|
|
||||||
history=history,
|
|
||||||
temperature=0.85,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50
|
|
||||||
):
|
|
||||||
print(response[response_size:], end="")
|
|
||||||
response_size = len(response)
|
|
||||||
```
|
```
|
||||||
|
|
||||||
**非流式输出:**
|
### Demo
|
||||||
```python
|
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
model_dir = "your_model_parameter_dir"
|
Check out the demos in the `scripts/demo/` folder:
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
history = []
|
|
||||||
|
|
||||||
while True:
|
```bash
|
||||||
query = input(">> ")
|
# Download model weights (required before running demos)
|
||||||
if query == "!exit":
|
python scripts/demo/download.py # model → params/
|
||||||
break
|
|
||||||
|
|
||||||
response = model.generate(
|
# Single-turn interactive streaming prompt loop (no conversation history)
|
||||||
query=query,
|
python scripts/demo/stream_chat.py
|
||||||
history=history,
|
# Type your message after >>, type !exit to quit
|
||||||
temperature=0.85,
|
|
||||||
top_p=0.95,
|
# Batch generation (5 hardcoded prompts, non-streaming)
|
||||||
top_k=50
|
python scripts/demo/generate_batch.py
|
||||||
)
|
|
||||||
print(response)
|
# Single-prompt autoregressive streaming
|
||||||
|
python scripts/demo/generate_ar.py
|
||||||
```
|
```
|
||||||
|
|
||||||
**(2). 基于检索的生成(RAG):**
|
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).
|
||||||
|
|
||||||
```python
|
Watch a video walkthrough on [bilibili](https://www.bilibili.com/video/BV1fuLB6yEj6).
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
model_dir = "your_model_parameter_dir"
|
---
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
|
|
||||||
retrieved_content = model.retrieve_generate(
|
See [Documentation](#documentation) for full references beyond the examples above.
|
||||||
query=query,
|
|
||||||
retrieve_top_k=5,
|
#### Text Generation
|
||||||
temperature=0.6,
|
|
||||||
top_k=30,
|
Batch generation from a JSONL file:
|
||||||
top_p=0.95
|
|
||||||
)
|
```bash
|
||||||
print(retrieved_content)
|
python scripts/tools/generate.py \
|
||||||
|
--param_path ./params \
|
||||||
|
--input_json_file input.jsonl \
|
||||||
|
--output_json_file output.jsonl
|
||||||
```
|
```
|
||||||
|
|
||||||
|
#### Docker
|
||||||
|
|
||||||
|
Build and run with Docker (recommended for GPU environments):
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Build image
|
||||||
|
docker build -t astrai:latest .
|
||||||
|
|
||||||
|
# Run with GPU support
|
||||||
|
docker run --gpus all -it astrai:latest
|
||||||
|
|
||||||
|
# Run inference server
|
||||||
|
docker run --gpus all -p 8000:8000 astrai:latest \
|
||||||
|
python -m scripts.tools.server --port 8000 --device cuda
|
||||||
|
|
||||||
|
# Run with volume mount for data
|
||||||
|
docker run --gpus all -v /path/to/data:/data -it astrai:latest
|
||||||
|
|
||||||
|
# Docker Compose (GPU, default)
|
||||||
|
docker compose up -d
|
||||||
|
|
||||||
|
# Docker Compose CPU server profile (CUDA-only generation scripts/demos are unavailable)
|
||||||
|
docker compose --profile cpu up -d
|
||||||
|
|
||||||
|
# YAML-driven serving (see serve.yaml; up/run/down/logs/status...)
|
||||||
|
bash scripts/serve.sh up
|
||||||
|
```
|
||||||
|
|
||||||
|
> **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`.
|
||||||
|
|
||||||
|
#### HTTP API Examples
|
||||||
|
|
||||||
|
Additional request examples beyond the [Getting Started](#getting-started) flow:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 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}'
|
||||||
|
|
||||||
|
# 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}'
|
||||||
|
|
||||||
|
# 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"]}'
|
||||||
|
|
||||||
|
# Health check
|
||||||
|
curl http://localhost:8000/health
|
||||||
|
```
|
||||||
|
|
||||||
|
See [Inference Guide](docs/guides/inference.md) for SSE streaming format, error codes, and stats endpoint.
|
||||||
|
|
||||||
|
### Documentation
|
||||||
|
|
||||||
|
| Document | Description |
|
||||||
|
|----------|-------------|
|
||||||
|
| [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 |
|
||||||
|
| [Docker Serving](./docs/developer/docker-serving.md) | YAML-driven containerized serving (`serve.yaml`, `serve.sh`) |
|
||||||
|
| [Docker Training](./docs/developer/docker-training.md) | YAML-driven containerized training (`train.yaml`, `train.sh`) |
|
||||||
|
|
||||||
|
### Contributing
|
||||||
|
|
||||||
|
We welcome contributions! Please see our [Contributing Guidelines](CONTRIBUTING.md) for details.
|
||||||
|
|
||||||
|
1. Fork the repository.
|
||||||
|
2. Create a feature branch.
|
||||||
|
3. Commit your changes.
|
||||||
|
4. Open a Pull Request.
|
||||||
|
|
||||||
|
For major changes, please open an issue first to discuss what you would like to change.
|
||||||
|
|
||||||
|
### Community
|
||||||
|
|
||||||
|
- **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/ViperEkura)
|
||||||
|
|
||||||
|
### 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>
|
||||||
|
|||||||
@@ -1,220 +0,0 @@
|
|||||||
## 1. 为什么我要做这个项目?
|
|
||||||
|
|
||||||
现在市面上有很多大模型,比如GPT、LLaMA这些,动不动就是几十亿甚至上千亿参数。但说实话,这些模型对硬件要求太高了,普通开发者根本玩不起。我就想:**能不能做一个既好用又能在普通电脑上跑起来的模型呢?** 这其实也是目前大部分人的期望, 能有一个可以本地部署的ai小型项目,实现完全私有化并且有一定的智能能力。
|
|
||||||
|
|
||||||
于是就有了这个KHAOSZ项目,1B参数,中英双语,支持对话、文本生成、RAG检索,而且训练代码都是开源的!
|
|
||||||
|
|
||||||
## 2. 系统架构
|
|
||||||
|
|
||||||
系统分为以下板块
|
|
||||||
|
|
||||||
```mermaid
|
|
||||||
graph LR
|
|
||||||
%% 样式定义
|
|
||||||
classDef config fill:#e1f5fe,stroke:#01579b;
|
|
||||||
classDef trainer fill:#f3e5f5,stroke:#4a148c;
|
|
||||||
classDef data fill:#e8f5e8,stroke:#1b5e20;
|
|
||||||
classDef model fill:#fff3e0,stroke:#e65100;
|
|
||||||
classDef inference fill:#fce4ec,stroke:#880e4f;
|
|
||||||
classDef parallel fill:#e0f2f1,stroke:#004d40;
|
|
||||||
|
|
||||||
%% 配置模块
|
|
||||||
subgraph Config["Config(配置模块)"]
|
|
||||||
C1[model_config.py]
|
|
||||||
C2[train_config.py]
|
|
||||||
C3[scheduler_config.py]
|
|
||||||
end
|
|
||||||
class Config config;
|
|
||||||
|
|
||||||
%% 训练器模块
|
|
||||||
subgraph Trainer["Trainer(训练器模块)"]
|
|
||||||
T1[trainer.py]
|
|
||||||
T2[train_content.py]
|
|
||||||
T3[schedule.py]
|
|
||||||
T4[strategy.py]
|
|
||||||
T5[train_callback.py]
|
|
||||||
end
|
|
||||||
class Trainer trainer;
|
|
||||||
|
|
||||||
%% 数据模块
|
|
||||||
subgraph Data["Data(数据模块)"]
|
|
||||||
D1[dataset.py]
|
|
||||||
D2[sampler.py]
|
|
||||||
D3[mmap.py]
|
|
||||||
D4[tokenizer.py]
|
|
||||||
D5[checkpoint.py]
|
|
||||||
end
|
|
||||||
class Data data;
|
|
||||||
|
|
||||||
%% 模型模块
|
|
||||||
subgraph Model["Model(模型模块)"]
|
|
||||||
M1[transformer.py]
|
|
||||||
M2[module.py]
|
|
||||||
end
|
|
||||||
class Model model;
|
|
||||||
|
|
||||||
%% 推理模块
|
|
||||||
subgraph Inference["Inference(推理模块)"]
|
|
||||||
I1[generator.py]
|
|
||||||
I2[core.py]
|
|
||||||
end
|
|
||||||
class Inference inference;
|
|
||||||
|
|
||||||
%% 并行模块
|
|
||||||
subgraph Parallel["Parallel(并行模块)"]
|
|
||||||
P1[setup.py]
|
|
||||||
P2[module.py]
|
|
||||||
end
|
|
||||||
class Parallel parallel;
|
|
||||||
|
|
||||||
%% 配置依赖
|
|
||||||
C2 -.-> T1
|
|
||||||
C1 -.-> M1
|
|
||||||
C3 -.-> T3
|
|
||||||
|
|
||||||
%% 训练器内部依赖
|
|
||||||
T1 --> T5
|
|
||||||
T1 --> T2
|
|
||||||
T2 --> T3
|
|
||||||
T2 --> T4
|
|
||||||
|
|
||||||
%% 数据流
|
|
||||||
D1 --> D2
|
|
||||||
D1 --> D3
|
|
||||||
D1 --> D4
|
|
||||||
D1 --> D5
|
|
||||||
|
|
||||||
%% 模型依赖
|
|
||||||
M1 --> M2
|
|
||||||
|
|
||||||
%% 推理依赖
|
|
||||||
I1 --> I2
|
|
||||||
|
|
||||||
%% 跨模块依赖
|
|
||||||
T2 -.-> M1
|
|
||||||
I1 -.-> M1
|
|
||||||
T2 -.-> D1
|
|
||||||
T1 -.-> P1
|
|
||||||
```
|
|
||||||
|
|
||||||
|
|
||||||
### 1. 配置管理(/config/)
|
|
||||||
- **模型配置**:定义模型结构参数(如层数、头数、维度等),通过 `ModelConfig` 统一管理。
|
|
||||||
- **训练配置**:设置训练参数(如批次大小、训练阶段 PT/SFT/DPO、优化器等),由 `TrainConfig` 加载。
|
|
||||||
- **调度配置**:控制学习率策略(如余弦退火)和训练进度。
|
|
||||||
|
|
||||||
### 2. 硬件与并行(/parallel/)
|
|
||||||
- **分布式初始化**:通过 `setup_parallel` 函数,根据配置初始化多卡/多机训练环境。
|
|
||||||
|
|
||||||
### 3. 数据处理(/data/)
|
|
||||||
- **高效加载**:使用内存映射(mmap)技术加载超大语料,避免内存溢出,实现零拷贝读取。
|
|
||||||
|
|
||||||
### 4. 模型与训练(/model/, /trainer/)
|
|
||||||
- **统一模型架构**:基于 Transformer,支持灵活配置不同规模(如7B、13B)。
|
|
||||||
- **策略化训练器**:`Trainer` 根据训练阶段(PT/SFT/DPO)自动切换训练策略,复用同一训练循环。
|
|
||||||
- **训练上下文管理**:统一管理模型、优化器、调度器和指标,支持多阶段无缝衔接。
|
|
||||||
|
|
||||||
### 5. 推理服务(/inference/, /utils/)
|
|
||||||
- **统一生成接口**:提供同步、批量、流式生成方法,适配所有训练阶段。
|
|
||||||
- **KV缓存优化**:在自回归生成中缓存 Key/Value,昇腾XPU下利用高速片上内存加速。
|
|
||||||
- **RAG支持**:结合检索器和嵌入模型,从外部知识库注入相关信息,提升回答质量。
|
|
||||||
- **智能文本分割**:
|
|
||||||
- **结构优先分割**:按标题、段落等切分;
|
|
||||||
- **语义分割**:基于句子嵌入相似度,确保片段语义完整,提升微调效果。
|
|
||||||
|
|
||||||
|
|
||||||
## 3. 训练流程
|
|
||||||
|
|
||||||
常见大语言模型(Large Language Model, LLM)的训练流程通常包含三个阶段:**预训练(Pre-training, PT)**、**监督微调(Supervised Fine-Tuning, SFT)** 以及 **基于人类反馈的强化学习(Reinforcement Learning from Human Feedback, RLHF)**。本系统设计支持全流程无缝衔接,通过模块化策略实现不同训练阶段的高效切换与状态管理,确保模型能力从通用语言理解逐步对齐至符合人类偏好的对话与指令执行。
|
|
||||||
|
|
||||||
### **2.1 预训练阶段**
|
|
||||||
|
|
||||||
预训练阶段旨在构建模型的基础语言能力与通用知识表示。该阶段在大规模、无标注的语料库(通常涵盖数百GB至数TB的文本数据)上进行自监督学习。模型架构基于标准的Transformer Decoder,通过掩码语言建模(如因果语言建模)目标进行训练,使模型能够学习词汇、语法、语义及蕴含于文本中的世界知识。
|
|
||||||
|
|
||||||
**核心公式:因果语言建模(Causal Language Modeling)**
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{PT}} = - \sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta)
|
|
||||||
$$
|
|
||||||
|
|
||||||
**符号说明:**
|
|
||||||
|
|
||||||
- $T$:序列长度
|
|
||||||
- $x_t$:序列中第 $ t $ 个词元(token)
|
|
||||||
- $x_{<t}$:位置 $ t $ 之前的所有词元
|
|
||||||
- $\theta$:模型参数
|
|
||||||
- $P(x_t \mid x_{<t}; \theta)$:模型在给定上文条件下预测下一个词元的概率
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
本阶段的核心在于利用分布式的并行计算资源,实现模型参数的稳定优化。训练器模块中的`PTStrategy`策略,专门负责管理预训练特有的数据采样、长序列分段与梯度累积逻辑。同时,硬件适配模块会根据运行环境(如华为昇腾NPU集群或标准GPU集群)自动选择最优的并行通信后端(如HCCL或NCCL),并进行计算图优化,以最大化硬件利用率和训练吞吐量。
|
|
||||||
|
|
||||||
另外系统通过数据模块中的高效内存映射加载器(`MmapFileHandler`),实现海量数据的零拷贝读取,以克服传统IO瓶颈。
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
### **2.2 监督微调阶段**
|
|
||||||
|
|
||||||
预训练模型虽具备强大的语言生成能力,但尚未对齐至遵循人类指令、进行安全有益对话的行为模式。监督微调阶段旨在弥合这一差距。该阶段使用由人工精心编写的、高质量的“指令-响应”配对数据集。
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
**核心公式:序列到序列条件语言建模**
|
|
||||||
|
|
||||||
设完整序列 $S = [s_1, s_2, \ldots, s_{P+L}]$,其中:
|
|
||||||
|
|
||||||
- 前 $P$ 个token是prompt 以及对应控制token: $X = [s_1, \ldots, s_P]$
|
|
||||||
- 后 $L$ 个token是response以及对应控制token: $Y = [s_{P+1}, \ldots, s_{P+L}]$
|
|
||||||
|
|
||||||
损失函数为:
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{SFT}} = - \sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
|
|
||||||
$$
|
|
||||||
|
|
||||||
训练器模块将动态切换到`SFTStrategy`策略。此策略的核心是引入序列级的监督学习目标,例如预测给定指令下完整、正确的响应序列。训练上下文管理器(`TrainContext`)负责平滑地从PT阶段检查点加载模型状态,并初始化新的优化器和学习率调度器。本阶段不仅优化模型参数,更重要的是引导模型学习“对话”这一特定任务范式,使其输出风格、内容与格式均符合人类期望。
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
### **2.3 基于人类反馈的强化学习阶段**
|
|
||||||
|
|
||||||
为生成更具帮助性、无害性且符合人类偏好的高质量输出,系统进一步集成强化学习阶段。传统的RLHF流程包括**奖励模型训练**与**策略模型微调**两个核心步骤。系统支持以直接偏好优化(Direct Preference Optimization,DPO)算法为代表的策略微调,并针对稳定性与收敛性进行了多项工程优化。
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#### **2.3.1 传统 RLHF(奖励模型训练)**
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{RM}} = -\mathbb{E}_{(x, y_w, y_l) \sim D} \left[ \log \sigma\left( r_\phi(x, y_w) - r_\phi(x, y_l) \right) \right]
|
|
||||||
$$
|
|
||||||
|
|
||||||
**符号说明:**
|
|
||||||
|
|
||||||
- $r_\phi(x, y)$:参数为 $phi$ 的奖励模型给出的标量分数
|
|
||||||
- $y_w, y_l $:同一提示 $ x $ 下的优选和劣选回答
|
|
||||||
- $\sigma $:sigmoid 函数
|
|
||||||
- $\mathcal{D} $:人类偏好数据集
|
|
||||||
|
|
||||||
|
|
||||||
#### **2.3.2 DPO 直接偏好优化**(推荐)
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{DPO}} = -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]
|
|
||||||
$$
|
|
||||||
|
|
||||||
**符号说明:**
|
|
||||||
|
|
||||||
- $\pi_\theta(y \mid x) $:当前策略模型生成回答的概率
|
|
||||||
- $\pi_{\text{ref}}(y \mid x) $:参考模型生成回答的概率
|
|
||||||
- $\beta $:温度参数(通常设为 0.1-0.5)
|
|
||||||
- 注意:隐式学习奖励函数 $r(x, y) = \beta \log \frac{\pi_\theta(y \mid x)}{\pi_{\text{ref}}(y \mid x)} $
|
|
||||||
|
|
||||||
|
|
||||||
在本阶段,训练器模块启用`RLHFStrategy`策略(或类似的`DPOStrategy`直接偏好优化策略)。该策略管理一个复杂的训练循环,其中包含策略模型(待优化的LLM)、参考模型(通常为SFT后的模型快照)和奖励模型。系统流程如下:
|
|
||||||
|
|
||||||
1. **偏好数据收集与奖励建模**:首先,通过收集人类标注员对同一提示词下多个模型生成结果的排序偏好数据,训练一个独立的奖励模型(Reward Model, RM)。该模型学习为生成文本输出一个标量奖励分数,以量化其符合人类偏好的程度。
|
|
||||||
2. **策略优化**:随后,使用奖励模型作为优化信号,通过强化学习算法对SFT模型(作为策略)进行微调。策略优化的目标是最大化从奖励模型获得的期望累计奖励,同时通过KL散度惩罚项约束策略模型与参考模型的输出分布不过度偏离,以防止模式崩溃并保持生成多样性。训练上下文管理器在此阶段同时维护策略模型、参考模型和奖励模型(或价值函数模型)的状态,并协调复杂的多阶段梯度计算。
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
通过上述三阶段的递进式训练,模型完成了从通用语言基座到专业化、高对齐度对话智能体的进化。系统通过统一的`Trainer`接口和策略模式设计,使得各阶段训练在代码层面高度复用,在流程层面清晰解耦,为大规模语言模型的研发与迭代提供了高效、灵活且可扩展的工程基础。
|
|
||||||
@@ -1,89 +0,0 @@
|
|||||||
## 模型介绍
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
### 1. 模型搭建
|
|
||||||
|
|
||||||
本模型采用Transformer架构, 使用GQA(q_head=24, kv_head=4) 机制,相较于传统的MHA可以节省KV cache 的显存占用(但是目前没有做KV cache),通过堆叠24层Transformer实现模型的搭建, 参数量为1.0b。Transformer 是自回归模型, 是通过计算前面所有的token的关系得到下一个token的概率分布
|
|
||||||
|
|
||||||

|
|
||||||
|
|
||||||
什么是自回归模型呢, 在把句子拆分成token之后, 模型会预测下一个token的概率分布。这意味着模型会根据给定的上下文(即已经出现的tokens序列),计算出下一个可能的token及其对应的概率。
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#### 1. 自回归
|
|
||||||
|
|
||||||
假设我们有一个句子被拆分成如下tokens列表:
|
|
||||||
|
|
||||||
```
|
|
||||||
["你好", "," "今天", "天气"]
|
|
||||||
```
|
|
||||||
|
|
||||||
接下来,模型会基于这个序列预测下一个可能出现的token。这通常以概率分布的形式给出,比如:
|
|
||||||
|
|
||||||
```
|
|
||||||
-> {"token": "不错", "probability": 0.4}
|
|
||||||
-> {"token": "晴朗", "probability": 0.2}
|
|
||||||
-> ......
|
|
||||||
```
|
|
||||||
|
|
||||||
这里,“不错”和“晴朗”是两个可能跟随在“天气”之后的tokens,并且给出了每个token成为下一个token的可能性大小。
|
|
||||||
|
|
||||||
之后,我们通过采样(通过top_k, top_p, temperature参数调整采样后的结果)得到下一个token并且将下一个token加入序列作为输入
|
|
||||||
|
|
||||||
```
|
|
||||||
["你好", "," "今天", "天气", "不错"]
|
|
||||||
```
|
|
||||||
|
|
||||||
之后都是在重复这个流程, 直到遇到控制流程结束的token(<|end_of_seqence|>)模型停止处理(一般模型都会设置控制token, 不然模型会一直输出到显存爆炸)。
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#### 2. 因果掩码
|
|
||||||
|
|
||||||
transformer 中采用注意力机制,输入的形状一般为[bsz, seq_len], 输出为[bsz, seq_len,n_dim], 为了实现预测下一个token, 模型的输入和输出必须错开来一个位置。模型预测的target必须错开一个位置, 在训练的时候我们也采用错开一个位置的方法
|
|
||||||
|
|
||||||
```
|
|
||||||
sequence : [[1, 2, 3, 4, 5, 6]]
|
|
||||||
input_ids: [[1, 2, 3, 4, 5]]
|
|
||||||
target_ids: [[2, 3, 4, 5, 6]]
|
|
||||||
```
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
注意力得分计算的公式为
|
|
||||||
|
|
||||||
|
|
||||||
$$ s_{ij} = softmax(\frac{q_i^Tk_j}{\sqrt{d_k}}) $$
|
|
||||||
$$ s_{ij} := s_{ij} + mask_{ij} $$
|
|
||||||
|
|
||||||
|
|
||||||
其中注意力得分代表了模型对两个token之间相似程度的关注程度
|
|
||||||
|
|
||||||
对于decoder only结构的模型, 为了防止模型从未来的位置偷到信息, 在注意力的计算过程中需要增加掩码,我们需要在注意力得分计算之前应用一个掩码。这个掩码通常是一个下三角矩阵,对于长度为n的序列,它的形状是[n, n]。下面以一个长度为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]]
|
|
||||||
```
|
|
||||||
|
|
||||||
在这个矩阵中,0表示可以注意到的位置,而-inf表示应该被掩盖(即不应注意到)的位置。因为这个句子保证了注意力得分中 $j > i$ 的部分通过softmax 之后由`inf` 变成0, 也就是模型不能看到未来的信息
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#### 3. 旋转位置编码
|
|
||||||
|
|
||||||
旋转位置编码(Rotary Position Embedding, RoPE)是一种为了解决Transformer模型中缺乏对序列位置信息直接建模的问题而设计的位置编码方法。与传统的位置编码(如正弦和余弦函数的位置编码)不同,RoPE通过将位置信息直接嵌入到查询(Query, Q)和键(Key, K)向量中来实现,使得模型能够更自然地处理序列中的相对位置关系。
|
|
||||||
|
|
||||||
|
|
||||||
$$ 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 $$
|
|
||||||
|
|
||||||
其中的 $R_{i-j}$ 控制了模型的不同token 在不同相对距离上注意力的衰减,在 $i - j$ 绝对值越大的时候, 衰减的程度越强, 通过这种方式能让模型学习到相对位置关系, 从而使得模型可以扩展和适应长序列
|
|
||||||
@@ -1,27 +0,0 @@
|
|||||||
## kv_cache 实现
|
|
||||||
|
|
||||||
根据注意力的计算公式
|
|
||||||
|
|
||||||
$$
|
|
||||||
\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*}
|
|
||||||
$$
|
|
||||||
|
|
||||||
由于模型是自回归模型, 我们只用求序列最后一个部分,也就是说 $ i $ 的下标是确定的, 是序列最后一个元素, 我们求的是 $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*}
|
|
||||||
$$
|
|
||||||
|
|
||||||
如果我们把式子展开
|
|
||||||
|
|
||||||
$$
|
|
||||||
o_n = \sum_j \text{softmax}\left(\frac{q_n k_{j}}{\sqrt{d_k}}\right)v_{j}
|
|
||||||
$$
|
|
||||||
|
|
||||||
以上表达式只有k和v存在长度下标, 而 $q$ 没有, 所以计算过程中 $q$ 的输入是确定的上次输入的最后一个token, 而 $k, v$ 是需要对不同长度的部分进行缓存的,同时缓存的时候应该注意位置编码的计算应该在kvcache的计算之前进行,否则会存在位置编码的计算错误
|
|
||||||
Binary file not shown.
|
Before Width: | Height: | Size: 21 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 11 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 590 KiB |
@@ -0,0 +1,95 @@
|
|||||||
|
__version__ = "1.3.13"
|
||||||
|
__author__ = "ViperEkura"
|
||||||
|
|
||||||
|
from astrai.config import (
|
||||||
|
AutoRegressiveLMConfig,
|
||||||
|
BaseModelConfig,
|
||||||
|
ConfigFactory,
|
||||||
|
EncoderConfig,
|
||||||
|
PipelineConfig,
|
||||||
|
TrainConfig,
|
||||||
|
)
|
||||||
|
from astrai.dataset import (
|
||||||
|
BaseDataset,
|
||||||
|
DatasetFactory,
|
||||||
|
RDSampler,
|
||||||
|
Store,
|
||||||
|
StoreFactory,
|
||||||
|
)
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.inference import InferenceEngine, get_app, run_server, sample
|
||||||
|
from astrai.inference.network import ProtocolHandler
|
||||||
|
from astrai.inference.runtime.sample import SamplingPipeline
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"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()
|
||||||
@@ -0,0 +1,25 @@
|
|||||||
|
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__ = [
|
||||||
|
"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)
|
||||||
@@ -0,0 +1,178 @@
|
|||||||
|
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 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
|
||||||
|
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
|
||||||
|
rope_scaling: Optional[dict] = None
|
||||||
|
attn_type: str = "gqa"
|
||||||
|
num_attention_heads: Optional[int] = None
|
||||||
|
num_key_value_heads: Optional[int] = None
|
||||||
|
use_qk_norm: Optional[bool] = None
|
||||||
|
use_gated_attention: Optional[bool] = None
|
||||||
|
kv_lora_rank: Optional[int] = None
|
||||||
|
qk_nope_head_dim: Optional[int] = None
|
||||||
|
qk_rope_head_dim: Optional[int] = None
|
||||||
|
ffn_type: str = "mlp"
|
||||||
|
n_routed_experts: Optional[int] = None
|
||||||
|
n_shared_experts: Optional[int] = None
|
||||||
|
n_activated_experts: Optional[int] = None
|
||||||
|
topk_method: Optional[str] = None
|
||||||
|
moe_intermediate_size: Optional[int] = None
|
||||||
|
shared_expert_intermediate_size: Optional[int] = None
|
||||||
|
norm_topk_prob: bool = True
|
||||||
|
decoder_sparse_step: int = 1
|
||||||
|
mlp_only_layers: Optional[list[int]] = None
|
||||||
|
moe_aux_loss_coef: float = 0.01
|
||||||
|
|
||||||
|
@field_validator("attn_type")
|
||||||
|
def _validate_attn_type(cls, v: str) -> str:
|
||||||
|
if v not in _ATTN_TYPES:
|
||||||
|
raise ValueError(
|
||||||
|
f"attn_type must be one of {sorted(_ATTN_TYPES)}, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("ffn_type")
|
||||||
|
def _validate_ffn_type(cls, v: str) -> str:
|
||||||
|
if v not in _FFN_TYPES:
|
||||||
|
raise ValueError(f"ffn_type must be one of {sorted(_FFN_TYPES)}, got {v!r}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("decoder_sparse_step")
|
||||||
|
def _validate_decoder_sparse_step(cls, v: int) -> int:
|
||||||
|
if v < 1:
|
||||||
|
raise ValueError(f"decoder_sparse_step must be at least 1, got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
@ConfigFactory.register("embedding")
|
||||||
|
class EncoderConfig(BaseModelConfig):
|
||||||
|
"""Configuration for embedding encoder model.
|
||||||
|
|
||||||
|
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)
|
||||||
@@ -0,0 +1,220 @@
|
|||||||
|
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
|
||||||
|
|
||||||
|
TRAIN_TYPES = frozenset({"seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"})
|
||||||
|
PARALLEL_MODES = frozenset({"none", "ddp", "fsdp"})
|
||||||
|
BACKENDS = frozenset({"nccl", "gloo"})
|
||||||
|
START_METHODS = frozenset({"spawn", "fork", "forkserver"})
|
||||||
|
_COMPILE_MODES = frozenset({"default", "reduce-overhead", "max-autotune"})
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(config=ConfigDict(arbitrary_types_allowed=True))
|
||||||
|
class TrainConfig(BaseConfig):
|
||||||
|
"""Training configuration.
|
||||||
|
|
||||||
|
Combines hyperparameters with runtime objects (model_fn, dataset, etc.).
|
||||||
|
Only JSON-serializable fields are written to checkpoint meta via to_dict().
|
||||||
|
|
||||||
|
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 {}.
|
||||||
|
strategy_kwargs (Dict[str, Any]): Extra strategy 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)
|
||||||
|
strategy_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
@field_validator("strategy")
|
||||||
|
def _validate_strategy(cls, v: str) -> str:
|
||||||
|
if v not in TRAIN_TYPES:
|
||||||
|
raise ValueError(
|
||||||
|
f"strategy must be one of {sorted(TRAIN_TYPES)}, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("parallel_mode")
|
||||||
|
def _validate_parallel_mode(cls, v: str) -> str:
|
||||||
|
if v not in PARALLEL_MODES:
|
||||||
|
raise ValueError(
|
||||||
|
f"parallel_mode must be one of {sorted(PARALLEL_MODES)}, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("backend")
|
||||||
|
def _validate_backend(cls, v: str) -> str:
|
||||||
|
if v not in BACKENDS:
|
||||||
|
raise ValueError(f"backend must be one of {sorted(BACKENDS)}, got {v!r}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("start_method")
|
||||||
|
def _validate_start_method(cls, v: str) -> str:
|
||||||
|
if v not in START_METHODS:
|
||||||
|
raise ValueError(
|
||||||
|
f"start_method must be one of {sorted(START_METHODS)}, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("compile_mode")
|
||||||
|
def _validate_compile_mode(cls, v: Optional[str]) -> Optional[str]:
|
||||||
|
if v is not None and v not in _COMPILE_MODES:
|
||||||
|
raise ValueError(
|
||||||
|
f"compile_mode must be one of {sorted(_COMPILE_MODES)} or None, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator(
|
||||||
|
"n_epoch",
|
||||||
|
"batch_per_device",
|
||||||
|
"grad_accum_steps",
|
||||||
|
"ckpt_interval",
|
||||||
|
"val_step",
|
||||||
|
"rollout_interval",
|
||||||
|
"rollout_max_tokens",
|
||||||
|
)
|
||||||
|
def _validate_positive_int(cls, v: int) -> int:
|
||||||
|
if v <= 0:
|
||||||
|
raise ValueError(f"must be positive, got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("rollout_temperature")
|
||||||
|
def _validate_positive_float(cls, v: float) -> float:
|
||||||
|
if v <= 0:
|
||||||
|
raise ValueError(f"must be positive, got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("rollout_top_p")
|
||||||
|
def _validate_top_p(cls, v: float) -> float:
|
||||||
|
if not 0 < v <= 1:
|
||||||
|
raise ValueError(f"rollout_top_p must be in (0, 1], got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator(
|
||||||
|
"rollout_top_k", "num_workers", "neftune_alpha", "moe_aux_loss_coef"
|
||||||
|
)
|
||||||
|
def _validate_non_negative(cls, v):
|
||||||
|
if v < 0:
|
||||||
|
raise ValueError(f"must be non-negative, got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("max_grad_norm")
|
||||||
|
def _validate_max_grad_norm(cls, v: Optional[float]) -> Optional[float]:
|
||||||
|
if v is not None and v <= 0:
|
||||||
|
raise ValueError(f"max_grad_norm must be positive or None, got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("val_split")
|
||||||
|
def _validate_val_split(cls, v: Optional[float]) -> Optional[float]:
|
||||||
|
if v is not None and not 0 < v < 1:
|
||||||
|
raise ValueError(f"val_split must be in (0, 1) or None, got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def _validate_online_strategy(self) -> "TrainConfig":
|
||||||
|
if self.strategy.startswith("online_") and self.reward_model_fn is None:
|
||||||
|
raise ValueError(
|
||||||
|
f"reward_model_fn is required for online RL strategy {self.strategy!r}"
|
||||||
|
)
|
||||||
|
return self
|
||||||
@@ -0,0 +1,37 @@
|
|||||||
|
from astrai.dataset.dataset import (
|
||||||
|
BaseDataset,
|
||||||
|
DatasetFactory,
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"BaseDataset",
|
||||||
|
"DatasetFactory",
|
||||||
|
"dpo_collate_fn",
|
||||||
|
"grpo_collate_fn",
|
||||||
|
"Store",
|
||||||
|
"Streamable",
|
||||||
|
"Recordable",
|
||||||
|
"StoreFactory",
|
||||||
|
"MmapStore",
|
||||||
|
"JsonlStore",
|
||||||
|
"detect_format",
|
||||||
|
"save_bin",
|
||||||
|
"load_bin",
|
||||||
|
"RDSampler",
|
||||||
|
]
|
||||||
@@ -0,0 +1,535 @@
|
|||||||
|
"""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.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from functools import partial
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Callable, Dict, List, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
|
from astrai.dataset.storage import (
|
||||||
|
Store,
|
||||||
|
StoreFactory,
|
||||||
|
detect_format,
|
||||||
|
)
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.preprocessing.transform import TokenizeTransform
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
_DEFAULT_MESSAGES_CONFIG = {
|
||||||
|
"version": 1,
|
||||||
|
"input": {"sections": [{"field": "messages", "action": "$role", "template": True}]},
|
||||||
|
"mask": {"system": "mask", "user": "mask", "assistant": "train"},
|
||||||
|
"mask_default": "mask",
|
||||||
|
"output": {"position_ids_mode": "doc_reset"},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _build_jsonl_transform(
|
||||||
|
path: str, tokenizer_path: Optional[str] = None
|
||||||
|
) -> Optional["TokenizeTransform"]:
|
||||||
|
"""Auto-build a TokenizeTransform for JSONL eager loading.
|
||||||
|
|
||||||
|
Reads ``dataset_config.json`` from the data dir if present, or
|
||||||
|
falls back to the built-in chatml SFT config when *tokenizer_path*
|
||||||
|
is provided.
|
||||||
|
"""
|
||||||
|
root = Path(path)
|
||||||
|
config_path = root / "dataset_config.json" if root.is_dir() else None
|
||||||
|
if config_path is not None and config_path.exists():
|
||||||
|
return TokenizeTransform.from_config_file(str(config_path))
|
||||||
|
if tokenizer_path:
|
||||||
|
config = PipelineConfig.from_dict(_DEFAULT_MESSAGES_CONFIG)
|
||||||
|
return TokenizeTransform(config, tokenizer_path)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def dpo_tokenize(
|
||||||
|
record: dict,
|
||||||
|
tokenizer,
|
||||||
|
max_len: int = 2048,
|
||||||
|
) -> Optional[dict]:
|
||||||
|
"""Tokenize one DPO record into chosen/rejected + masks.
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
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}
|
||||||
|
]
|
||||||
|
|
||||||
|
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
|
||||||
|
)
|
||||||
|
|
||||||
|
full_ch = ch_ids[:max_len]
|
||||||
|
full_re = re_ids[:max_len]
|
||||||
|
|
||||||
|
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]
|
||||||
|
|
||||||
|
return {
|
||||||
|
"chosen": full_ch,
|
||||||
|
"rejected": full_re,
|
||||||
|
"chosen_mask": ch_mask,
|
||||||
|
"rejected_mask": re_mask,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
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 _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 dataset types.
|
||||||
|
|
||||||
|
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).
|
||||||
|
"""
|
||||||
|
|
||||||
|
required_keys: List[str] = []
|
||||||
|
|
||||||
|
def __init__(self, store: Store):
|
||||||
|
super().__init__()
|
||||||
|
self.store: Store = store
|
||||||
|
validate_keys(store, self.required_keys)
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
return len(self.store)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def keys(self) -> List[str]:
|
||||||
|
return self.store.keys
|
||||||
|
|
||||||
|
@property
|
||||||
|
def token_count(self) -> int:
|
||||||
|
return self.store.token_count
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
|
||||||
|
class DatasetFactory(BaseFactory["BaseDataset"]):
|
||||||
|
"""Factory for creating dataset instances by train-type.
|
||||||
|
|
||||||
|
Use :meth:`DatasetFactory.register("custom")` to register new
|
||||||
|
dataset classes; they must inherit from :class:`BaseDataset`.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def load(
|
||||||
|
cls,
|
||||||
|
train_type: str,
|
||||||
|
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: 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.
|
||||||
|
"""
|
||||||
|
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
|
||||||
|
|
||||||
|
processor = cls._maybe_build_processor(
|
||||||
|
train_type, storage_type, tokenizer_path, max_len
|
||||||
|
)
|
||||||
|
|
||||||
|
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(
|
||||||
|
"JSONL dataset config not found. Expected "
|
||||||
|
"dataset_config.json alongside *.jsonl files, pass "
|
||||||
|
"tokenizer_path= for the built-in messages config, or "
|
||||||
|
"use processor= for lazy on-the-fly tokenisation."
|
||||||
|
)
|
||||||
|
store.load(load_path, transform=transform, **kwargs)
|
||||||
|
else:
|
||||||
|
store.load(load_path, **kwargs)
|
||||||
|
|
||||||
|
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.
|
||||||
|
|
||||||
|
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.
|
||||||
|
|
||||||
|
Stream mode: ``store.fetch(begin, end, "sequence")`` returns the
|
||||||
|
input window; the +1 shifted call returns the next-token target.
|
||||||
|
"""
|
||||||
|
|
||||||
|
required_keys = ["sequence"]
|
||||||
|
|
||||||
|
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.
|
||||||
|
|
||||||
|
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.
|
||||||
|
"""
|
||||||
|
|
||||||
|
required_keys = ["sequence", "loss_mask", "position_ids"]
|
||||||
|
|
||||||
|
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):
|
||||||
|
"""Record-structured dataset for Direct Preference Optimization.
|
||||||
|
|
||||||
|
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.
|
||||||
|
|
||||||
|
Two loading paths (handled by :class:`DatasetFactory`):
|
||||||
|
|
||||||
|
- **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``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
required_keys = ["chosen", "rejected", "chosen_mask", "rejected_mask"]
|
||||||
|
|
||||||
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
|
return {
|
||||||
|
"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 offline Group Relative Policy Optimization.
|
||||||
|
|
||||||
|
Each sample is one prompt with its group of responses and scalar
|
||||||
|
rewards — an independent training unit with no windowing or stride.
|
||||||
|
|
||||||
|
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]:
|
||||||
|
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.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),
|
||||||
|
}
|
||||||
@@ -1,20 +1,28 @@
|
|||||||
import torch
|
|
||||||
import torch.distributed as dist
|
|
||||||
|
|
||||||
from torch.utils.data import Dataset, Sampler
|
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
|
from torch.utils.data import Dataset, Sampler
|
||||||
|
|
||||||
|
|
||||||
|
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.
|
||||||
|
"""
|
||||||
|
|
||||||
class ResumableDistributedSampler(Sampler[int]):
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
data_source: Dataset,
|
data_source: Dataset,
|
||||||
start_epoch: int=0,
|
start_epoch: int = 0,
|
||||||
start_iter: int=0,
|
start_iter: int = 0,
|
||||||
seed: int=42,
|
seed: int = 42,
|
||||||
drop_last: bool=False,
|
drop_last: bool = False,
|
||||||
shuffle: bool=True,
|
shuffle: bool = True,
|
||||||
process_group: Optional[dist.ProcessGroup]=None,
|
process_group: Optional[dist.ProcessGroup] = None,
|
||||||
):
|
):
|
||||||
self.epoch = start_epoch
|
self.epoch = start_epoch
|
||||||
self.iter = start_iter
|
self.iter = start_iter
|
||||||
@@ -40,9 +48,10 @@ class ResumableDistributedSampler(Sampler[int]):
|
|||||||
self.drop_last = drop_last
|
self.drop_last = drop_last
|
||||||
self.shuffle = shuffle
|
self.shuffle = shuffle
|
||||||
|
|
||||||
offset = 0 if drop_last else self.num_replicas - 1
|
offset = 0 if drop_last else self.num_replicas - 1
|
||||||
self.num_samples_per_replica = (self.num_samples + offset) // self.num_replicas
|
self.num_samples_per_replica = (self.num_samples + offset) // self.num_replicas
|
||||||
self.total_size = self.num_samples_per_replica * 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
|
self._indices = None
|
||||||
|
|
||||||
@@ -58,10 +67,10 @@ class ResumableDistributedSampler(Sampler[int]):
|
|||||||
padding_size = self.total_size - len(indices)
|
padding_size = self.total_size - len(indices)
|
||||||
indices += indices[:padding_size]
|
indices += indices[:padding_size]
|
||||||
|
|
||||||
local_indices = indices[self.rank:self.total_size:self.num_replicas]
|
local_indices = indices[self.rank : self.total_size : self.num_replicas]
|
||||||
|
|
||||||
self.iter = self.iter % self.num_samples_per_replica
|
self.iter = self.iter % self.num_samples_per_replica
|
||||||
self._indices = local_indices[self.iter:]
|
self._indices = local_indices[self.iter :]
|
||||||
|
|
||||||
def __iter__(self):
|
def __iter__(self):
|
||||||
if self._indices is None:
|
if self._indices is None:
|
||||||
@@ -73,6 +82,12 @@ class ResumableDistributedSampler(Sampler[int]):
|
|||||||
|
|
||||||
self.epoch += 1
|
self.epoch += 1
|
||||||
self._indices = None
|
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):
|
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._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,818 @@
|
|||||||
|
"""Attention backend abstraction with context-manager switching.
|
||||||
|
|
||||||
|
The backend encapsulates KV cache I/O and attention computation. The
|
||||||
|
attention module (GQA/MLA) keeps projections, rotary, QK-norm, gating,
|
||||||
|
and output projection; the backend handles everything from "write K/V
|
||||||
|
to cache" through "SDPA output".
|
||||||
|
|
||||||
|
Usage — mirroring ``torch.nn.attention.sdpa_kernel``:
|
||||||
|
|
||||||
|
from astrai.extension import attn_backend, ATTN_BACKEND
|
||||||
|
|
||||||
|
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
|
||||||
|
engine.generate("hello")
|
||||||
|
|
||||||
|
# or with an instance:
|
||||||
|
with attn_backend(TorchNativeBackend()):
|
||||||
|
...
|
||||||
|
|
||||||
|
# or the shorthand (instance is itself a context manager):
|
||||||
|
with TorchNativeBackend():
|
||||||
|
...
|
||||||
|
|
||||||
|
Thread-safe via ``contextvars`` — each scheduler thread gets its own
|
||||||
|
active backend. Backend resolution follows a strict precedence:
|
||||||
|
|
||||||
|
1. explicit ``attn_backend(...)`` context (wins over everything),
|
||||||
|
2. the process-wide ``ASTR_BACKEND`` environment override,
|
||||||
|
3. an implicit default picked from the available backends
|
||||||
|
(cuda > flash > torch).
|
||||||
|
|
||||||
|
Capability is polymorphic: every backend declares ``available()``
|
||||||
|
(machine-level) and ``supports_call(...)`` (per-call), so adding a new
|
||||||
|
backend requires no changes to the resolution logic. Training calls
|
||||||
|
(``fwd=None``, no KV cache) resolve through the same priority list: the
|
||||||
|
CUDA cache kernels cannot run without a cache, so they fall back to
|
||||||
|
flash (when it can handle the call — mask-free/causal only) and finally
|
||||||
|
to the reference ``TorchNativeBackend``.
|
||||||
|
|
||||||
|
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
||||||
|
(blhd). The backend returns ``[batch, seq_len, n_heads * head_dim]``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import contextvars
|
||||||
|
import enum
|
||||||
|
import functools
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import threading
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from typing import TYPE_CHECKING, Dict, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.extension.loader import is_available
|
||||||
|
from astrai.extension.ops.attention import (
|
||||||
|
attn_paged_decode,
|
||||||
|
attn_paged_prefill,
|
||||||
|
)
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
|
||||||
|
try:
|
||||||
|
import flash_attn as _flash_attn
|
||||||
|
except Exception:
|
||||||
|
_flash_attn = None
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from astrai.inference.cache import KVCache
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
_default_backend_lock = threading.Lock()
|
||||||
|
_env_backend_name: Optional[str] = None
|
||||||
|
_env_backend: Optional["AttentionBackend"] = None
|
||||||
|
_current_backend: contextvars.ContextVar[Optional["AttentionBackend"]] = (
|
||||||
|
contextvars.ContextVar("attn_backend", default=None)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Backends are stateless — one canonical instance per class, created lazily
|
||||||
|
# and reused everywhere (resolution, fallback, context managers).
|
||||||
|
_singletons: Dict[type, "AttentionBackend"] = {}
|
||||||
|
|
||||||
|
|
||||||
|
@functools.lru_cache(maxsize=1)
|
||||||
|
def flash_attn_available() -> bool:
|
||||||
|
if not torch.cuda.is_available():
|
||||||
|
return False
|
||||||
|
fa = _flash_attn
|
||||||
|
if fa is None:
|
||||||
|
return False
|
||||||
|
|
||||||
|
try:
|
||||||
|
major = int(fa.__version__.split(".")[0])
|
||||||
|
cc = torch.cuda.get_device_capability()
|
||||||
|
cc_num = cc[0] * 10 + cc[1]
|
||||||
|
except Exception:
|
||||||
|
major, cc_num = 0, 0
|
||||||
|
if (major >= 3 and cc_num < 90) or (major < 3 and 0 < cc_num < 70):
|
||||||
|
return False
|
||||||
|
|
||||||
|
try:
|
||||||
|
if not hasattr(fa, "flash_attn_func"):
|
||||||
|
return False
|
||||||
|
x = torch.zeros(1, 1, 1, 64, device="cuda", dtype=torch.bfloat16)
|
||||||
|
out = fa.flash_attn_func(x, x, x, causal=True)
|
||||||
|
return bool(torch.isfinite(out).all().item())
|
||||||
|
except Exception:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
class ATTN_BACKEND(enum.Enum):
|
||||||
|
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
|
||||||
|
|
||||||
|
TORCH_NATIVE = "torch_native"
|
||||||
|
CUDA = "cuda"
|
||||||
|
FLASH = "flash"
|
||||||
|
|
||||||
|
|
||||||
|
def _instance(backend_cls: type) -> "AttentionBackend":
|
||||||
|
"""Return the canonical singleton instance for a backend class.
|
||||||
|
|
||||||
|
Backends hold no per-instance state, so a single cached instance is
|
||||||
|
safe and avoids per-call allocation on the attention hot path.
|
||||||
|
"""
|
||||||
|
backend = _singletons.get(backend_cls)
|
||||||
|
if backend is None:
|
||||||
|
backend = backend_cls()
|
||||||
|
_singletons[backend_cls] = backend
|
||||||
|
return backend
|
||||||
|
|
||||||
|
|
||||||
|
@functools.lru_cache(maxsize=1)
|
||||||
|
def _priority_backends() -> Tuple["AttentionBackend", ...]:
|
||||||
|
"""Available backends in priority order: cuda -> flash -> torch.
|
||||||
|
|
||||||
|
Computed once (machine availability cannot change at runtime) and
|
||||||
|
cached forever; the tuple always ends with ``TorchNativeBackend``,
|
||||||
|
which is unconditionally available.
|
||||||
|
"""
|
||||||
|
return tuple(
|
||||||
|
_instance(cls)
|
||||||
|
for cls in (CudaBackend, FlashAttnBackend, TorchNativeBackend)
|
||||||
|
if cls.available()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_default_backend() -> "AttentionBackend":
|
||||||
|
"""Pick the highest-priority available backend (cuda -> flash -> torch).
|
||||||
|
|
||||||
|
Resolved lazily on first use and cached via ``_priority_backends``.
|
||||||
|
Per-call capability fallback happens in ``attention()``, so the
|
||||||
|
default is safe for training and fp32 models.
|
||||||
|
"""
|
||||||
|
return _priority_backends()[0]
|
||||||
|
|
||||||
|
|
||||||
|
def _environment_backend() -> Optional["AttentionBackend"]:
|
||||||
|
"""Resolve the process-wide ``ASTR_BACKEND`` override, if configured."""
|
||||||
|
global _env_backend, _env_backend_name
|
||||||
|
name = os.environ.get("ASTR_BACKEND", "").strip().lower()
|
||||||
|
if not name:
|
||||||
|
return None
|
||||||
|
if name != _env_backend_name:
|
||||||
|
with _default_backend_lock:
|
||||||
|
if name != _env_backend_name:
|
||||||
|
try:
|
||||||
|
_env_backend = _resolve_backend(name)
|
||||||
|
except (ValueError, RuntimeError):
|
||||||
|
_env_backend = None
|
||||||
|
logger.warning(
|
||||||
|
"ASTR_BACKEND=%r is not a registered attention backend; "
|
||||||
|
"falling back to default resolution",
|
||||||
|
name,
|
||||||
|
)
|
||||||
|
_env_backend_name = name
|
||||||
|
return _env_backend
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_backend(
|
||||||
|
backend: Optional[Union[str, ATTN_BACKEND, "AttentionBackend", type]] = None,
|
||||||
|
) -> "AttentionBackend":
|
||||||
|
"""Resolve a backend configuration to its canonical instance.
|
||||||
|
|
||||||
|
Accepts a registered name, ``ATTN_BACKEND`` enum value, backend class,
|
||||||
|
or instance. Names/classes resolve to the shared singleton; a caller
|
||||||
|
may still pass its own instance to opt out of sharing.
|
||||||
|
"""
|
||||||
|
if backend is not None:
|
||||||
|
if isinstance(backend, ATTN_BACKEND):
|
||||||
|
return _instance(AttentionBackendFactory.get_component_class(backend.value))
|
||||||
|
if isinstance(backend, str):
|
||||||
|
return _instance(AttentionBackendFactory.get_component_class(backend))
|
||||||
|
if isinstance(backend, type) and issubclass(backend, AttentionBackend):
|
||||||
|
return _instance(backend)
|
||||||
|
if isinstance(backend, AttentionBackend):
|
||||||
|
return backend
|
||||||
|
raise TypeError(
|
||||||
|
f"expected a registered name, ATTN_BACKEND, AttentionBackend type, "
|
||||||
|
f"or instance, got {type(backend).__name__}"
|
||||||
|
)
|
||||||
|
return _resolve_default_backend()
|
||||||
|
|
||||||
|
|
||||||
|
def get_backend(
|
||||||
|
use_default: bool = True,
|
||||||
|
) -> Optional["AttentionBackend"]:
|
||||||
|
"""Resolve the active backend: explicit context > env > default.
|
||||||
|
|
||||||
|
An ``attn_backend(...)`` context is the caller's explicit choice and
|
||||||
|
always wins. ``ASTR_BACKEND`` is a process-wide override consulted
|
||||||
|
only when no context is set. Pass ``use_default=False`` at request
|
||||||
|
submission to retain only an environment override or the caller's
|
||||||
|
:func:`attn_backend` value.
|
||||||
|
"""
|
||||||
|
context_backend = _current_backend.get()
|
||||||
|
if context_backend is not None:
|
||||||
|
return context_backend
|
||||||
|
env_backend = _environment_backend()
|
||||||
|
if env_backend is not None:
|
||||||
|
return env_backend
|
||||||
|
return _resolve_default_backend() if use_default else None
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def attn_backend(backend: Union[str, ATTN_BACKEND, "AttentionBackend", type]):
|
||||||
|
"""Context manager to select an attention backend.
|
||||||
|
|
||||||
|
Mirrors ``torch.nn.attention.sdpa_kernel``. Accepts an
|
||||||
|
registered name, ``ATTN_BACKEND`` enum value, backend class, or instance.
|
||||||
|
|
||||||
|
Examples::
|
||||||
|
|
||||||
|
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
|
||||||
|
...
|
||||||
|
with attn_backend(TorchNativeBackend):
|
||||||
|
...
|
||||||
|
with attn_backend(TorchNativeBackend()):
|
||||||
|
...
|
||||||
|
"""
|
||||||
|
instance = _resolve_backend(backend)
|
||||||
|
token = _current_backend.set(instance)
|
||||||
|
try:
|
||||||
|
yield instance
|
||||||
|
finally:
|
||||||
|
_current_backend.reset(token)
|
||||||
|
|
||||||
|
|
||||||
|
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
||||||
|
"""Expand KV heads to match Q heads for GQA."""
|
||||||
|
if n_rep == 1:
|
||||||
|
return x
|
||||||
|
n_heads, head_dim = x.shape[-2:]
|
||||||
|
return (
|
||||||
|
x.unsqueeze(-2)
|
||||||
|
.expand(*x.shape[:-2], n_heads, n_rep, head_dim)
|
||||||
|
.reshape(*x.shape[:-2], n_heads * n_rep, head_dim)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def attention(
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"] = None,
|
||||||
|
layer_id: int = 0,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
fwd: Optional[str] = None,
|
||||||
|
backend: Optional[Union[str, ATTN_BACKEND, "AttentionBackend", type]] = None,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Functional attention entry point — mirrors ``F.scaled_dot_product_attention``.
|
||||||
|
|
||||||
|
Delegates to the active backend. ``backend`` (optional) is an explicit
|
||||||
|
escape hatch; when omitted the backend is resolved as
|
||||||
|
explicit context > ``ASTR_BACKEND`` env > default (cuda > flash > torch).
|
||||||
|
Handles KV cache I/O, GQA head expansion, and causal masking so the
|
||||||
|
caller only needs to provide projected q/k/v.
|
||||||
|
|
||||||
|
Training calls (``fwd=None``, ``kv_cache=None``) resolve through the
|
||||||
|
same capability chain — the CUDA cache kernels cannot run without a
|
||||||
|
cache, so they fall back to flash (mask-free/causal calls only) and
|
||||||
|
finally to torch SDPA. An explicitly-selected backend that cannot
|
||||||
|
handle the call raises — an implicit one falls back down the priority
|
||||||
|
list to the first capable backend.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [batch, q_len, n_heads, head_dim] (blhd)
|
||||||
|
k: [batch, q_len, n_kv_heads, head_dim] (blhd)
|
||||||
|
v: [batch, q_len, n_kv_heads, head_dim] (blhd)
|
||||||
|
kv_cache: cache dataclass, or None for training (no cache).
|
||||||
|
layer_id: transformer layer index for buffer access.
|
||||||
|
attn_mask: pre-built attention mask (SDPA-compatible).
|
||||||
|
is_causal: whether to apply causal masking.
|
||||||
|
fwd: "prefill" / "decode" for inference, None for training.
|
||||||
|
backend: optional explicit backend (name, enum, class, or instance).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, q_len, n_heads * head_dim]
|
||||||
|
"""
|
||||||
|
if backend is not None:
|
||||||
|
selected = _resolve_backend(backend)
|
||||||
|
explicit = True
|
||||||
|
else:
|
||||||
|
context_backend = _current_backend.get()
|
||||||
|
explicit = context_backend is not None
|
||||||
|
# Resolve through the same chain as inference: explicit context >
|
||||||
|
# ASTR_BACKEND env > default. Training calls (fwd=None, no cache)
|
||||||
|
# land on the CUDA backend and fall back by capability below —
|
||||||
|
# flash when it can handle the call, else torch SDPA.
|
||||||
|
selected = get_backend()
|
||||||
|
assert selected is not None
|
||||||
|
|
||||||
|
if not selected.supports_call(q, kv_cache, attn_mask, is_causal, fwd):
|
||||||
|
if explicit:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Explicitly-set backend {type(selected).__name__} cannot "
|
||||||
|
f"handle this attention call (shape={q.shape}, "
|
||||||
|
f"dtype={q.dtype}, kv_cache={'none' if kv_cache is None else 'present'}, "
|
||||||
|
f"attn_mask={'none' if attn_mask is None else 'present'}). "
|
||||||
|
f"Remove the attn_backend() context or switch to a compatible backend."
|
||||||
|
)
|
||||||
|
selected = next(
|
||||||
|
(
|
||||||
|
candidate
|
||||||
|
for candidate in _priority_backends()
|
||||||
|
if candidate.supports_call(q, kv_cache, attn_mask, is_causal, fwd)
|
||||||
|
),
|
||||||
|
_instance(TorchNativeBackend),
|
||||||
|
)
|
||||||
|
return selected.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal, fwd)
|
||||||
|
|
||||||
|
|
||||||
|
class AttentionBackend(ABC):
|
||||||
|
"""Abstract base for attention computation strategies.
|
||||||
|
|
||||||
|
Subclasses implement ``fwd_decode`` (q_len == 1, with cache) and
|
||||||
|
``fwd_prefill`` (q_len > 1, with or without cache). The public
|
||||||
|
``forward`` method dispatches based on q_len.
|
||||||
|
|
||||||
|
Capability contract — every backend declares:
|
||||||
|
|
||||||
|
* ``available()`` — machine-level: can this backend exist here
|
||||||
|
(kernel ``.so`` loaded, flash-attn present, GPU available)?
|
||||||
|
Used once to build the default priority list.
|
||||||
|
* ``supports_call(q, kv_cache, attn_mask, is_causal, fwd)`` — can this
|
||||||
|
backend run this *specific* call (shape/dtype/cache/mask)? Used by
|
||||||
|
``attention()`` for the per-call fallback. Resolution logic never
|
||||||
|
checks concrete backend types, so adding a backend requires no
|
||||||
|
changes outside its own class.
|
||||||
|
|
||||||
|
Three equivalent ways to activate a backend::
|
||||||
|
|
||||||
|
with attn_backend(ATTN_BACKEND.TORCH_NATIVE): # enum
|
||||||
|
...
|
||||||
|
with attn_backend(TorchNativeBackend): # class
|
||||||
|
...
|
||||||
|
with TorchNativeBackend(): # instance
|
||||||
|
...
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __enter__(self) -> "AttentionBackend":
|
||||||
|
self._token = _current_backend.set(self)
|
||||||
|
return self
|
||||||
|
|
||||||
|
def __exit__(self, *exc) -> None:
|
||||||
|
_current_backend.reset(self._token)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
@abstractmethod
|
||||||
|
def available(cls) -> bool:
|
||||||
|
"""Return True if this backend can run on the current machine.
|
||||||
|
|
||||||
|
Checks static availability only (compiled kernels, optional
|
||||||
|
packages, GPU presence) — not call-specific constraints.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def supports_call(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
attn_mask: Optional[Tensor],
|
||||||
|
is_causal: bool,
|
||||||
|
fwd: Optional[str],
|
||||||
|
) -> bool:
|
||||||
|
"""Return True if this backend can run this specific attention call.
|
||||||
|
|
||||||
|
Called on the canonical singleton instance (or a caller-provided
|
||||||
|
one); must be side-effect free.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
fwd: Optional[str] = None,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Dispatch to decode or extend based on q_len.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [batch, q_len, n_heads, head_dim]
|
||||||
|
k: [batch, q_len, n_kv_heads, head_dim]
|
||||||
|
v: [batch, q_len, n_kv_heads, head_dim]
|
||||||
|
kv_cache: cache dataclass, or None for training (no cache).
|
||||||
|
layer_id: transformer layer index for buffer access.
|
||||||
|
attn_mask: pre-built attention mask compatible with SDPA.
|
||||||
|
is_causal: whether to apply causal masking.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, q_len, n_heads * head_dim]
|
||||||
|
"""
|
||||||
|
if fwd == "decode":
|
||||||
|
return self.fwd_decode(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
|
if fwd == "prefill" or fwd is None:
|
||||||
|
return self.fwd_prefill(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
|
raise ValueError(f"unsupported attention forward mode: {fwd}")
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def fwd_decode(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Single-token decode with KV cache."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def fwd_prefill(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Multi-token prefill or training forward."""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def supports_graph() -> bool:
|
||||||
|
"""Return True if this backend supports CUDA-graph capture.
|
||||||
|
|
||||||
|
Override in subclasses that can run under ``torch.cuda.graph``.
|
||||||
|
|
||||||
|
Called on the *active* backend instance (or its class) — a cheap
|
||||||
|
boolean check with no side-effects.
|
||||||
|
"""
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
class AttentionBackendFactory(BaseFactory[AttentionBackend]):
|
||||||
|
"""Factory for registered attention backends."""
|
||||||
|
|
||||||
|
|
||||||
|
@AttentionBackendFactory.register(ATTN_BACKEND.TORCH_NATIVE.value)
|
||||||
|
class TorchNativeBackend(AttentionBackend):
|
||||||
|
"""Reference backend using torch SDPA with indirect KV cache indexing.
|
||||||
|
|
||||||
|
Writes new K/V into the cache buffers, gathers the full sequence K/V
|
||||||
|
via ``req_to_token`` indirect indexing, then calls
|
||||||
|
``F.scaled_dot_product_attention``.
|
||||||
|
|
||||||
|
For training (``kv_cache is None``), skips cache I/O entirely and
|
||||||
|
runs SDPA directly on the projected q/k/v.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def available(cls) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
def supports_call(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
attn_mask: Optional[Tensor],
|
||||||
|
is_causal: bool,
|
||||||
|
fwd: Optional[str],
|
||||||
|
) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
def fwd_decode(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
|
|
||||||
|
def fwd_prefill(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
|
|
||||||
|
def _forward(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
if q.ndim == 4:
|
||||||
|
n_rep = q.size(2) // k.size(2)
|
||||||
|
if n_rep > 1:
|
||||||
|
k = repeat_kv(k, n_rep)
|
||||||
|
v = repeat_kv(v, n_rep)
|
||||||
|
return (
|
||||||
|
F.scaled_dot_product_attention(
|
||||||
|
q.permute(0, 2, 1, 3),
|
||||||
|
k.permute(0, 2, 1, 3),
|
||||||
|
v.permute(0, 2, 1, 3),
|
||||||
|
attn_mask,
|
||||||
|
is_causal=is_causal,
|
||||||
|
)
|
||||||
|
.permute(0, 2, 1, 3)
|
||||||
|
.contiguous()
|
||||||
|
)
|
||||||
|
|
||||||
|
if kv_cache is None or kv_cache.qo_indptr is None:
|
||||||
|
raise ValueError("packed attention requires KV cache metadata")
|
||||||
|
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||||
|
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||||
|
outputs = []
|
||||||
|
n_rep = q.size(1) // k.size(1)
|
||||||
|
for i in range(kv_cache.req_pool_indices.numel()):
|
||||||
|
q_start = int(kv_cache.qo_indptr[i])
|
||||||
|
q_end = int(kv_cache.qo_indptr[i + 1])
|
||||||
|
indices = kv_cache.req_to_token[
|
||||||
|
kv_cache.req_pool_indices[i], : kv_cache.seq_lens[i]
|
||||||
|
]
|
||||||
|
k_i = kv_cache.k_buffer[layer_id, indices]
|
||||||
|
v_i = kv_cache.v_buffer[layer_id, indices]
|
||||||
|
if n_rep > 1:
|
||||||
|
k_i = repeat_kv(k_i, n_rep)
|
||||||
|
v_i = repeat_kv(v_i, n_rep)
|
||||||
|
q_len = q_end - q_start
|
||||||
|
kv_len = k_i.size(0)
|
||||||
|
q_pos = torch.arange(kv_len - q_len, kv_len, device=q.device)
|
||||||
|
causal_mask = q_pos[:, None] >= torch.arange(kv_len, device=q.device)
|
||||||
|
out = F.scaled_dot_product_attention(
|
||||||
|
q[q_start:q_end].transpose(0, 1).unsqueeze(0),
|
||||||
|
k_i.transpose(0, 1).unsqueeze(0),
|
||||||
|
v_i.transpose(0, 1).unsqueeze(0),
|
||||||
|
attn_mask=causal_mask,
|
||||||
|
)
|
||||||
|
outputs.append(out.squeeze(0).transpose(0, 1))
|
||||||
|
return torch.cat(outputs)
|
||||||
|
|
||||||
|
|
||||||
|
@AttentionBackendFactory.register(ATTN_BACKEND.CUDA.value)
|
||||||
|
class CudaBackend(AttentionBackend):
|
||||||
|
"""CUDA kernel backend with direct KV cache access.
|
||||||
|
|
||||||
|
Decode path: writes K/V to the flat pool, then calls
|
||||||
|
``attn_paged_decode`` with req_to_token + kv_indptr.
|
||||||
|
|
||||||
|
Prefill path: writes K/V to the flat pool, then calls
|
||||||
|
``attn_paged_prefill`` with ragged-batch support via qo_indptr +
|
||||||
|
kv_indptr.
|
||||||
|
|
||||||
|
``kv_cache is None`` (training) raises — the per-call fallback to
|
||||||
|
torch SDPA for training / fp32 / unsupported head_dim happens in the
|
||||||
|
``attention()`` entry point.
|
||||||
|
|
||||||
|
Raises ``RuntimeError`` if the required kernel is not available.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Head dims supported by the CUDA kernels (single source of truth).
|
||||||
|
HEAD_DIMS = (32, 64, 128, 256)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def available(cls) -> bool:
|
||||||
|
return (
|
||||||
|
torch.cuda.is_available()
|
||||||
|
and is_available("attn_paged_decode")
|
||||||
|
and is_available("attn_paged_prefill")
|
||||||
|
)
|
||||||
|
|
||||||
|
def supports_call(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
attn_mask: Optional[Tensor],
|
||||||
|
is_causal: bool,
|
||||||
|
fwd: Optional[str],
|
||||||
|
) -> bool:
|
||||||
|
# The CUDA kernels are bf16-only, support head_dim in
|
||||||
|
# HEAD_DIMS, and need a KV cache (decode/prefill); everything
|
||||||
|
# else falls back down the priority list to torch.
|
||||||
|
return (
|
||||||
|
fwd in ("prefill", "decode")
|
||||||
|
and kv_cache is not None
|
||||||
|
and q.ndim == 3
|
||||||
|
and q.dtype == torch.bfloat16
|
||||||
|
and q.size(-1) in self.HEAD_DIMS
|
||||||
|
and is_available(f"attn_paged_{fwd}")
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def supports_graph() -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
def fwd_decode(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
if kv_cache is None:
|
||||||
|
raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
|
||||||
|
|
||||||
|
kv_indptr = kv_cache.kv_indptr
|
||||||
|
|
||||||
|
out = attn_paged_decode(
|
||||||
|
q,
|
||||||
|
kv_cache.k_buffer[layer_id],
|
||||||
|
kv_cache.v_buffer[layer_id],
|
||||||
|
kv_cache.req_to_token,
|
||||||
|
kv_cache.req_pool_indices,
|
||||||
|
kv_indptr,
|
||||||
|
new_k=k,
|
||||||
|
new_v=v,
|
||||||
|
is_causal=True,
|
||||||
|
o_part_buf=kv_cache.decode_o_part,
|
||||||
|
ml_part_buf=kv_cache.decode_ml_part,
|
||||||
|
out_buf=kv_cache.decode_out,
|
||||||
|
)
|
||||||
|
return out
|
||||||
|
|
||||||
|
def fwd_prefill(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
if kv_cache is None:
|
||||||
|
raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
|
||||||
|
|
||||||
|
loc = kv_cache.out_cache_loc
|
||||||
|
kv_cache.k_buffer[layer_id, loc] = k
|
||||||
|
kv_cache.v_buffer[layer_id, loc] = v
|
||||||
|
|
||||||
|
out = attn_paged_prefill(
|
||||||
|
q,
|
||||||
|
kv_cache.k_buffer[layer_id],
|
||||||
|
kv_cache.v_buffer[layer_id],
|
||||||
|
kv_cache.req_to_token,
|
||||||
|
kv_cache.req_pool_indices,
|
||||||
|
kv_cache.kv_indptr,
|
||||||
|
kv_cache.qo_indptr,
|
||||||
|
kv_cache.q_tile_to_batch,
|
||||||
|
kv_cache.q_tile_to_index,
|
||||||
|
attn_mask,
|
||||||
|
is_causal=is_causal,
|
||||||
|
)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
@AttentionBackendFactory.register(ATTN_BACKEND.FLASH.value)
|
||||||
|
class FlashAttnBackend(AttentionBackend):
|
||||||
|
"""FlashAttention backend via the optional ``flash-attn`` package.
|
||||||
|
|
||||||
|
Decode (q_len=1, contiguous cache): uses ``flash_attn_with_kvcache``,
|
||||||
|
which reads K/V directly from the flat pool via cache_batch_idx +
|
||||||
|
cache_seqlens — no materialized KV gather.
|
||||||
|
|
||||||
|
Prefill / non-contiguous decode: falls back to KV gather +
|
||||||
|
``flash_attn_func``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def available(cls) -> bool:
|
||||||
|
return flash_attn_available()
|
||||||
|
|
||||||
|
def supports_call(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
attn_mask: Optional[Tensor],
|
||||||
|
is_causal: bool,
|
||||||
|
fwd: Optional[str],
|
||||||
|
) -> bool:
|
||||||
|
if not self.available():
|
||||||
|
return False
|
||||||
|
if q.dtype not in (torch.float16, torch.bfloat16):
|
||||||
|
return False
|
||||||
|
if fwd is not None:
|
||||||
|
return q.ndim == 3 and hasattr(_flash_attn, "flash_attn_varlen_func")
|
||||||
|
# Dense (training) path: flash_attn_func cannot apply a custom
|
||||||
|
# mask, so only mask-free calls are supported — ``is_causal`` is
|
||||||
|
# a flag, not a mask. Masked training (SFT/DPO/GRPO) must fall
|
||||||
|
# back to TorchNativeBackend instead of silently ignoring the mask.
|
||||||
|
return attn_mask is None
|
||||||
|
|
||||||
|
def fwd_decode(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
return self._forward_packed(q, k, v, kv_cache, layer_id)
|
||||||
|
|
||||||
|
def fwd_prefill(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional["KVCache"],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
if q.ndim == 3:
|
||||||
|
return self._forward_packed(q, k, v, kv_cache, layer_id)
|
||||||
|
return self._forward_dense(q, k, v, attn_mask, is_causal)
|
||||||
|
|
||||||
|
def _forward_dense(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
n_rep = q.size(2) // k.size(2)
|
||||||
|
if n_rep > 1:
|
||||||
|
k = repeat_kv(k, n_rep)
|
||||||
|
v = repeat_kv(v, n_rep)
|
||||||
|
|
||||||
|
if attn_mask is not None:
|
||||||
|
raise ValueError(
|
||||||
|
"FlashAttnBackend cannot handle a custom attention mask; "
|
||||||
|
"use a causal mask or select TorchNativeBackend."
|
||||||
|
)
|
||||||
|
fa = _flash_attn
|
||||||
|
if fa is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
"FlashAttnBackend requires the optional 'flash-attn' package. "
|
||||||
|
"Install with `pip install flash-attn`."
|
||||||
|
)
|
||||||
|
out = fa.flash_attn_func(
|
||||||
|
q.contiguous(),
|
||||||
|
k.contiguous(),
|
||||||
|
v.contiguous(),
|
||||||
|
causal=is_causal,
|
||||||
|
)
|
||||||
|
return out.contiguous()
|
||||||
|
|
||||||
|
def _forward_packed(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: "KVCache",
|
||||||
|
layer_id: int,
|
||||||
|
) -> Tensor:
|
||||||
|
fa = _flash_attn
|
||||||
|
if fa is None or not hasattr(fa, "flash_attn_varlen_func"):
|
||||||
|
raise RuntimeError("packed inference requires flash_attn_varlen_func")
|
||||||
|
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||||
|
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||||
|
page_table = kv_cache.req_to_token[
|
||||||
|
kv_cache.req_pool_indices, : kv_cache.max_len
|
||||||
|
]
|
||||||
|
positions = torch.arange(kv_cache.max_len, device=q.device)
|
||||||
|
indices = page_table[positions.unsqueeze(0) < kv_cache.seq_lens.unsqueeze(1)]
|
||||||
|
k_flat = kv_cache.k_buffer[layer_id, indices].contiguous()
|
||||||
|
v_flat = kv_cache.v_buffer[layer_id, indices].contiguous()
|
||||||
|
out = fa.flash_attn_varlen_func(
|
||||||
|
q.contiguous(),
|
||||||
|
k_flat,
|
||||||
|
v_flat,
|
||||||
|
kv_cache.qo_indptr,
|
||||||
|
kv_cache.kv_indptr,
|
||||||
|
int((kv_cache.qo_indptr[1:] - kv_cache.qo_indptr[:-1]).max()),
|
||||||
|
int(kv_cache.seq_lens.max()),
|
||||||
|
dropout_p=0.0,
|
||||||
|
causal=True,
|
||||||
|
)
|
||||||
|
return out
|
||||||
@@ -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,484 @@
|
|||||||
|
"""FP8 training: scaling recipes, per-tensor state, and aten::linear dispatch.
|
||||||
|
|
||||||
|
Layered (see ``ops/fp8.py`` for the CUDA interface adapter):
|
||||||
|
1. ``ops.fp8`` — the only module touching the pybind.
|
||||||
|
2. This module (strategy layer): scaling *recipes* (TE-style delayed scaling
|
||||||
|
or dynamic current-amax scaling), per-tensor scales + amax history, and the
|
||||||
|
``fp8_autocast`` context manager (like ``torch.autocast``).
|
||||||
|
3. aten::linear integration: registers the CUDA + AutogradCUDA impls.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
from astrai.extension.fp8 import fp8_autocast
|
||||||
|
with fp8_autocast(enabled=True, fp8_format="hybrid"):
|
||||||
|
logits = model(input_ids)
|
||||||
|
loss.backward() # fp8 backward runs anywhere; fwd captured state on the node
|
||||||
|
|
||||||
|
Format defaults follow the ecosystem consensus: E4M3 forward / E5M2 backward
|
||||||
|
("hybrid"); every operand's scale is a quantization step derived from its amax
|
||||||
|
history by the active recipe.
|
||||||
|
|
||||||
|
The context mirrors ``torch.autocast`` (``autocast_mode.py``): the active
|
||||||
|
``(enabled, recipe, fp8_format)`` triple is thread-local (a ``contextvars``
|
||||||
|
``ContextVar``, absent outside any region), and the manager is class-based and
|
||||||
|
reentrant with nested ``enabled=False`` disabling dispatch inside it. The module
|
||||||
|
targets *training*: every step quantizes x/w/g fresh (no weight-cast cache — the
|
||||||
|
optimizer bumps the weight version each step, so a torch-style cached_cast would
|
||||||
|
miss anyway), and the per-operand scales come from the delayed/dynamic recipe.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import functools
|
||||||
|
from contextvars import ContextVar, Token
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from enum import Enum
|
||||||
|
from typing import Dict, List, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch.library import Library
|
||||||
|
|
||||||
|
from astrai.extension.ops.fp8 import (
|
||||||
|
linear_backward_fp8,
|
||||||
|
linear_forward_fp8,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Max representable value per FP8 format (E4M3: 448, E5M2: 57344).
|
||||||
|
FP8_MAX = {"e4m3": 448.0, "e5m2": 57344.0}
|
||||||
|
|
||||||
|
|
||||||
|
class FP8Format(str, Enum):
|
||||||
|
"""Per-direction FP8 format. HYBRID = E4M3 forward / E5M2 backward."""
|
||||||
|
|
||||||
|
E4M3 = "e4m3"
|
||||||
|
E5M2 = "e5m2"
|
||||||
|
HYBRID = "hybrid"
|
||||||
|
|
||||||
|
def fwd(self) -> str:
|
||||||
|
return "e4m3" if self is FP8Format.HYBRID else self.value
|
||||||
|
|
||||||
|
def bwd(self) -> str:
|
||||||
|
return "e5m2" if self is FP8Format.HYBRID else self.value
|
||||||
|
|
||||||
|
|
||||||
|
class FP8Recipe:
|
||||||
|
"""Scale-from-amax policy: ``scale = (amax / FP8_MAX[fmt]) / 2^margin``.
|
||||||
|
|
||||||
|
``scale_from_history`` receives the operand's amax tensor (a ring window for
|
||||||
|
delayed scaling, the current amax for dynamic scaling) and returns the
|
||||||
|
quantization step. Subclasses set ``history_len`` / ``margin``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
history_len: int = 16
|
||||||
|
margin: int = 0
|
||||||
|
|
||||||
|
def scale_from_history(self, amax: torch.Tensor, fmt: str) -> torch.Tensor:
|
||||||
|
peak = amax.max()
|
||||||
|
return ((peak / FP8_MAX[fmt]) / (2**self.margin)).clamp_min(1e-12)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class DelayedScaling(FP8Recipe):
|
||||||
|
"""TE-style delayed scaling: max over the amax history window (amax from
|
||||||
|
*previous* steps; the window trades responsiveness against stability)."""
|
||||||
|
|
||||||
|
history_len: int = 16
|
||||||
|
margin: int = 0
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class DynamicScaling(FP8Recipe):
|
||||||
|
"""Current-amax scaling (torchao DYNAMIC): measure, then quantize. No
|
||||||
|
history — the scale is derived from the same-step amax, at an extra pass."""
|
||||||
|
|
||||||
|
history_len: int = 1
|
||||||
|
margin: int = 0
|
||||||
|
|
||||||
|
|
||||||
|
class _ScaleRing:
|
||||||
|
"""One operand's delayed-scaling state: a float32 buffer
|
||||||
|
``[hist[n] | scale | counter]`` (views). The quantize kernel's last-finishing
|
||||||
|
block records the measured amax into ``hist[idx]``, reduces the window and
|
||||||
|
publishes the next scale entirely on device — the Python-side write/max/write
|
||||||
|
chain is gone. The counter slot stays int32-zero (float bits) between
|
||||||
|
launches; ``idx`` advances host-side each step.
|
||||||
|
"""
|
||||||
|
|
||||||
|
__slots__ = ("recipe", "state", "hist", "scale", "idx", "initialized")
|
||||||
|
|
||||||
|
def __init__(self, device: torch.device, recipe: FP8Recipe):
|
||||||
|
self.recipe = recipe
|
||||||
|
n = recipe.history_len
|
||||||
|
self.state = torch.zeros(n + 2, device=device, dtype=torch.float32)
|
||||||
|
self.hist = self.state[:n]
|
||||||
|
self.scale = self.state[n : n + 1]
|
||||||
|
self.idx = 0
|
||||||
|
self.initialized = False
|
||||||
|
|
||||||
|
def advance(self) -> None:
|
||||||
|
"""Rotate to the next history slot after an in-kernel finalize."""
|
||||||
|
self.idx = (self.idx + 1) % self.hist.numel()
|
||||||
|
|
||||||
|
def seed(self, t: torch.Tensor, fmt: str) -> None:
|
||||||
|
amax = t.abs().amax().to(torch.float32).clamp_min(1e-12)
|
||||||
|
self.hist.fill_(amax)
|
||||||
|
self.scale.copy_(self.recipe.scale_from_history(self.hist, fmt))
|
||||||
|
self.initialized = True
|
||||||
|
|
||||||
|
|
||||||
|
class FP8TensorMeta:
|
||||||
|
"""Per-weight delayed-scaling state: one ring per operand role (``w``/``x``/
|
||||||
|
``g``). Fused kernels record amax while quantizing, so the scale used at step
|
||||||
|
N reflects amax from steps < N. DynamicScaling never allocates a meta — it
|
||||||
|
measures the current amax inline.
|
||||||
|
"""
|
||||||
|
|
||||||
|
__slots__ = ("w", "x", "g")
|
||||||
|
|
||||||
|
def __init__(self, device: torch.device, recipe: FP8Recipe):
|
||||||
|
self.w = _ScaleRing(device, recipe)
|
||||||
|
self.x = _ScaleRing(device, recipe)
|
||||||
|
self.g = _ScaleRing(device, recipe)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class _ActiveConfig:
|
||||||
|
"""The immutable (enabled, recipe, format) triple of one open region."""
|
||||||
|
|
||||||
|
enabled: bool
|
||||||
|
recipe: FP8Recipe
|
||||||
|
fp8_format: FP8Format
|
||||||
|
|
||||||
|
|
||||||
|
# Thread-local active configuration (torch's autocast TLS analog): set by
|
||||||
|
# fp8_autocast on __enter__, absent outside any region. Autograd engine
|
||||||
|
# threads run backwards with their own empty context — fine, since backward
|
||||||
|
# only reads state captured on ctx at forward time.
|
||||||
|
_active_config: ContextVar[Optional[_ActiveConfig]] = ContextVar(
|
||||||
|
"astrai_fp8_active_config", default=None
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class FP8State:
|
||||||
|
"""Global fp8 training state: per-tensor metas + out-of-region defaults.
|
||||||
|
|
||||||
|
The active ``(enabled, recipe, fp8_format)`` triple is a ``ContextVar`` set
|
||||||
|
by ``fp8_autocast``. The properties below read that active config when a
|
||||||
|
region is open and the global defaults otherwise; the setters (and
|
||||||
|
``fp8_linear_enable``) write the global defaults — the persistent switch
|
||||||
|
applying outside any region. The metas registry is shared across threads
|
||||||
|
(GIL-protected); fp8 backward runs on autograd engine threads and only
|
||||||
|
touches metas captured on ``ctx`` at forward time.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self.default_enabled = False
|
||||||
|
self.default_recipe: FP8Recipe = DelayedScaling()
|
||||||
|
self.default_format: FP8Format = FP8Format.HYBRID
|
||||||
|
self._metas: Dict[tuple, FP8TensorMeta] = {}
|
||||||
|
|
||||||
|
# Active-config views (region config if open, else the defaults).
|
||||||
|
@property
|
||||||
|
def enabled(self) -> bool:
|
||||||
|
cfg = _active_config.get()
|
||||||
|
return cfg.enabled if cfg is not None else self.default_enabled
|
||||||
|
|
||||||
|
@property
|
||||||
|
def recipe(self) -> FP8Recipe:
|
||||||
|
cfg = _active_config.get()
|
||||||
|
return cfg.recipe if cfg is not None else self.default_recipe
|
||||||
|
|
||||||
|
@property
|
||||||
|
def fp8_format(self) -> FP8Format:
|
||||||
|
cfg = _active_config.get()
|
||||||
|
return cfg.fp8_format if cfg is not None else self.default_format
|
||||||
|
|
||||||
|
# Persistent (out-of-region) defaults.
|
||||||
|
@enabled.setter
|
||||||
|
def enabled(self, value: bool) -> None:
|
||||||
|
self.default_enabled = bool(value)
|
||||||
|
|
||||||
|
@recipe.setter
|
||||||
|
def recipe(self, value: FP8Recipe) -> None:
|
||||||
|
self.default_recipe = value
|
||||||
|
|
||||||
|
@fp8_format.setter
|
||||||
|
def fp8_format(self, value: FP8Format) -> None:
|
||||||
|
self.default_format = FP8Format(value)
|
||||||
|
|
||||||
|
def get_weight_meta(self, w: torch.Tensor) -> FP8TensorMeta:
|
||||||
|
key = (w.data_ptr(), w.shape, w.dtype)
|
||||||
|
meta = self._metas.get(key)
|
||||||
|
if meta is None:
|
||||||
|
meta = FP8TensorMeta(w.device, self.recipe)
|
||||||
|
self._metas[key] = meta
|
||||||
|
return meta
|
||||||
|
|
||||||
|
def reset(self) -> None:
|
||||||
|
self.default_enabled = False
|
||||||
|
self._metas.clear()
|
||||||
|
|
||||||
|
|
||||||
|
# Process-wide singleton; per-thread/per-region state lives in _active_config.
|
||||||
|
_state = FP8State()
|
||||||
|
|
||||||
|
|
||||||
|
def fp8_state() -> FP8State:
|
||||||
|
return _state
|
||||||
|
|
||||||
|
|
||||||
|
def _active() -> Optional[_ActiveConfig]:
|
||||||
|
"""The active config when fp8 dispatch is on, else ``None`` (fast guard).
|
||||||
|
|
||||||
|
A region config wins (honoring nested ``enabled=False`` regions); with no
|
||||||
|
region open this falls back to the persistent global switch
|
||||||
|
(``fp8_linear_enable``), so that flag still routes aten::linear to fp8.
|
||||||
|
"""
|
||||||
|
cfg = _active_config.get()
|
||||||
|
if cfg is not None:
|
||||||
|
return cfg if cfg.enabled else None
|
||||||
|
if _state.default_enabled:
|
||||||
|
return _ActiveConfig(True, _state.default_recipe, _state.default_format)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _current_config() -> _ActiveConfig:
|
||||||
|
"""Like ``_active()`` but always returns a config (disabled regions and
|
||||||
|
out-of-region direct calls resolve to the global defaults)."""
|
||||||
|
cfg = _active_config.get()
|
||||||
|
if cfg is not None:
|
||||||
|
return cfg
|
||||||
|
return _ActiveConfig(
|
||||||
|
_state.default_enabled, _state.default_recipe, _state.default_format
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class fp8_autocast:
|
||||||
|
"""Autocast-style context: fp8 linear dispatch on this thread.
|
||||||
|
|
||||||
|
Mirrors ``torch.autocast`` — a class-based, reentrant, nestable context
|
||||||
|
over thread-local state::
|
||||||
|
|
||||||
|
with fp8_autocast(enabled=True, fp8_format="hybrid"):
|
||||||
|
logits = model(input_ids) # aten::linear -> fp8 path
|
||||||
|
loss.backward() # fp8 backward; state was captured at forward time
|
||||||
|
|
||||||
|
Nesting follows torch: each ``__enter__`` pushes the new active config, each
|
||||||
|
``__exit__`` restores the previous one, and a nested ``enabled=False`` region
|
||||||
|
simply disables dispatch inside it. The instance doubles as a decorator.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
enabled: bool = True,
|
||||||
|
update_interval: int = 16,
|
||||||
|
recipe: Optional[FP8Recipe] = None,
|
||||||
|
fp8_format: str = "hybrid",
|
||||||
|
margin: int = 0,
|
||||||
|
):
|
||||||
|
if recipe is None:
|
||||||
|
recipe = DelayedScaling(history_len=update_interval, margin=margin)
|
||||||
|
self._config = _ActiveConfig(bool(enabled), recipe, FP8Format(fp8_format))
|
||||||
|
self._tokens: List[Token] = []
|
||||||
|
|
||||||
|
def __enter__(self) -> "fp8_autocast":
|
||||||
|
self._tokens.append(_active_config.set(self._config))
|
||||||
|
return self
|
||||||
|
|
||||||
|
def __exit__(self, exc_type, exc_val, exc_tb) -> bool:
|
||||||
|
token = self._tokens.pop()
|
||||||
|
_active_config.reset(token)
|
||||||
|
return False
|
||||||
|
|
||||||
|
def __call__(self, func):
|
||||||
|
@functools.wraps(func)
|
||||||
|
def decorate(*args, **kwargs):
|
||||||
|
with self:
|
||||||
|
return func(*args, **kwargs)
|
||||||
|
|
||||||
|
return decorate
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Strategy-level forward / backward (called from the aten::linear impl)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _dynamic_scale(t: torch.Tensor, recipe: FP8Recipe, fmt: str) -> torch.Tensor:
|
||||||
|
amax = t.abs().amax().to(torch.float32).clamp_min(1e-12)
|
||||||
|
return recipe.scale_from_history(amax, fmt)
|
||||||
|
|
||||||
|
|
||||||
|
_zero_bias: Dict[Optional[int], torch.Tensor] = {}
|
||||||
|
|
||||||
|
|
||||||
|
def _empty_bias(x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Per-device cached 0-element bf16 bias (the binding only checks numel —
|
||||||
|
never mutated), saving a CUDA allocation per bias-less linear."""
|
||||||
|
key = x.device.index
|
||||||
|
t = _zero_bias.get(key)
|
||||||
|
if t is None:
|
||||||
|
t = torch.empty(0, device=x.device, dtype=torch.bfloat16)
|
||||||
|
_zero_bias[key] = t
|
||||||
|
return t
|
||||||
|
|
||||||
|
|
||||||
|
def fp8_linear_forward(
|
||||||
|
x: torch.Tensor, w: torch.Tensor, bias=None, cfg: Optional[_ActiveConfig] = None
|
||||||
|
):
|
||||||
|
"""Scaled fp8 linear forward (called from the aten::linear impl).
|
||||||
|
|
||||||
|
Pure FP8 path for both recipes: quantize x/w with the active scales, run the
|
||||||
|
pre-quantized GEMM. Delayed scaling finalizes the rings inside the quantize
|
||||||
|
kernels (amax folded into the window, next scale published on device);
|
||||||
|
dynamic scaling measures the current amax itself. Training quantizes the
|
||||||
|
weight every step (the optimizer bumps its version, so there is no cast
|
||||||
|
cache, matching ``cached_cast``-less behavior).
|
||||||
|
"""
|
||||||
|
state = fp8_state()
|
||||||
|
if cfg is None:
|
||||||
|
cfg = _current_config()
|
||||||
|
fmt = cfg.fp8_format.fwd()
|
||||||
|
margin = cfg.recipe.margin
|
||||||
|
if bias is None:
|
||||||
|
bias = _empty_bias(x)
|
||||||
|
if isinstance(cfg.recipe, DynamicScaling): # measure-then-quantize, no state
|
||||||
|
sx = _dynamic_scale(x.reshape(-1, w.size(1)), cfg.recipe, fmt)
|
||||||
|
sw = _dynamic_scale(w, cfg.recipe, fmt)
|
||||||
|
out, *_ = linear_forward_fp8(x, w, bias, sx, sw, fmt)
|
||||||
|
return out
|
||||||
|
|
||||||
|
meta = state.get_weight_meta(w)
|
||||||
|
if not meta.w.initialized:
|
||||||
|
meta.w.seed(w, fmt)
|
||||||
|
if not meta.x.initialized:
|
||||||
|
meta.x.seed(x, fmt)
|
||||||
|
# The kernel finalizes each ring in-kernel (overwriting the scale slot), so
|
||||||
|
# the w/x scales are taken from the ring before the quantize.
|
||||||
|
if w.dtype is not torch.bfloat16: # static pre-quantized weight
|
||||||
|
w_arg, sw_arg, w_ring = w, meta.w.scale, None
|
||||||
|
else:
|
||||||
|
w_arg, sw_arg, w_ring = w, meta.w.scale, meta.w.state
|
||||||
|
out, _x8, _w8, _ax, _aw = linear_forward_fp8(
|
||||||
|
x,
|
||||||
|
w_arg,
|
||||||
|
bias,
|
||||||
|
meta.x.scale,
|
||||||
|
sw_arg,
|
||||||
|
fmt,
|
||||||
|
None,
|
||||||
|
meta.x.state,
|
||||||
|
meta.x.idx,
|
||||||
|
margin,
|
||||||
|
w_ring,
|
||||||
|
meta.w.idx,
|
||||||
|
margin,
|
||||||
|
)
|
||||||
|
meta.x.advance()
|
||||||
|
if w_ring is not None:
|
||||||
|
meta.w.advance()
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
class _LinearFp8(torch.autograd.Function):
|
||||||
|
"""The fp8 linear forward/backward pair (standard Function style).
|
||||||
|
|
||||||
|
The forward runs inside ``fp8_autocast`` and captures the active
|
||||||
|
fmt/recipe/meta on ``ctx``; the backward reads only that captured state, so
|
||||||
|
``loss.backward()`` may run after the context exits. The gradient is
|
||||||
|
quantized once (E5M2 in hybrid) and both dX/dW GEMMs share it; the output
|
||||||
|
masks come from ``needs_input_grad``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def forward(ctx, x, w, bias):
|
||||||
|
cfg = _current_config()
|
||||||
|
out = fp8_linear_forward(x, w, bias, cfg)
|
||||||
|
ctx.save_for_backward(x, w)
|
||||||
|
ctx.fmt_bwd = cfg.fp8_format.bwd()
|
||||||
|
ctx.recipe = cfg.recipe
|
||||||
|
ctx.is_dynamic = isinstance(cfg.recipe, DynamicScaling)
|
||||||
|
ctx.meta = None if ctx.is_dynamic else _state.get_weight_meta(w)
|
||||||
|
return out
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
@torch.autograd.function.once_differentiable
|
||||||
|
def backward(ctx, g):
|
||||||
|
x, w = ctx.saved_tensors
|
||||||
|
fmt = ctx.fmt_bwd
|
||||||
|
# Per-recipe scale/ring selection; both branches share one call below.
|
||||||
|
if ctx.is_dynamic:
|
||||||
|
sg = _dynamic_scale(g, ctx.recipe, fmt)
|
||||||
|
sw = _dynamic_scale(w, ctx.recipe, fmt)
|
||||||
|
sx = _dynamic_scale(x, ctx.recipe, fmt)
|
||||||
|
ring, idx = None, 0
|
||||||
|
else:
|
||||||
|
meta = ctx.meta
|
||||||
|
if not meta.g.initialized:
|
||||||
|
meta.g.seed(g, fmt)
|
||||||
|
sg, ring, idx = meta.g.scale, meta.g.state, meta.g.idx
|
||||||
|
sw, sx = meta.w.scale, meta.x.scale
|
||||||
|
grad_x, grad_w, grad_b, _amax_g = linear_backward_fp8(
|
||||||
|
g,
|
||||||
|
x,
|
||||||
|
w,
|
||||||
|
list(ctx.needs_input_grad),
|
||||||
|
sg,
|
||||||
|
sw,
|
||||||
|
sx,
|
||||||
|
fmt,
|
||||||
|
ring,
|
||||||
|
idx,
|
||||||
|
ctx.recipe.margin,
|
||||||
|
)
|
||||||
|
if not ctx.is_dynamic:
|
||||||
|
meta.g.advance() # the g quantize kernel finalized the ring in-kernel
|
||||||
|
return grad_x, grad_w, grad_b if ctx.needs_input_grad[2] else None
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# aten::linear integration
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def fp8_linear_enable(enabled: bool = True) -> None:
|
||||||
|
"""Toggle fp8 dispatch for aten::linear globally (the out-of-region default;
|
||||||
|
``fp8_autocast`` regions override it thread-locally)."""
|
||||||
|
fp8_state().default_enabled = enabled
|
||||||
|
|
||||||
|
|
||||||
|
def fp8_linear_enabled() -> bool:
|
||||||
|
"""Whether fp8 dispatch is active right now (region config or global)."""
|
||||||
|
return _active() is not None
|
||||||
|
|
||||||
|
|
||||||
|
def _fp8_supported(x: torch.Tensor, w: torch.Tensor) -> bool:
|
||||||
|
"""Shape guard for the fp8 path. Unlike a strict 16-alignment requirement,
|
||||||
|
the kernels handle unaligned M/N via boundary checks (slower but correct) —
|
||||||
|
so no whole-call bf16 fallback for small decode batches. Only the K-dimension
|
||||||
|
contraction must match and the weight must be 2D."""
|
||||||
|
return x.dim() >= 2 and w.dim() == 2 and x.size(-1) == w.size(1)
|
||||||
|
|
||||||
|
|
||||||
|
def _linear_cuda_impl(x: torch.Tensor, w: torch.Tensor, bias=None):
|
||||||
|
if (
|
||||||
|
_active() is not None
|
||||||
|
and x.dtype is torch.bfloat16
|
||||||
|
and w.dtype is torch.bfloat16
|
||||||
|
and _fp8_supported(x, w)
|
||||||
|
):
|
||||||
|
return _LinearFp8.apply(x, w, bias)
|
||||||
|
return torch.ops.aten.linear.default.redispatch(
|
||||||
|
torch._C.DispatchKeySet(torch._C.DispatchKey.CompositeImplicitAutograd),
|
||||||
|
x,
|
||||||
|
w,
|
||||||
|
bias,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
_lib = Library("aten", "IMPL", "CUDA")
|
||||||
|
_lib.impl("linear", _linear_cuda_impl)
|
||||||
|
# Also replace torch's generated linear autograd formula (which would call
|
||||||
|
# aten::linear_backward after the fp8_autocast region exits). The fp8 backward
|
||||||
|
# is owned by _LinearFp8 with state captured at forward time, so loss.backward()
|
||||||
|
# works wherever it is called; the CUDA registration still covers inference_mode.
|
||||||
|
_lib_autograd = Library("aten", "IMPL", "AutogradCUDA")
|
||||||
|
_lib_autograd.impl("linear", _linear_cuda_impl)
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
"""Compiled CUDA kernel modules (``*.so``) live here, kept separate from Python source."""
|
||||||
@@ -0,0 +1,83 @@
|
|||||||
|
"""Dynamic discovery and loading of compiled CUDA kernel modules.
|
||||||
|
|
||||||
|
Each kernel is built by the CMake build in ``csrc/CMakeLists.txt`` into a
|
||||||
|
``.so`` placed in ``astrai/extension/lib/`` — the module name equals the
|
||||||
|
``.so`` name equals the pybind name (e.g. ``attn_decode``, defined via
|
||||||
|
``TORCH_EXTENSION_NAME``). ``KERNEL_NAMES`` is discovered automatically from
|
||||||
|
the ``.so`` files present, so adding a kernel to the CMake ``KERNELS``
|
||||||
|
registry needs no change here.
|
||||||
|
|
||||||
|
Loading is **lazy and centralized**: module names are discovered eagerly
|
||||||
|
(cheap glob), but each ``.so`` is imported on first use via the single
|
||||||
|
``get_module`` accessor, then cached. The wrapper modules (``ops/*.py``) never
|
||||||
|
touch the internals or keep their own caches — they call ``get_module(name)``
|
||||||
|
(or ``is_available(name)`` when a torch fallback is acceptable). A kernel that
|
||||||
|
failed to build (or is running on a CPU-only machine) is ``None`` in the cache,
|
||||||
|
so ``is_available`` returns ``False`` and ``get_module`` raises a clear error.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import glob
|
||||||
|
import importlib
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_LIB_DIR = os.path.join(os.path.dirname(__file__), "lib")
|
||||||
|
|
||||||
|
|
||||||
|
def _discover_kernel_names() -> list[str]:
|
||||||
|
"""Return the module names of the compiled kernel ``.so`` files in lib/."""
|
||||||
|
names: list[str] = []
|
||||||
|
for path in glob.glob(os.path.join(_LIB_DIR, "*.so")):
|
||||||
|
# strip the "<soabi>.so" suffix, e.g. attn_decode.cpython-312-...so
|
||||||
|
names.append(os.path.basename(path).split(".", 1)[0])
|
||||||
|
return sorted(names)
|
||||||
|
|
||||||
|
|
||||||
|
KERNEL_NAMES = _discover_kernel_names()
|
||||||
|
|
||||||
|
_available: dict[str, bool] = {}
|
||||||
|
_modules: dict[str, object] = {}
|
||||||
|
|
||||||
|
|
||||||
|
def _try_load(name: str) -> object:
|
||||||
|
"""Import and cache the ``name`` kernel module (lazy, one attempt).
|
||||||
|
|
||||||
|
Returns the module, or ``None`` if it is unavailable. Cached so each
|
||||||
|
``.so`` is imported at most once per process.
|
||||||
|
"""
|
||||||
|
if name not in _modules:
|
||||||
|
try:
|
||||||
|
_modules[name] = importlib.import_module(
|
||||||
|
f".lib.{name}", package=__package__
|
||||||
|
)
|
||||||
|
_available[name] = True
|
||||||
|
except ImportError:
|
||||||
|
logger.warning("kernel '%s' failed to import; marking unavailable", name)
|
||||||
|
_modules[name] = None
|
||||||
|
_available[name] = False
|
||||||
|
return _modules[name]
|
||||||
|
|
||||||
|
|
||||||
|
def is_available(name: str) -> bool:
|
||||||
|
"""Return ``True`` if the compiled kernel ``name`` could be loaded."""
|
||||||
|
if name not in _available:
|
||||||
|
_try_load(name)
|
||||||
|
return _available.get(name, False)
|
||||||
|
|
||||||
|
|
||||||
|
def get_module(name: str) -> object:
|
||||||
|
"""Return the loaded kernel module for ``name``, importing it on first use.
|
||||||
|
|
||||||
|
Raises ``RuntimeError`` if the kernel is unavailable (not built, or failed
|
||||||
|
to import) — callers that can tolerate a torch fallback should check
|
||||||
|
``is_available(name)`` first instead.
|
||||||
|
"""
|
||||||
|
mod = _try_load(name)
|
||||||
|
if mod is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"CUDA kernel '{name}' is not available. "
|
||||||
|
f"Build with CSRC_KERNELS=true (or use the torch-native fallback)."
|
||||||
|
)
|
||||||
|
return mod
|
||||||
@@ -0,0 +1,19 @@
|
|||||||
|
"""Stateless wrappers around compiled extension kernels."""
|
||||||
|
|
||||||
|
from astrai.extension.ops.attention import (
|
||||||
|
TensorLayout,
|
||||||
|
attn_decode,
|
||||||
|
attn_paged_decode,
|
||||||
|
attn_paged_prefill,
|
||||||
|
attn_prefill,
|
||||||
|
)
|
||||||
|
from astrai.extension.ops.rotary import rotary_emb
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"TensorLayout",
|
||||||
|
"attn_decode",
|
||||||
|
"attn_paged_decode",
|
||||||
|
"attn_paged_prefill",
|
||||||
|
"attn_prefill",
|
||||||
|
"rotary_emb",
|
||||||
|
]
|
||||||
@@ -0,0 +1,192 @@
|
|||||||
|
"""Attention kernel wrapper functions - one entry point per compiled kernel.
|
||||||
|
|
||||||
|
Each wrapper calls its CUDA kernel directly. If the kernel is not
|
||||||
|
available, raises ``RuntimeError``. Fallback to torch SDPA is the
|
||||||
|
responsibility of the attention backend, not this module.
|
||||||
|
|
||||||
|
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
||||||
|
(blhd). Scale is always ``1/sqrt(head_dim)``.
|
||||||
|
|
||||||
|
Interface (all functions):
|
||||||
|
is_causal: True = causal mask; False = non-causal
|
||||||
|
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
|
||||||
|
"""
|
||||||
|
|
||||||
|
import enum
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.extension.loader import get_module
|
||||||
|
|
||||||
|
|
||||||
|
class TensorLayout(enum.IntEnum):
|
||||||
|
"""Q/K/V tensor layout, mirrors the C++ ``TensorLayout`` enum in ``attn_common.h``.
|
||||||
|
|
||||||
|
Kernels internally operate on BHLD; BLHD inputs are transposed at entry.
|
||||||
|
"""
|
||||||
|
|
||||||
|
BHLD = 0 # [batch, n_heads, seq_len, head_dim]
|
||||||
|
BLHD = 1 # [batch, seq_len, n_heads, head_dim]
|
||||||
|
|
||||||
|
|
||||||
|
def attn_decode(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
mask: Optional[torch.Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""GQA decode attention (q_len == 1).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [batch, 1, n_heads, head_dim] (blhd, bf16)
|
||||||
|
k: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||||
|
v: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||||
|
mask: 2D [batch, kv_len] or 3D [batch, 1, kv_len] (bool, True=keep)
|
||||||
|
is_causal: apply causal mask
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, 1, n_heads, head_dim] (blhd, bf16)
|
||||||
|
"""
|
||||||
|
mod = get_module("attn_decode")
|
||||||
|
causal_offset = (k.size(1) - 1) if is_causal else -1
|
||||||
|
return mod.attn_decode(
|
||||||
|
q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def attn_prefill(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
mask: Optional[torch.Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""GQA prefill attention (q_len > 1).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [batch, q_len, n_heads, head_dim] (blhd, bf16)
|
||||||
|
k: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||||
|
v: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||||
|
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
|
||||||
|
is_causal: apply causal mask
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, q_len, n_heads, head_dim] (blhd, bf16)
|
||||||
|
"""
|
||||||
|
mod = get_module("attn_prefill")
|
||||||
|
causal_offset = (k.size(1) - q.size(1)) if is_causal else -1
|
||||||
|
return mod.attn_prefill(
|
||||||
|
q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def attn_paged_decode(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k_cache: torch.Tensor,
|
||||||
|
v_cache: torch.Tensor,
|
||||||
|
req_to_token: torch.Tensor,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
kv_indptr: torch.Tensor,
|
||||||
|
new_k: Optional[torch.Tensor] = None,
|
||||||
|
new_v: Optional[torch.Tensor] = None,
|
||||||
|
mask: Optional[torch.Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
o_part_buf: Optional[torch.Tensor] = None,
|
||||||
|
ml_part_buf: Optional[torch.Tensor] = None,
|
||||||
|
out_buf: Optional[torch.Tensor] = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""SGLang-style paged decode (q_len == 1, flat KV pool).
|
||||||
|
|
||||||
|
Reads K/V directly from a flat pool [size, kv_head, head_dim] via
|
||||||
|
req_to_token indirect indexing. Each request has its own seq_len
|
||||||
|
(from kv_indptr), eliminating padding waste.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [batch, n_heads, head_dim] (bf16, 3D — no seq dim)
|
||||||
|
k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat)
|
||||||
|
v_cache: same as k_cache
|
||||||
|
req_to_token: [num_reqs, max_context_len] (int32) — token -> slot
|
||||||
|
req_pool_indices: [batch] (int32) — rows into req_to_token
|
||||||
|
kv_indptr: [batch+1] (int32) — prefix sum of per-request seq_lens
|
||||||
|
new_k: current-token K to append, [batch, n_kv_heads, head_dim]
|
||||||
|
new_v: current-token V to append, same shape as new_k
|
||||||
|
mask: 2D [batch, max_context_len] (bool, True=keep) or None
|
||||||
|
is_causal: apply causal mask
|
||||||
|
o_part_buf: pre-allocated split-KV o partial buffer (workflow bypass)
|
||||||
|
ml_part_buf: pre-allocated split-KV m/l buffer (workflow bypass)
|
||||||
|
out_buf: pre-allocated output buffer [batch, n_heads, head_dim] (graph-safe)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, n_heads, head_dim] (bf16, 3D)
|
||||||
|
"""
|
||||||
|
mod = get_module("attn_paged_decode")
|
||||||
|
causal_offset = 0 if is_causal else -1
|
||||||
|
return mod.attn_paged_decode(
|
||||||
|
q,
|
||||||
|
k_cache,
|
||||||
|
v_cache,
|
||||||
|
req_to_token,
|
||||||
|
req_pool_indices,
|
||||||
|
kv_indptr,
|
||||||
|
new_k=new_k,
|
||||||
|
new_v=new_v,
|
||||||
|
mask=mask,
|
||||||
|
causal_offset=causal_offset,
|
||||||
|
o_part_buf=o_part_buf,
|
||||||
|
ml_part_buf=ml_part_buf,
|
||||||
|
out_buf=out_buf,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def attn_paged_prefill(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k_cache: torch.Tensor,
|
||||||
|
v_cache: torch.Tensor,
|
||||||
|
req_to_token: torch.Tensor,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
kv_indptr: torch.Tensor,
|
||||||
|
qo_indptr: torch.Tensor,
|
||||||
|
q_tile_to_batch: torch.Tensor,
|
||||||
|
q_tile_to_index: torch.Tensor,
|
||||||
|
mask: Optional[torch.Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""SGLang-style paged prefill (ragged batch, flat KV pool).
|
||||||
|
|
||||||
|
Reads K/V directly from a flat pool [size, kv_head, head_dim] via
|
||||||
|
req_to_token. Supports ragged batches: each request has its own
|
||||||
|
q_len and kv_len, addressed via qo_indptr and kv_indptr.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [total_q, n_heads, head_dim] (bf16, 3D — flattened across requests)
|
||||||
|
k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat)
|
||||||
|
v_cache: same as k_cache
|
||||||
|
req_to_token: [num_reqs, max_context_len] (int32)
|
||||||
|
req_pool_indices: [batch] (int32)
|
||||||
|
kv_indptr: [batch+1] (int32) — prefix sum of per-request kv_lens
|
||||||
|
qo_indptr: [batch+1] (int32) — prefix sum of per-request q_lens
|
||||||
|
q_tile_to_batch: [num_q_tiles] (int32) — request index per Q tile
|
||||||
|
q_tile_to_index: [num_q_tiles] (int32) — local Q tile index per request
|
||||||
|
mask: 4D [batch, 1, q_len, kv_len] (bool, True=keep) or None
|
||||||
|
is_causal: apply causal mask
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[total_q, n_heads, head_dim] (bf16, 3D)
|
||||||
|
"""
|
||||||
|
mod = get_module("attn_paged_prefill")
|
||||||
|
causal_offset = 0 if is_causal else -1
|
||||||
|
return mod.attn_paged_prefill(
|
||||||
|
q,
|
||||||
|
k_cache,
|
||||||
|
v_cache,
|
||||||
|
req_to_token,
|
||||||
|
req_pool_indices,
|
||||||
|
kv_indptr,
|
||||||
|
qo_indptr,
|
||||||
|
q_tile_to_batch,
|
||||||
|
q_tile_to_index,
|
||||||
|
mask,
|
||||||
|
causal_offset=causal_offset,
|
||||||
|
)
|
||||||
@@ -0,0 +1,241 @@
|
|||||||
|
"""FP8 CUDA kernel interface adapter (the only module touching the pybind).
|
||||||
|
|
||||||
|
Isolates the ``fp8_ops`` CUDA extension behind stable Python primitives:
|
||||||
|
|
||||||
|
- ``quantize_bf16(x, scale, fmt) -> (x8, amax)`` — BF16 → FP8 with fused amax
|
||||||
|
- ``mm_fp8(a8, b8, sa, sb) -> out`` — pre-quantized FP8 GEMM (BF16 output)
|
||||||
|
- ``linear_forward_fp8(x, w, bias, sx, sw) -> (out, x8, w8, amax_x, amax_w)``
|
||||||
|
- ``linear_backward_fp8(g, x, w, masks, sg, sw, sx, fmt) -> (gx, gw, gb, amax_g)``
|
||||||
|
|
||||||
|
Scale semantics: scales are *quantization steps* — the value divided out when
|
||||||
|
quantizing (``x8 = x / scale``). Every primitive computes its own inverse
|
||||||
|
internally; callers never pass ``scale_inv``. ``amax`` values are *returned*,
|
||||||
|
never passed as output arguments. ``fmt`` is ``"e4m3"`` or ``"e5m2"``.
|
||||||
|
|
||||||
|
Policy (scales, amax history, delayed scaling, autocast) lives in ``fp8.py``;
|
||||||
|
this module is stateless.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import List, Optional, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch.library import custom_op
|
||||||
|
|
||||||
|
from astrai.extension.loader import get_module
|
||||||
|
|
||||||
|
# fmt string -> kernel int (0 = E4M3, 1 = E5M2)
|
||||||
|
_FMT_TO_INT = {"e4m3": 0, "e5m2": 1}
|
||||||
|
|
||||||
|
|
||||||
|
def _fmt_int(fmt: str) -> int:
|
||||||
|
try:
|
||||||
|
return _FMT_TO_INT[fmt]
|
||||||
|
except KeyError:
|
||||||
|
raise ValueError(f"unsupported fp8 format {fmt!r} (expected 'e4m3' or 'e5m2')")
|
||||||
|
|
||||||
|
|
||||||
|
def _fmt_dtype(fmt: str) -> torch.dtype:
|
||||||
|
return torch.float8_e5m2 if _fmt_int(fmt) else torch.float8_e4m3fn
|
||||||
|
|
||||||
|
|
||||||
|
@custom_op("custom::fp8_quantize", mutates_args=())
|
||||||
|
def fp8_quantize(
|
||||||
|
x: torch.Tensor, scale: torch.Tensor, fmt: int
|
||||||
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""BF16 -> FP8 quantize with fused amax; returns ``(x8, amax)``."""
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_quantize.register_fake
|
||||||
|
def _fp8_quantize_fake(x, scale, fmt):
|
||||||
|
dtype = torch.float8_e5m2 if fmt else torch.float8_e4m3fn
|
||||||
|
return (
|
||||||
|
torch.empty(x.shape, device=x.device, dtype=dtype),
|
||||||
|
torch.empty(1, device=x.device, dtype=torch.float32),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_quantize.register_kernel("cuda")
|
||||||
|
def _fp8_quantize_cuda(x, scale, fmt):
|
||||||
|
if x.dtype != torch.bfloat16:
|
||||||
|
raise TypeError(f"fp8 quantize requires bf16 input, got {x.dtype}")
|
||||||
|
return get_module("fp8_ops").quantize_bf16(x, scale, int(fmt))
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_quantize.register_kernel("cpu")
|
||||||
|
def _fp8_quantize_cpu(x, scale, fmt):
|
||||||
|
x8 = (x.float() / scale).to(_fmt_dtype("e5m2" if fmt else "e4m3"))
|
||||||
|
amax = x.abs().amax().float().reshape(1).clamp_min(1e-12)
|
||||||
|
return x8, amax
|
||||||
|
|
||||||
|
|
||||||
|
@custom_op("custom::fp8_gemm", mutates_args=())
|
||||||
|
def fp8_gemm(
|
||||||
|
a: torch.Tensor,
|
||||||
|
b: torch.Tensor,
|
||||||
|
sa: torch.Tensor,
|
||||||
|
sb: torch.Tensor,
|
||||||
|
out_dtype: int = 0,
|
||||||
|
out_scale: Optional[torch.Tensor] = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""FP8 GEMM: ``a @ b * (sa * sb)`` with FP32 accumulation.
|
||||||
|
|
||||||
|
``out_dtype``: 0 = BF16 (default), 1 = FP8 E4M3 (requires ``out_scale``,
|
||||||
|
the quantization step for the output — mirrors ``torch._scaled_mm``).
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_gemm.register_fake
|
||||||
|
def _fp8_gemm_fake(a, b, sa, sb, out_dtype=0, out_scale=None):
|
||||||
|
dtype = torch.float8_e4m3fn if out_dtype else torch.bfloat16
|
||||||
|
return torch.empty((a.size(0), b.size(1)), device=a.device, dtype=dtype)
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_gemm.register_kernel("cuda")
|
||||||
|
def _fp8_gemm_cuda(a, b, sa, sb, out_dtype=0, out_scale=None):
|
||||||
|
if a.dtype != b.dtype or a.dtype not in (torch.float8_e4m3fn, torch.float8_e5m2):
|
||||||
|
raise TypeError(
|
||||||
|
f"fp8 GEMM requires matching fp8 inputs, got {a.dtype}/{b.dtype}"
|
||||||
|
)
|
||||||
|
return get_module("fp8_ops").mm_fp8(a, b, sa, sb, int(out_dtype), out_scale)
|
||||||
|
|
||||||
|
|
||||||
|
@fp8_gemm.register_kernel("cpu")
|
||||||
|
def _fp8_gemm_cpu(a, b, sa, sb, out_dtype=0, out_scale=None):
|
||||||
|
acc = a.float() @ b.float() * sa * sb
|
||||||
|
if out_dtype:
|
||||||
|
os_ = 1.0 if out_scale is None else out_scale
|
||||||
|
return (acc * os_).to(torch.float8_e4m3fn)
|
||||||
|
return acc.to(torch.bfloat16)
|
||||||
|
|
||||||
|
|
||||||
|
def quantize_bf16(
|
||||||
|
x: torch.Tensor, scale: torch.Tensor, fmt: str = "e4m3"
|
||||||
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""BF16 -> FP8 quantize with fused amax; returns ``(x8, amax)``.
|
||||||
|
|
||||||
|
``scale`` is the quantization step (device scalar); ``fmt`` selects
|
||||||
|
E4M3 or E5M2. ``amax`` is a fresh 1-element float32 tensor — the caller
|
||||||
|
never clears it.
|
||||||
|
"""
|
||||||
|
return fp8_quantize(x, scale, _fmt_int(fmt))
|
||||||
|
|
||||||
|
|
||||||
|
def mm_fp8(
|
||||||
|
a: torch.Tensor,
|
||||||
|
b: torch.Tensor,
|
||||||
|
sa: torch.Tensor,
|
||||||
|
sb: torch.Tensor,
|
||||||
|
out_dtype: str = "bf16",
|
||||||
|
out_scale: Optional[torch.Tensor] = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Pre-quantized FP8 GEMM: ``a @ b * (sa * sb)``.
|
||||||
|
|
||||||
|
``a``/``b`` must be FP8 tensors of the same format (E4M3 or E5M2);
|
||||||
|
``sa``/``sb`` are their quantization steps. ``out_dtype`` is ``"bf16"``
|
||||||
|
(default) or ``"e4m3"`` — FP8 output for layer-to-layer pipelines, which
|
||||||
|
requires ``out_scale`` (the output quantization step).
|
||||||
|
"""
|
||||||
|
if out_dtype not in ("bf16", "e4m3"):
|
||||||
|
raise ValueError(
|
||||||
|
f"unsupported out_dtype {out_dtype!r} (expected 'bf16' or 'e4m3')"
|
||||||
|
)
|
||||||
|
return fp8_gemm(a, b, sa, sb, int(out_dtype == "e4m3"), out_scale)
|
||||||
|
|
||||||
|
|
||||||
|
def linear_forward_fp8(
|
||||||
|
x: torch.Tensor,
|
||||||
|
w: torch.Tensor,
|
||||||
|
bias: Optional[torch.Tensor],
|
||||||
|
sx: torch.Tensor,
|
||||||
|
sw: torch.Tensor,
|
||||||
|
fmt: str = "e4m3",
|
||||||
|
bias_scale: Optional[torch.Tensor] = None,
|
||||||
|
x_ring: Optional[torch.Tensor] = None,
|
||||||
|
x_ring_idx: int = 0,
|
||||||
|
x_ring_margin: int = 0,
|
||||||
|
w_ring: Optional[torch.Tensor] = None,
|
||||||
|
w_ring_idx: int = 0,
|
||||||
|
w_ring_margin: int = 0,
|
||||||
|
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||||
|
"""Pure FP8 linear forward: quantize x/w to ``fmt``, pre-quantized GEMM.
|
||||||
|
|
||||||
|
Returns ``(out, x8, w8, amax_x, amax_w)`` — the quantized operands are
|
||||||
|
handed back so the policy layer can cache the weight quantization while
|
||||||
|
the weight tensor is unchanged (torch autocast's cached_cast analog).
|
||||||
|
``x8`` is ``[M, K]`` and ``w8`` is ``[N, K]`` (the passed-in ``w`` itself
|
||||||
|
on the pre-quantized path). ``bias`` may be ``None``. For static fp8
|
||||||
|
inference, ``w`` and ``bias`` may arrive pre-quantized to ``fmt``
|
||||||
|
(produced by :func:`quantize_bf16` with their scales as ``sw`` /
|
||||||
|
``bias_scale``); a pre-quantized ``bias`` requires ``bias_scale``, and
|
||||||
|
its ``amax_w`` comes back 0. The bias is fused into the GEMM epilogue.
|
||||||
|
``x_ring`` / ``w_ring`` (delayed scaling) are ``[hist | scale | counter]``
|
||||||
|
float32 buffers the quantize kernels finalize in-kernel: the measured
|
||||||
|
amax lands in ``hist[idx]`` and the next step's scale is published on
|
||||||
|
device, replacing the eager hist/max/scale update chain.
|
||||||
|
"""
|
||||||
|
fmt8 = _fmt_dtype(fmt)
|
||||||
|
if x.dtype != torch.bfloat16 or w.dtype not in (torch.bfloat16, fmt8):
|
||||||
|
raise TypeError(
|
||||||
|
f"fp8 forward requires bf16 x and bf16-or-{fmt} w, got {x.dtype}/{w.dtype}"
|
||||||
|
)
|
||||||
|
if bias is None:
|
||||||
|
bias = torch.empty(0, device=x.device, dtype=x.dtype)
|
||||||
|
return get_module("fp8_ops").linear_forward_fp8(
|
||||||
|
x,
|
||||||
|
w,
|
||||||
|
bias,
|
||||||
|
sx,
|
||||||
|
sw,
|
||||||
|
_fmt_int(fmt),
|
||||||
|
bias_scale,
|
||||||
|
x_ring,
|
||||||
|
x_ring_idx,
|
||||||
|
x_ring_margin,
|
||||||
|
w_ring,
|
||||||
|
w_ring_idx,
|
||||||
|
w_ring_margin,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def linear_backward_fp8(
|
||||||
|
g: torch.Tensor,
|
||||||
|
x: torch.Tensor,
|
||||||
|
w: torch.Tensor,
|
||||||
|
masks: List[bool],
|
||||||
|
sg: torch.Tensor,
|
||||||
|
sw: torch.Tensor,
|
||||||
|
sx: torch.Tensor,
|
||||||
|
fmt: str = "e5m2",
|
||||||
|
g_ring: Optional[torch.Tensor] = None,
|
||||||
|
g_ring_idx: int = 0,
|
||||||
|
g_ring_margin: int = 0,
|
||||||
|
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||||
|
"""FP8 linear backward; returns ``(grad_input, grad_weight, grad_bias, amax_g)``.
|
||||||
|
|
||||||
|
The gradient (and the transposed w/x operands) are quantized to ``fmt``
|
||||||
|
(default E5M2 — larger dynamic range for gradients) and the two GEMMs run
|
||||||
|
as FP8 tensor-core products sharing a single gradient quantization.
|
||||||
|
``g_ring`` (delayed scaling) is a ``[hist | scale | counter]`` buffer the
|
||||||
|
g quantize kernel finalizes in-kernel (see :func:`linear_forward_fp8`).
|
||||||
|
"""
|
||||||
|
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 get_module("fp8_ops").linear_backward_fp8(
|
||||||
|
g,
|
||||||
|
x,
|
||||||
|
w,
|
||||||
|
list(masks),
|
||||||
|
sg,
|
||||||
|
sw,
|
||||||
|
sx,
|
||||||
|
_fmt_int(fmt),
|
||||||
|
g_ring,
|
||||||
|
g_ring_idx,
|
||||||
|
g_ring_margin,
|
||||||
|
)
|
||||||
@@ -0,0 +1,31 @@
|
|||||||
|
"""Rotary embedding CUDA kernel wrapper.
|
||||||
|
|
||||||
|
Calls the compiled CUDA kernel directly. If the kernel is not available,
|
||||||
|
raises ``RuntimeError``. Fallback to torch complex multiply is the
|
||||||
|
responsibility of ``astrai.extension.backend.rotary.apply_rotary_emb``.
|
||||||
|
|
||||||
|
Layout: x is packed [tokens, n_heads, head_dim] or dense
|
||||||
|
[batch, seq_len, n_heads, head_dim]. ``freqs_cis`` has matching token axes.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.extension.loader import get_module
|
||||||
|
|
||||||
|
|
||||||
|
def rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Fused rotary embedding kernel.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
x: packed 3D or dense 4D bf16 tensor.
|
||||||
|
freqs_cis: matching token axes followed by [head_dim/2, 2].
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tensor with the same shape as ``x``.
|
||||||
|
"""
|
||||||
|
mod = get_module("rotary_emb")
|
||||||
|
if not x.is_contiguous():
|
||||||
|
x = x.contiguous()
|
||||||
|
if not freqs_cis.is_contiguous():
|
||||||
|
freqs_cis = freqs_cis.contiguous()
|
||||||
|
return mod.rotary_emb(x, freqs_cis)
|
||||||
@@ -0,0 +1,153 @@
|
|||||||
|
"""Base factory with decorator-based registration and kwarg-filtered instantiation."""
|
||||||
|
|
||||||
|
import inspect
|
||||||
|
import sys
|
||||||
|
from abc import ABC
|
||||||
|
from typing import (
|
||||||
|
Callable,
|
||||||
|
Dict,
|
||||||
|
ForwardRef,
|
||||||
|
Generic,
|
||||||
|
List,
|
||||||
|
Optional,
|
||||||
|
Type,
|
||||||
|
TypeVar,
|
||||||
|
Union,
|
||||||
|
get_args,
|
||||||
|
get_origin,
|
||||||
|
)
|
||||||
|
|
||||||
|
T = TypeVar("T")
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_base_type(
|
||||||
|
arg: Union[Type, str, ForwardRef], factory_cls: type
|
||||||
|
) -> Optional[Type]:
|
||||||
|
"""Resolve the generic type-arg T to a concrete class.
|
||||||
|
|
||||||
|
- Concrete class (``BaseFactory[MyBase]``): returned directly.
|
||||||
|
- Forward reference (``BaseFactory["MyBase"]``): ``Base["X"]``
|
||||||
|
produces a ``ForwardRef("X")`` at class-creation time. We
|
||||||
|
extract the name and evaluate it in the factory module's
|
||||||
|
global namespace — the same mechanism ``typing.get_type_hints``
|
||||||
|
uses internally.
|
||||||
|
"""
|
||||||
|
if isinstance(arg, type):
|
||||||
|
return arg
|
||||||
|
|
||||||
|
if isinstance(arg, str):
|
||||||
|
name = arg
|
||||||
|
elif isinstance(arg, ForwardRef):
|
||||||
|
name = arg.__forward_arg__
|
||||||
|
else:
|
||||||
|
return None
|
||||||
|
|
||||||
|
mod = sys.modules.get(factory_cls.__module__)
|
||||||
|
if mod is None:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
return eval(name, vars(mod)) # noqa: S307
|
||||||
|
except NameError:
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_component(component_cls: Type, base: Optional[Type]) -> None:
|
||||||
|
"""Validate that *component_cls* inherits from *base*.
|
||||||
|
|
||||||
|
No-op when *base* is ``None`` (e.g. forward-ref resolution failed).
|
||||||
|
"""
|
||||||
|
if base is not None and not issubclass(component_cls, base):
|
||||||
|
raise TypeError(f"{component_cls.__name__} must inherit from {base.__name__}")
|
||||||
|
|
||||||
|
|
||||||
|
class BaseFactory(ABC, Generic[T]):
|
||||||
|
"""Generic factory with decorator-based registration.
|
||||||
|
|
||||||
|
Create a factory by subclassing with the desired base type::
|
||||||
|
|
||||||
|
class MyFactory(BaseFactory[MyBase]):
|
||||||
|
pass
|
||||||
|
|
||||||
|
Register components with the ``register`` decorator::
|
||||||
|
|
||||||
|
@MyFactory.register("custom")
|
||||||
|
class CustomComponent(MyBase):
|
||||||
|
...
|
||||||
|
|
||||||
|
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.
|
||||||
|
"""
|
||||||
|
|
||||||
|
_entries: Dict[str, Type[T]]
|
||||||
|
|
||||||
|
def __init_subclass__(cls, **kwargs):
|
||||||
|
super().__init_subclass__(**kwargs)
|
||||||
|
for orig_base in getattr(cls, "__orig_bases__", ()):
|
||||||
|
if get_origin(orig_base) is BaseFactory:
|
||||||
|
(arg,) = get_args(orig_base)
|
||||||
|
cls._entries = {}
|
||||||
|
cls._component_base = _resolve_base_type(arg, cls)
|
||||||
|
return
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def register(cls, name: str) -> Callable[[Type[T]], Type[T]]:
|
||||||
|
"""Decorator to register a component class.
|
||||||
|
|
||||||
|
Validates that the decorated class inherits from the generic
|
||||||
|
type parameter ``T`` declared on the factory.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def decorator(component_cls: Type[T]) -> Type[T]:
|
||||||
|
_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, filtering kwargs to match
|
||||||
|
the component's ``__init__`` signature.
|
||||||
|
"""
|
||||||
|
component_cls = cls._entries.get(name)
|
||||||
|
if component_cls is None:
|
||||||
|
raise ValueError(
|
||||||
|
f"Unknown component: '{name}'. Supported types: {sorted(cls._entries)}"
|
||||||
|
)
|
||||||
|
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 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[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."""
|
||||||
|
return name in cls._entries
|
||||||
@@ -0,0 +1,33 @@
|
|||||||
|
"""Inference module for continuous batching.
|
||||||
|
|
||||||
|
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 InferenceEngine
|
||||||
|
from astrai.inference.network import get_app, run_server
|
||||||
|
from astrai.inference.runtime.executor import Executor
|
||||||
|
from astrai.inference.runtime.sample import sample
|
||||||
|
from astrai.inference.scheduler import InferenceScheduler
|
||||||
|
from astrai.inference.task import STOP, Task, TaskManager, TaskStatus
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"InferenceEngine",
|
||||||
|
"InferenceScheduler",
|
||||||
|
"Executor",
|
||||||
|
"STOP",
|
||||||
|
"Task",
|
||||||
|
"TaskManager",
|
||||||
|
"TaskStatus",
|
||||||
|
"sample",
|
||||||
|
"get_app",
|
||||||
|
"run_server",
|
||||||
|
]
|
||||||
Vendored
+27
@@ -0,0 +1,27 @@
|
|||||||
|
"""KV cache subsystem: buffers, strategies, pool management."""
|
||||||
|
|
||||||
|
from astrai.inference.cache.buffer import KVCache, KVStorage, ReqToTokenPool
|
||||||
|
from astrai.inference.cache.pool import PagePool, TaskCacheManager, page_hash
|
||||||
|
from astrai.inference.cache.strategy import (
|
||||||
|
AllocationStrategy,
|
||||||
|
Allocator,
|
||||||
|
ContiguousStrategy,
|
||||||
|
PagedStrategy,
|
||||||
|
RadixCache,
|
||||||
|
TaskCacheState,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"KVCache",
|
||||||
|
"KVStorage",
|
||||||
|
"ReqToTokenPool",
|
||||||
|
"Allocator",
|
||||||
|
"RadixCache",
|
||||||
|
"TaskCacheState",
|
||||||
|
"AllocationStrategy",
|
||||||
|
"ContiguousStrategy",
|
||||||
|
"PagedStrategy",
|
||||||
|
"PagePool",
|
||||||
|
"TaskCacheManager",
|
||||||
|
"page_hash",
|
||||||
|
]
|
||||||
Vendored
+96
@@ -0,0 +1,96 @@
|
|||||||
|
"""Physical KV cache buffers.
|
||||||
|
|
||||||
|
Layer 1 — ``KVStorage``: flat token-level K/V GPU buffers [n_layers, size, n_kv_heads, head_dim]
|
||||||
|
Layer 2 — ``ReqToTokenPool``: index table [req_idx, pos] → physical token slot
|
||||||
|
Layer 3 — ``KVCache``: pure dataclass passed to the model for direct buffer access
|
||||||
|
|
||||||
|
These classes have no knowledge of tasks, allocation policies, or scheduling.
|
||||||
|
They are the "dumb" physical storage layer.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import threading
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import List, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
class ReqToTokenPool:
|
||||||
|
"""Maps [req_idx, pos] → physical token slot in KV storage.
|
||||||
|
|
||||||
|
Each row is one request; each column is a sequence position. The value
|
||||||
|
at [req_idx, pos] is the flat index into the KV storage buffers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, size: int, max_context_len: int, device: torch.device):
|
||||||
|
self.size = size
|
||||||
|
self.max_context_len = max_context_len
|
||||||
|
self.req_to_token = torch.zeros(
|
||||||
|
(size, max_context_len), dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
self.free_slots = list(range(size))
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
def alloc(self, num_reqs: int) -> Optional[List[int]]:
|
||||||
|
with self._lock:
|
||||||
|
if num_reqs > len(self.free_slots):
|
||||||
|
return None
|
||||||
|
slots = self.free_slots[:num_reqs]
|
||||||
|
self.free_slots = self.free_slots[num_reqs:]
|
||||||
|
return slots
|
||||||
|
|
||||||
|
def free(self, req_indices: List[int]):
|
||||||
|
with self._lock:
|
||||||
|
self.free_slots.extend(req_indices)
|
||||||
|
|
||||||
|
def write(self, indices, values):
|
||||||
|
self.req_to_token[indices] = values
|
||||||
|
|
||||||
|
|
||||||
|
class KVStorage:
|
||||||
|
"""Token-level KV cache storage.
|
||||||
|
|
||||||
|
Buffers: ``[n_layers, size, n_kv_heads, head_dim]``. Each token occupies
|
||||||
|
one slot indexed by ``ReqToTokenPool``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
size: int,
|
||||||
|
n_layers: int,
|
||||||
|
n_kv_heads: int,
|
||||||
|
head_dim: int,
|
||||||
|
device: torch.device,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
):
|
||||||
|
self.size = size
|
||||||
|
self.k_buffer = torch.empty(
|
||||||
|
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
|
||||||
|
)
|
||||||
|
self.v_buffer = torch.empty(
|
||||||
|
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class KVCache:
|
||||||
|
"""Pure data struct passed to model for KV cache I/O.
|
||||||
|
|
||||||
|
The attention layer does raw buffer indexing — no methods, no abstraction.
|
||||||
|
"""
|
||||||
|
|
||||||
|
k_buffer: Tensor
|
||||||
|
v_buffer: Tensor
|
||||||
|
req_to_token: Tensor
|
||||||
|
req_pool_indices: Tensor
|
||||||
|
seq_lens: Tensor
|
||||||
|
out_cache_loc: Tensor
|
||||||
|
max_len: int = 0
|
||||||
|
kv_indptr: Optional[Tensor] = None
|
||||||
|
qo_indptr: Optional[Tensor] = None
|
||||||
|
q_tile_to_batch: Optional[Tensor] = None
|
||||||
|
q_tile_to_index: Optional[Tensor] = None
|
||||||
|
decode_o_part: Optional[Tensor] = None
|
||||||
|
decode_ml_part: Optional[Tensor] = None
|
||||||
|
decode_out: Optional[Tensor] = None
|
||||||
Vendored
+382
@@ -0,0 +1,382 @@
|
|||||||
|
"""KV cache orchestration: PagePool + TaskCacheManager.
|
||||||
|
|
||||||
|
PagePool owns the physical buffers (``KVStorage`` + ``ReqToTokenPool``)
|
||||||
|
and wires them to an allocation strategy. It assembles the ``KVCache``
|
||||||
|
dataclass passed to the model forward.
|
||||||
|
|
||||||
|
TaskCacheManager owns the ``task_id`` → ``TaskCacheState`` mapping and
|
||||||
|
delegates physical slot allocation to the strategy, and KV bind to the pool.
|
||||||
|
|
||||||
|
See ``cache_buffer.py`` for the raw buffer primitives and ``cache_strategy.py``
|
||||||
|
for the allocation policies.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Dict, List, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.inference.cache.buffer import KVCache, KVStorage, ReqToTokenPool
|
||||||
|
from astrai.inference.cache.strategy import (
|
||||||
|
AllocationStrategy,
|
||||||
|
Allocator,
|
||||||
|
ContiguousStrategy,
|
||||||
|
PagedStrategy,
|
||||||
|
RadixCache,
|
||||||
|
TaskCacheState,
|
||||||
|
)
|
||||||
|
from astrai.inference.workspace import Q_TILE_ROWS, InferenceWorkspace
|
||||||
|
|
||||||
|
# Re-export everything so existing ``from astrai.inference.cache import ...``
|
||||||
|
# continues to work unchanged after the file split.
|
||||||
|
__all__ = [
|
||||||
|
"KVCache",
|
||||||
|
"KVStorage",
|
||||||
|
"ReqToTokenPool",
|
||||||
|
"Allocator",
|
||||||
|
"RadixCache",
|
||||||
|
"AllocationStrategy",
|
||||||
|
"ContiguousStrategy",
|
||||||
|
"PagedStrategy",
|
||||||
|
"PagePool",
|
||||||
|
"TaskCacheManager",
|
||||||
|
"TaskCacheState",
|
||||||
|
"page_hash",
|
||||||
|
]
|
||||||
|
|
||||||
|
# ---- helpers ----
|
||||||
|
|
||||||
|
|
||||||
|
def page_hash(
|
||||||
|
token_ids: List[int], page_idx: int, page_size: int, parent_hash: int = 0
|
||||||
|
) -> int:
|
||||||
|
start = page_idx * page_size
|
||||||
|
end = min(start + page_size, len(token_ids))
|
||||||
|
h = parent_hash
|
||||||
|
for i in range(start, end):
|
||||||
|
h = (h * 31 + token_ids[i]) & 0xFFFFFFFFFFFFFFFF
|
||||||
|
return h
|
||||||
|
|
||||||
|
|
||||||
|
def _is_steady_increment(
|
||||||
|
prev_sig: Optional[tuple],
|
||||||
|
prev_vals: Optional[List[int]],
|
||||||
|
cur_sig: tuple,
|
||||||
|
cur_vals: List[int],
|
||||||
|
) -> bool:
|
||||||
|
return (
|
||||||
|
prev_sig is not None
|
||||||
|
and prev_vals is not None
|
||||||
|
and prev_sig == cur_sig
|
||||||
|
and len(prev_vals) == len(cur_vals)
|
||||||
|
and all(c == p + 1 for c, p in zip(cur_vals, prev_vals))
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ---- task-scoped bind state ----
|
||||||
|
@dataclass
|
||||||
|
class _BindState:
|
||||||
|
"""Cached bind metadata for steady-state decode increment detection."""
|
||||||
|
|
||||||
|
sig: tuple
|
||||||
|
seq_lens: List[int]
|
||||||
|
|
||||||
|
|
||||||
|
# ---- pool + manager ----
|
||||||
|
|
||||||
|
|
||||||
|
class PagePool:
|
||||||
|
"""Physical KV cache: buffers + req-to-token table + allocation strategy + bind.
|
||||||
|
|
||||||
|
Does not know about tasks — task lifecycle is managed by
|
||||||
|
:class:`TaskCacheManager`, which holds a reference to this pool.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
n_layers: int,
|
||||||
|
n_kv_heads: int,
|
||||||
|
head_dim: int,
|
||||||
|
max_batch_size: int,
|
||||||
|
max_seq_len: int,
|
||||||
|
device: torch.device,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
page_size: int = 1,
|
||||||
|
n_tokens: Optional[int] = None,
|
||||||
|
):
|
||||||
|
self.page_size = page_size
|
||||||
|
self.max_batch_size = max_batch_size
|
||||||
|
self.max_seq_len = max_seq_len
|
||||||
|
self.device = device
|
||||||
|
self.dtype = dtype
|
||||||
|
self.n_layers = n_layers
|
||||||
|
self.n_kv_heads = n_kv_heads
|
||||||
|
self.head_dim = head_dim
|
||||||
|
|
||||||
|
self.contiguous = n_tokens is None
|
||||||
|
self.n_tokens = max_batch_size * max_seq_len if self.contiguous else n_tokens
|
||||||
|
if self.n_tokens > torch.iinfo(torch.int32).max:
|
||||||
|
raise ValueError("KV cache token count exceeds the int32 slot index limit")
|
||||||
|
|
||||||
|
self._storage = KVStorage(
|
||||||
|
self.n_tokens, n_layers, n_kv_heads, head_dim, device, dtype
|
||||||
|
)
|
||||||
|
self._req_pool = ReqToTokenPool(max_batch_size, max_seq_len, device)
|
||||||
|
|
||||||
|
if self.contiguous:
|
||||||
|
for i in range(max_batch_size):
|
||||||
|
self._req_pool.req_to_token[i] = torch.arange(
|
||||||
|
i * max_seq_len,
|
||||||
|
(i + 1) * max_seq_len,
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
self._strategy: AllocationStrategy = ContiguousStrategy()
|
||||||
|
else:
|
||||||
|
n_pages = self.n_tokens // page_size
|
||||||
|
alloc = Allocator(n_pages)
|
||||||
|
prefix = RadixCache(page_size) if page_size > 1 else None
|
||||||
|
if prefix is not None:
|
||||||
|
alloc.on_evict = prefix.evict
|
||||||
|
self._strategy = PagedStrategy(
|
||||||
|
alloc, prefix, page_size, self._req_pool, device
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def strategy(self) -> AllocationStrategy:
|
||||||
|
return self._strategy
|
||||||
|
|
||||||
|
@property
|
||||||
|
def req_pool(self) -> ReqToTokenPool:
|
||||||
|
return self._req_pool
|
||||||
|
|
||||||
|
def bind_tasks(
|
||||||
|
self,
|
||||||
|
req_indices: List[int],
|
||||||
|
seq_lens: List[int],
|
||||||
|
workspace: InferenceWorkspace,
|
||||||
|
device: Optional[torch.device] = None,
|
||||||
|
start_pos: Optional[int] = None,
|
||||||
|
incremental: bool = False,
|
||||||
|
) -> KVCache:
|
||||||
|
"""Assemble the ``KVCache`` metadata for a batch of tasks.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
req_indices: request slot indices (from ``ReqToTokenPool``).
|
||||||
|
seq_lens: current sequence length per task.
|
||||||
|
workspace: pre-allocated fixed-shape buffers (CUDA-graph safe).
|
||||||
|
start_pos: if set, produce **prefill** cache (full q_len range).
|
||||||
|
If ``None``, produce **decode** cache (last position).
|
||||||
|
incremental: if ``True``, reuse workspace state from previous step
|
||||||
|
by incrementing counters in-place (decode hot path).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
``KVCache`` dataclass with the correct output shapes for the
|
||||||
|
attention backend (prefill: ``[B, q_len]``, decode: ``[B, 1]``).
|
||||||
|
"""
|
||||||
|
if device is None:
|
||||||
|
device = workspace.device
|
||||||
|
b = len(req_indices)
|
||||||
|
|
||||||
|
rpi_buf = workspace.req_pool_indices
|
||||||
|
sl_buf = workspace.seq_lens
|
||||||
|
kvp_buf = workspace.kv_indptr
|
||||||
|
inc_buf = workspace.inc
|
||||||
|
ocl_buf = workspace.out_cache_loc
|
||||||
|
|
||||||
|
if incremental:
|
||||||
|
sl_buf[:b] += 1
|
||||||
|
kvp_buf[: b + 1] += inc_buf[: b + 1]
|
||||||
|
else:
|
||||||
|
rpi_buf[:b].copy_(
|
||||||
|
torch.tensor(req_indices, dtype=torch.int32, device=device)
|
||||||
|
)
|
||||||
|
sl_buf[:b].copy_(torch.tensor(seq_lens, dtype=torch.long, device=device))
|
||||||
|
kvp_buf[: b + 1].zero_()
|
||||||
|
kvp_buf[1 : b + 1] = sl_buf[:b].cumsum(0).to(torch.int32)
|
||||||
|
|
||||||
|
req_pool_indices = rpi_buf[:b]
|
||||||
|
seq_lens_t = sl_buf[:b]
|
||||||
|
kv_indptr = kvp_buf[: b + 1]
|
||||||
|
|
||||||
|
if start_pos is not None:
|
||||||
|
# Packed prefill concatenates each request's query tokens.
|
||||||
|
q_lens = [seq_len - start_pos for seq_len in seq_lens]
|
||||||
|
if any(q_len <= 0 for q_len in q_lens):
|
||||||
|
raise ValueError("prefill sequence lengths must exceed start_pos")
|
||||||
|
out_cache_loc = torch.cat(
|
||||||
|
[
|
||||||
|
self._req_pool.req_to_token[
|
||||||
|
req_pool_indices[i], start_pos : seq_lens[i]
|
||||||
|
]
|
||||||
|
for i in range(b)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
workspace.qo_indptr[: b + 1].zero_()
|
||||||
|
workspace.qo_indptr[1 : b + 1].copy_(
|
||||||
|
torch.tensor(q_lens, dtype=torch.int32, device=device).cumsum(0)
|
||||||
|
)
|
||||||
|
qo_indptr = workspace.qo_indptr[: b + 1]
|
||||||
|
tile_batches = []
|
||||||
|
tile_indices = []
|
||||||
|
for batch, q_len in enumerate(q_lens):
|
||||||
|
n_tiles = (q_len + Q_TILE_ROWS - 1) // Q_TILE_ROWS
|
||||||
|
tile_batches.extend([batch] * n_tiles)
|
||||||
|
tile_indices.extend(range(n_tiles))
|
||||||
|
n_tiles = len(tile_batches)
|
||||||
|
workspace.q_tile_to_batch[:n_tiles].copy_(
|
||||||
|
torch.tensor(tile_batches, dtype=torch.int32, device=device)
|
||||||
|
)
|
||||||
|
workspace.q_tile_to_index[:n_tiles].copy_(
|
||||||
|
torch.tensor(tile_indices, dtype=torch.int32, device=device)
|
||||||
|
)
|
||||||
|
q_tile_to_batch = workspace.q_tile_to_batch[:n_tiles]
|
||||||
|
q_tile_to_index = workspace.q_tile_to_index[:n_tiles]
|
||||||
|
decode_o_part = decode_ml_part = decode_out = None
|
||||||
|
else:
|
||||||
|
# ---- decode: out_cache_loc is a single column (last position) ----
|
||||||
|
write_pos = seq_lens_t - 1
|
||||||
|
loc = self._req_pool.req_to_token[req_pool_indices, write_pos].unsqueeze(-1)
|
||||||
|
ocl_buf[:b].copy_(loc)
|
||||||
|
out_cache_loc = ocl_buf[:b].reshape(-1)
|
||||||
|
workspace.qo_indptr[: b + 1].copy_(inc_buf[: b + 1])
|
||||||
|
qo_indptr = workspace.qo_indptr[: b + 1]
|
||||||
|
q_tile_to_batch = q_tile_to_index = None
|
||||||
|
decode_o_part = getattr(workspace, "decode_o_part", None)
|
||||||
|
decode_ml_part = getattr(workspace, "decode_ml_part", None)
|
||||||
|
decode_out = getattr(workspace, "decode_out", None)
|
||||||
|
|
||||||
|
return KVCache(
|
||||||
|
k_buffer=self._storage.k_buffer,
|
||||||
|
v_buffer=self._storage.v_buffer,
|
||||||
|
req_to_token=self._req_pool.req_to_token,
|
||||||
|
req_pool_indices=req_pool_indices,
|
||||||
|
seq_lens=seq_lens_t,
|
||||||
|
out_cache_loc=out_cache_loc,
|
||||||
|
max_len=max(seq_lens),
|
||||||
|
kv_indptr=kv_indptr,
|
||||||
|
qo_indptr=qo_indptr,
|
||||||
|
q_tile_to_batch=q_tile_to_batch,
|
||||||
|
q_tile_to_index=q_tile_to_index,
|
||||||
|
decode_o_part=decode_o_part,
|
||||||
|
decode_ml_part=decode_ml_part,
|
||||||
|
decode_out=decode_out,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TaskCacheManager:
|
||||||
|
"""Task ↔ KV slot lifecycle manager.
|
||||||
|
|
||||||
|
Sole owner of ``task_id → TaskCacheState``. Delegates physical slot
|
||||||
|
allocation to the strategy (via ``pool.strategy``) and KV bind to
|
||||||
|
``pool.bind_tasks()``.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
pool = PagePool(...)
|
||||||
|
mgr = TaskCacheManager(pool)
|
||||||
|
mgr.task_alloc("req_1", [101, 202, 303])
|
||||||
|
...
|
||||||
|
kv = mgr.bind(["req_1"], workspace)
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, pool: PagePool):
|
||||||
|
self._pool = pool
|
||||||
|
self._strategy = pool.strategy
|
||||||
|
self._req_pool = pool.req_pool
|
||||||
|
self._max_seq_len = pool.max_seq_len
|
||||||
|
self._states: Dict[str, TaskCacheState] = {}
|
||||||
|
self._bind_state: Optional[_BindState] = None
|
||||||
|
self._bind_was_steady = False
|
||||||
|
|
||||||
|
# -- public task lifecycle --
|
||||||
|
|
||||||
|
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
||||||
|
self._bind_state = None
|
||||||
|
req_slots = self._req_pool.alloc(1)
|
||||||
|
if req_slots is None:
|
||||||
|
return False
|
||||||
|
state = TaskCacheState(req_idx=req_slots[0])
|
||||||
|
self._states[task_id] = state
|
||||||
|
if not self._strategy.alloc(state, prompt_ids):
|
||||||
|
self._rollback(state, task_id)
|
||||||
|
return False
|
||||||
|
self._strategy.write_indices(state, prompt_ids)
|
||||||
|
state.length = len(prompt_ids)
|
||||||
|
return True
|
||||||
|
|
||||||
|
def task_free(self, task_id: str):
|
||||||
|
self._bind_state = None
|
||||||
|
state = self._states.pop(task_id, None)
|
||||||
|
if state is None:
|
||||||
|
return
|
||||||
|
self._strategy.free(state)
|
||||||
|
self._req_pool.free([state.req_idx])
|
||||||
|
|
||||||
|
def task_extend(self, task_id: str, pos: int) -> bool:
|
||||||
|
state = self._states.get(task_id)
|
||||||
|
if state is None or pos >= self._max_seq_len:
|
||||||
|
return False
|
||||||
|
if not self._strategy.extend(state, pos):
|
||||||
|
return False
|
||||||
|
state.length = pos + 1
|
||||||
|
return True
|
||||||
|
|
||||||
|
def task_cached(self, task_id: str) -> int:
|
||||||
|
state = self._states.get(task_id)
|
||||||
|
return state.cached if state is not None else 0
|
||||||
|
|
||||||
|
def task_record_hashes(
|
||||||
|
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
||||||
|
):
|
||||||
|
state = self._states.get(task_id)
|
||||||
|
if state is not None:
|
||||||
|
self._strategy.record_hashes(state, prompt_ids, start_logical_page)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def task_cacheable_ids(task_id: str, prompt_ids: List[int], output_ids: List[int]):
|
||||||
|
return list(prompt_ids) + list(output_ids[:-1])
|
||||||
|
|
||||||
|
# -- bind (assemble KVCache for the model forward) --
|
||||||
|
|
||||||
|
def bind(
|
||||||
|
self,
|
||||||
|
task_ids: List[str],
|
||||||
|
workspace: InferenceWorkspace,
|
||||||
|
device: Optional[torch.device] = None,
|
||||||
|
start_pos: Optional[int] = None,
|
||||||
|
) -> KVCache:
|
||||||
|
"""Build ``KVCache`` for an ordered list of task IDs."""
|
||||||
|
states = [self._states[tid] for tid in task_ids]
|
||||||
|
req_indices = [s.req_idx for s in states]
|
||||||
|
seq_lens = [s.length for s in states]
|
||||||
|
sig = tuple(req_indices)
|
||||||
|
|
||||||
|
prev = self._bind_state
|
||||||
|
incremental = (
|
||||||
|
start_pos is None
|
||||||
|
and prev is not None
|
||||||
|
and _is_steady_increment(prev.sig, prev.seq_lens, sig, seq_lens)
|
||||||
|
)
|
||||||
|
self._bind_state = _BindState(sig, list(seq_lens))
|
||||||
|
self._bind_was_steady = incremental
|
||||||
|
|
||||||
|
return self._pool.bind_tasks(
|
||||||
|
req_indices,
|
||||||
|
seq_lens,
|
||||||
|
workspace,
|
||||||
|
device=device,
|
||||||
|
start_pos=start_pos,
|
||||||
|
incremental=incremental,
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def bind_was_steady(self) -> bool:
|
||||||
|
return self._bind_was_steady
|
||||||
|
|
||||||
|
# -- internals --
|
||||||
|
|
||||||
|
def _rollback(self, state: TaskCacheState, task_id: str):
|
||||||
|
self._strategy.free(state)
|
||||||
|
self._req_pool.free([state.req_idx])
|
||||||
|
self._states.pop(task_id, None)
|
||||||
Vendored
+318
@@ -0,0 +1,318 @@
|
|||||||
|
"""KV cache allocation layer.
|
||||||
|
|
||||||
|
Encapsulates the physical slot allocation policy, isolated from GPU buffers
|
||||||
|
and task lifecycle management.
|
||||||
|
|
||||||
|
- ``TaskCacheState``: data contract between strategy and manager (per-task slot state)
|
||||||
|
- ``Allocator``: bitmask-based page allocator with LRU eviction
|
||||||
|
- ``RadixCache``: page-granular prefix index (exact token match)
|
||||||
|
- ``AllocationStrategy``: ABC for physical slot allocation
|
||||||
|
- ``ContiguousStrategy``: statically partitioned, no dynamic allocation
|
||||||
|
- ``PagedStrategy``: dynamic paged allocation from a shared pool
|
||||||
|
"""
|
||||||
|
|
||||||
|
import threading
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Callable, Dict, List, Optional, OrderedDict
|
||||||
|
|
||||||
|
from astrai.inference.cache.buffer import ReqToTokenPool
|
||||||
|
|
||||||
|
# ---- data contract: per-task slot state ----
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class TaskCacheState:
|
||||||
|
"""Per-task cache allocation state.
|
||||||
|
|
||||||
|
Co-locates all task-owned cache metadata so the alloc/free/extend
|
||||||
|
lifecycle is atomic. Owned by ``TaskCacheManager``, consumed by
|
||||||
|
every ``AllocationStrategy`` method.
|
||||||
|
"""
|
||||||
|
|
||||||
|
req_idx: int
|
||||||
|
length: int = 0
|
||||||
|
cached: int = 0
|
||||||
|
pages: List[int] = field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
|
# ---- allocation primitives ----
|
||||||
|
|
||||||
|
|
||||||
|
class Allocator:
|
||||||
|
"""Bitmask-based page allocator with ref-counting and LRU eviction."""
|
||||||
|
|
||||||
|
def __init__(self, n_pages: int):
|
||||||
|
self._free_mask = (1 << n_pages) - 1
|
||||||
|
self._refs: List[int] = [0] * n_pages
|
||||||
|
self._lru: OrderedDict[int, None] = OrderedDict()
|
||||||
|
self.on_evict: Optional[Callable[[int], None]] = None
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
def alloc(self) -> int:
|
||||||
|
with self._lock:
|
||||||
|
if self._free_mask:
|
||||||
|
lsb = self._free_mask & -self._free_mask
|
||||||
|
idx = lsb.bit_length() - 1
|
||||||
|
self._free_mask ^= lsb
|
||||||
|
self._refs[idx] = 1
|
||||||
|
return idx
|
||||||
|
if self._lru:
|
||||||
|
idx, _ = self._lru.popitem(last=False)
|
||||||
|
if self.on_evict:
|
||||||
|
self.on_evict(idx)
|
||||||
|
self._refs[idx] = 1
|
||||||
|
self._free_mask &= ~(1 << idx)
|
||||||
|
return idx
|
||||||
|
return -1
|
||||||
|
|
||||||
|
def free(self, idx: int, keep_cached: bool = False):
|
||||||
|
with self._lock:
|
||||||
|
self._refs[idx] -= 1
|
||||||
|
if self._refs[idx] == 0:
|
||||||
|
if keep_cached:
|
||||||
|
self._lru[idx] = None
|
||||||
|
else:
|
||||||
|
self._free_mask |= 1 << idx
|
||||||
|
|
||||||
|
def inc_ref(self, idx: int):
|
||||||
|
with self._lock:
|
||||||
|
self._refs[idx] += 1
|
||||||
|
self._lru.pop(idx, None)
|
||||||
|
|
||||||
|
def ref_count(self, idx: int) -> int:
|
||||||
|
with self._lock:
|
||||||
|
return self._refs[idx]
|
||||||
|
|
||||||
|
def touch(self, idx: int):
|
||||||
|
with self._lock:
|
||||||
|
if idx in self._lru:
|
||||||
|
self._lru.move_to_end(idx)
|
||||||
|
|
||||||
|
|
||||||
|
class RadixNode:
|
||||||
|
"""A page-aligned edge in the CPU-side prefix radix trie."""
|
||||||
|
|
||||||
|
__slots__ = ("parent", "children", "page_idx", "tokens", "lock_ref")
|
||||||
|
|
||||||
|
def __init__(self, parent=None, tokens=(), page_idx=None):
|
||||||
|
self.parent = parent
|
||||||
|
self.children: Dict[tuple, "RadixNode"] = {}
|
||||||
|
self.page_idx = page_idx
|
||||||
|
self.tokens = tuple(tokens)
|
||||||
|
self.lock_ref = 0
|
||||||
|
|
||||||
|
|
||||||
|
class RadixCache:
|
||||||
|
"""Page-granular radix prefix index with exact token matching."""
|
||||||
|
|
||||||
|
def __init__(self, page_size: int):
|
||||||
|
self._page_size = page_size
|
||||||
|
self._root = RadixNode()
|
||||||
|
self._page_to_node: Dict[int, RadixNode] = {}
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
def evict(self, idx: int):
|
||||||
|
with self._lock:
|
||||||
|
node = self._page_to_node.pop(idx, None)
|
||||||
|
if node is None:
|
||||||
|
return
|
||||||
|
node.page_idx = None
|
||||||
|
parent = node.parent
|
||||||
|
if parent is not None:
|
||||||
|
parent.children.pop(node.tokens, None)
|
||||||
|
|
||||||
|
def has_page(self, idx: int) -> bool:
|
||||||
|
with self._lock:
|
||||||
|
return idx in self._page_to_node
|
||||||
|
|
||||||
|
def lookup(self, token_ids: List[int]) -> List[int]:
|
||||||
|
with self._lock:
|
||||||
|
full_pages = len(token_ids) // self._page_size
|
||||||
|
hits: List[int] = []
|
||||||
|
node = self._root
|
||||||
|
for i in range(full_pages):
|
||||||
|
start = i * self._page_size
|
||||||
|
page_tokens = tuple(token_ids[start : start + self._page_size])
|
||||||
|
child = node.children.get(page_tokens)
|
||||||
|
if child is None or child.page_idx is None:
|
||||||
|
break
|
||||||
|
hits.append(child.page_idx)
|
||||||
|
node = child
|
||||||
|
return hits
|
||||||
|
|
||||||
|
def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int):
|
||||||
|
with self._lock:
|
||||||
|
full_pages = len(token_ids) // self._page_size
|
||||||
|
if logical_page_idx >= full_pages:
|
||||||
|
return
|
||||||
|
old = self._page_to_node.pop(page_idx, None)
|
||||||
|
if old is not None and old.parent is not None:
|
||||||
|
old.parent.children.pop(old.tokens, None)
|
||||||
|
|
||||||
|
node = self._root
|
||||||
|
for i in range(logical_page_idx + 1):
|
||||||
|
start = i * self._page_size
|
||||||
|
page_tokens = tuple(token_ids[start : start + self._page_size])
|
||||||
|
child = node.children.get(page_tokens)
|
||||||
|
if child is None:
|
||||||
|
child = RadixNode(node, page_tokens)
|
||||||
|
node.children[page_tokens] = child
|
||||||
|
node = child
|
||||||
|
if node.page_idx is not None and node.page_idx != page_idx:
|
||||||
|
replaced = node.page_idx
|
||||||
|
self._page_to_node.pop(replaced, None)
|
||||||
|
node.page_idx = page_idx
|
||||||
|
self._page_to_node[page_idx] = node
|
||||||
|
|
||||||
|
def release(self, pages: List[int]) -> None:
|
||||||
|
with self._lock:
|
||||||
|
for page_idx in pages:
|
||||||
|
node = self._page_to_node.get(page_idx)
|
||||||
|
if node is not None and node.lock_ref:
|
||||||
|
node.lock_ref -= 1
|
||||||
|
|
||||||
|
|
||||||
|
class AllocationStrategy(ABC):
|
||||||
|
"""Physical slot allocation policy.
|
||||||
|
|
||||||
|
Subclasses implement the actual allocation semantics. This ABC declares
|
||||||
|
the contract; there are no default implementations.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def alloc(self, state: TaskCacheState, prompt_ids: List[int]) -> bool: ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def free(self, state: TaskCacheState) -> None: ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def extend(self, state: TaskCacheState, pos: int) -> bool: ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None: ...
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def record_hashes(
|
||||||
|
self,
|
||||||
|
state: TaskCacheState,
|
||||||
|
prompt_ids: List[int],
|
||||||
|
start: int,
|
||||||
|
) -> None: ...
|
||||||
|
|
||||||
|
|
||||||
|
class ContiguousStrategy(AllocationStrategy):
|
||||||
|
"""Static contiguous allocation: slots are pre-assigned at pool init.
|
||||||
|
|
||||||
|
No dynamic allocation or prefix caching. All operations are no-ops
|
||||||
|
because ``ReqToTokenPool`` is pre-filled with contiguous ranges.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def alloc(self, state: TaskCacheState, prompt_ids: List[int]) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
def free(self, state: TaskCacheState) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def extend(self, state: TaskCacheState, pos: int) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def record_hashes(
|
||||||
|
self,
|
||||||
|
state: TaskCacheState,
|
||||||
|
prompt_ids: List[int],
|
||||||
|
start: int,
|
||||||
|
) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class PagedStrategy(AllocationStrategy):
|
||||||
|
"""Dynamic paged allocation from a shared bitmask pool.
|
||||||
|
|
||||||
|
``page_size`` is a parameter, not a separate strategy: at ``page_size=1``
|
||||||
|
each allocated page *is* one token slot (``page * 1 + 0``), and prefix
|
||||||
|
caching is simply disabled (``prefix=None``). The unified page formula
|
||||||
|
``pages[page_idx] * page_size + offset`` holds for both.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
alloc: Allocator,
|
||||||
|
prefix: Optional[RadixCache],
|
||||||
|
page_size: int,
|
||||||
|
req_pool: ReqToTokenPool,
|
||||||
|
device,
|
||||||
|
):
|
||||||
|
self._alloc = alloc
|
||||||
|
self._prefix = prefix
|
||||||
|
self._page_size = page_size
|
||||||
|
self._req_pool = req_pool
|
||||||
|
self._device = device
|
||||||
|
|
||||||
|
def alloc(self, state: TaskCacheState, prompt_ids: List[int]) -> bool:
|
||||||
|
if self._prefix is not None:
|
||||||
|
hits = self._prefix.lookup(prompt_ids)
|
||||||
|
state.cached = len(hits) * self._page_size
|
||||||
|
for p in hits:
|
||||||
|
self._alloc.inc_ref(p)
|
||||||
|
state.pages = list(hits)
|
||||||
|
|
||||||
|
remaining = len(prompt_ids) - state.cached
|
||||||
|
if remaining <= 0:
|
||||||
|
return True
|
||||||
|
n_new = (remaining + self._page_size - 1) // self._page_size
|
||||||
|
for _ in range(n_new):
|
||||||
|
p = self._alloc.alloc()
|
||||||
|
if p < 0:
|
||||||
|
return False
|
||||||
|
state.pages.append(p)
|
||||||
|
return True
|
||||||
|
|
||||||
|
def free(self, state: TaskCacheState) -> None:
|
||||||
|
if self._prefix is not None:
|
||||||
|
for p in state.pages:
|
||||||
|
keep = self._prefix.has_page(p)
|
||||||
|
self._alloc.free(p, keep_cached=keep)
|
||||||
|
if not keep:
|
||||||
|
self._prefix.evict(p)
|
||||||
|
else:
|
||||||
|
for p in state.pages:
|
||||||
|
self._alloc.free(p)
|
||||||
|
|
||||||
|
def extend(self, state: TaskCacheState, pos: int) -> bool:
|
||||||
|
page_idx = pos // self._page_size
|
||||||
|
if page_idx >= len(state.pages):
|
||||||
|
p = self._alloc.alloc()
|
||||||
|
if p < 0:
|
||||||
|
return False
|
||||||
|
state.pages.append(p)
|
||||||
|
offset = pos % self._page_size
|
||||||
|
self._req_pool.req_to_token[state.req_idx, pos] = (
|
||||||
|
state.pages[page_idx] * self._page_size + offset
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
def write_indices(self, state: TaskCacheState, prompt_ids: List[int]) -> None:
|
||||||
|
total = len(prompt_ids)
|
||||||
|
for pos in range(total):
|
||||||
|
page_idx = pos // self._page_size
|
||||||
|
offset = pos % self._page_size
|
||||||
|
if page_idx < len(state.pages):
|
||||||
|
self._req_pool.req_to_token[state.req_idx, pos] = (
|
||||||
|
state.pages[page_idx] * self._page_size + offset
|
||||||
|
)
|
||||||
|
|
||||||
|
def record_hashes(
|
||||||
|
self,
|
||||||
|
state: TaskCacheState,
|
||||||
|
prompt_ids: List[int],
|
||||||
|
start: int,
|
||||||
|
) -> None:
|
||||||
|
if self._prefix is None:
|
||||||
|
return
|
||||||
|
full = len(prompt_ids) // self._page_size
|
||||||
|
for i in range(start, min(full, len(state.pages))):
|
||||||
|
self._prefix.record(state.pages[i], prompt_ids, i)
|
||||||
@@ -0,0 +1,242 @@
|
|||||||
|
"""Unified inference engine for continuous batching."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import gc
|
||||||
|
import threading
|
||||||
|
from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
class GenerateResult:
|
||||||
|
"""Thread-safe token accumulator for streaming and non-streaming modes."""
|
||||||
|
|
||||||
|
def __init__(self, count: int = 1):
|
||||||
|
self._cond = threading.Condition()
|
||||||
|
self._event = threading.Event()
|
||||||
|
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):
|
||||||
|
with self._cond:
|
||||||
|
self.tokens.append((idx, token))
|
||||||
|
if token is not STOP:
|
||||||
|
self.results[idx] += token
|
||||||
|
else:
|
||||||
|
if not self._done[idx]:
|
||||||
|
self._done[idx] = True
|
||||||
|
self._completed += 1
|
||||||
|
self._cond.notify_all()
|
||||||
|
self._event.set()
|
||||||
|
|
||||||
|
def pop_all(self) -> List[Tuple[int, str]]:
|
||||||
|
with self._cond:
|
||||||
|
out = self.tokens.copy()
|
||||||
|
self.tokens.clear()
|
||||||
|
if not out:
|
||||||
|
self._event.clear()
|
||||||
|
return out
|
||||||
|
|
||||||
|
def wait(self, timeout: Optional[float] = None) -> bool:
|
||||||
|
return self._event.wait(timeout=timeout)
|
||||||
|
|
||||||
|
def wait_completion(self, timeout: float = 300.0):
|
||||||
|
with self._cond:
|
||||||
|
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]:
|
||||||
|
with self._cond:
|
||||||
|
return self.results.copy()
|
||||||
|
|
||||||
|
|
||||||
|
class InferenceEngine:
|
||||||
|
"""Unified inference engine backed by continuous-batching scheduler."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: nn.Module,
|
||||||
|
tokenizer: AutoTokenizer,
|
||||||
|
max_batch_size: int = 1,
|
||||||
|
max_seq_len: Optional[int] = None,
|
||||||
|
cache: Optional[PagePool] = None,
|
||||||
|
enable_cuda_graph: bool = True,
|
||||||
|
backend: Optional[Union[str, ATTN_BACKEND, AttentionBackend, type]] = None,
|
||||||
|
):
|
||||||
|
self.model = model
|
||||||
|
self.tokenizer = tokenizer
|
||||||
|
self.scheduler = InferenceScheduler(
|
||||||
|
model=self.model,
|
||||||
|
tokenizer=self.tokenizer,
|
||||||
|
max_batch_size=max_batch_size,
|
||||||
|
max_seq_len=max_seq_len,
|
||||||
|
cache=cache,
|
||||||
|
enable_cuda_graph=enable_cuda_graph,
|
||||||
|
backend=backend,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.scheduler.start()
|
||||||
|
|
||||||
|
def __enter__(self):
|
||||||
|
return self
|
||||||
|
|
||||||
|
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||||
|
self.shutdown()
|
||||||
|
return False
|
||||||
|
|
||||||
|
def generate(
|
||||||
|
self,
|
||||||
|
prompt: Union[str, List[str]],
|
||||||
|
stream: bool = False,
|
||||||
|
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,
|
||||||
|
) -> Union[Generator, str, List[str]]:
|
||||||
|
is_batch = isinstance(prompt, list)
|
||||||
|
prompts = prompt if is_batch else [prompt]
|
||||||
|
|
||||||
|
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: 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]:
|
||||||
|
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, next, sync_gen, None)
|
||||||
|
if token is None:
|
||||||
|
break
|
||||||
|
yield token
|
||||||
|
|
||||||
|
return _agen()
|
||||||
|
|
||||||
|
def _generate(
|
||||||
|
self,
|
||||||
|
prompts: List[str],
|
||||||
|
is_batch: bool,
|
||||||
|
stream: bool,
|
||||||
|
max_tokens: Optional[int],
|
||||||
|
temperature: float,
|
||||||
|
top_p: float,
|
||||||
|
top_k: int,
|
||||||
|
frequency_penalty: float,
|
||||||
|
rep_window: int,
|
||||||
|
) -> 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,
|
||||||
|
frequency_penalty=frequency_penalty,
|
||||||
|
rep_window=rep_window,
|
||||||
|
backend=request_backend,
|
||||||
|
stream_callback=lambda token, idx=i: result.append(token, idx),
|
||||||
|
)
|
||||||
|
for i, p in enumerate(prompts)
|
||||||
|
]
|
||||||
|
|
||||||
|
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]
|
||||||
|
|
||||||
|
remaining = n
|
||||||
|
finished = [False] * n
|
||||||
|
|
||||||
|
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]:
|
||||||
|
return self.scheduler.get_stats()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def backend_name(self) -> str:
|
||||||
|
return self.scheduler.backend_name
|
||||||
|
|
||||||
|
@property
|
||||||
|
def cuda_graph_enabled(self) -> bool:
|
||||||
|
return self.scheduler.cuda_graph_enabled
|
||||||
|
|
||||||
|
def shutdown(self):
|
||||||
|
self.scheduler.stop()
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
gc.collect()
|
||||||
@@ -0,0 +1,201 @@
|
|||||||
|
"""Unified per-task perf/stats: timing records, context-manager scopes, aggregate reporting."""
|
||||||
|
|
||||||
|
import time
|
||||||
|
from collections import deque
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any, Deque, Dict, Generator, List, Literal, Optional
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class TaskTiming:
|
||||||
|
"""Timestamp snapshots and computed metrics for one generation task.
|
||||||
|
|
||||||
|
Created by :class:`MetricsCollector` at task-registration time;
|
||||||
|
updated via ``record`` / ``mark_finished``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
task_id: str
|
||||||
|
arrival_time: float
|
||||||
|
prefill_start_time: Optional[float] = None
|
||||||
|
first_token_time: Optional[float] = None
|
||||||
|
finish_time: Optional[float] = None
|
||||||
|
input_tokens: int = 0
|
||||||
|
output_tokens: int = 0
|
||||||
|
_decode_steps: int = 0
|
||||||
|
_decode_total_s: float = 0.0
|
||||||
|
|
||||||
|
# derived metrics
|
||||||
|
|
||||||
|
@property
|
||||||
|
def queue_wait_ms(self) -> Optional[float]:
|
||||||
|
if self.prefill_start_time is not None:
|
||||||
|
return (self.prefill_start_time - self.arrival_time) * 1000
|
||||||
|
return None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def ttft_ms(self) -> Optional[float]:
|
||||||
|
if self.first_token_time is not None:
|
||||||
|
return (self.first_token_time - self.arrival_time) * 1000
|
||||||
|
return None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def prefill_tps(self) -> Optional[float]:
|
||||||
|
if self.prefill_start_time is not None and self.first_token_time is not None:
|
||||||
|
d = self.first_token_time - self.prefill_start_time
|
||||||
|
if d > 0 and self.input_tokens > 0:
|
||||||
|
return self.input_tokens / d
|
||||||
|
return None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def decode_tps(self) -> Optional[float]:
|
||||||
|
if self.first_token_time is not None and self.finish_time is not None:
|
||||||
|
d = self.finish_time - self.first_token_time
|
||||||
|
dt = self.output_tokens - 1
|
||||||
|
if dt > 0 and d > 0:
|
||||||
|
return dt / d
|
||||||
|
return None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def decode_avg_ms(self) -> Optional[float]:
|
||||||
|
if self._decode_steps > 0 and self._decode_total_s > 0:
|
||||||
|
return (self._decode_total_s / self._decode_steps) * 1000
|
||||||
|
return None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def e2e_latency_ms(self) -> Optional[float]:
|
||||||
|
if self.finish_time is not None:
|
||||||
|
return (self.finish_time - self.arrival_time) * 1000
|
||||||
|
return None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def total_tps(self) -> Optional[float]:
|
||||||
|
if self.finish_time is not None:
|
||||||
|
total = self.input_tokens + self.output_tokens
|
||||||
|
d = self.finish_time - self.arrival_time
|
||||||
|
if total > 0 and d > 0:
|
||||||
|
return total / d
|
||||||
|
return None
|
||||||
|
|
||||||
|
def to_dict(self) -> Dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"task_id": self.task_id,
|
||||||
|
"input_tokens": self.input_tokens,
|
||||||
|
"output_tokens": self.output_tokens,
|
||||||
|
"queue_wait_ms": (
|
||||||
|
round(self.queue_wait_ms, 2) if self.queue_wait_ms is not None else None
|
||||||
|
),
|
||||||
|
"ttft_ms": (round(self.ttft_ms, 2) if self.ttft_ms is not None else None),
|
||||||
|
"prefill_tps": (
|
||||||
|
round(self.prefill_tps, 2) if self.prefill_tps is not None else None
|
||||||
|
),
|
||||||
|
"decode_tps": (
|
||||||
|
round(self.decode_tps, 2) if self.decode_tps is not None else None
|
||||||
|
),
|
||||||
|
"decode_avg_ms": (
|
||||||
|
round(self.decode_avg_ms, 2) if self.decode_avg_ms is not None else None
|
||||||
|
),
|
||||||
|
"total_tps": (
|
||||||
|
round(self.total_tps, 2) if self.total_tps is not None else None
|
||||||
|
),
|
||||||
|
"e2e_latency_ms": (
|
||||||
|
round(self.e2e_latency_ms, 2)
|
||||||
|
if self.e2e_latency_ms is not None
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class MetricsCollector:
|
||||||
|
"""Single-owner perf/stats hub for all generation tasks.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
metrics = MetricsCollector()
|
||||||
|
metrics.register(task_id, arrival_time)
|
||||||
|
|
||||||
|
with metrics.record(task_ids, "prefill"):
|
||||||
|
run_prefill(...)
|
||||||
|
|
||||||
|
metrics.mark_finished(task_id, input_tokens, output_tokens)
|
||||||
|
|
||||||
|
stats = metrics.get_stats()
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, max_recent: int = 128):
|
||||||
|
self._timings: Dict[str, TaskTiming] = {}
|
||||||
|
self._completed: Deque[TaskTiming] = deque(maxlen=max_recent)
|
||||||
|
|
||||||
|
self._ttft_ms_sum = 0.0
|
||||||
|
self._ttft_ms_count = 0
|
||||||
|
self._decode_tps_sum = 0.0
|
||||||
|
self._decode_tps_count = 0
|
||||||
|
self._e2e_ms_sum = 0.0
|
||||||
|
self._e2e_ms_count = 0
|
||||||
|
|
||||||
|
def register(self, task_id: str):
|
||||||
|
"""Create a timing record for a newly-created task."""
|
||||||
|
self._timings[task_id] = TaskTiming(task_id=task_id, arrival_time=time.time())
|
||||||
|
|
||||||
|
def mark_finished(self, task_id: str, input_tokens: int, output_tokens: int):
|
||||||
|
"""Close timing for a finished/aborted task and move it to completed."""
|
||||||
|
timing = self._timings.pop(task_id, None)
|
||||||
|
if timing is None:
|
||||||
|
return
|
||||||
|
timing.finish_time = time.time()
|
||||||
|
timing.input_tokens = input_tokens
|
||||||
|
timing.output_tokens = output_tokens
|
||||||
|
self._completed.append(timing)
|
||||||
|
self._accumulate(timing)
|
||||||
|
|
||||||
|
# timing scopes
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def record(
|
||||||
|
self, task_ids: List[str], phase: Literal["prefill", "decode"]
|
||||||
|
) -> Generator[None, None, None]:
|
||||||
|
tic = time.time()
|
||||||
|
yield
|
||||||
|
toc = time.time()
|
||||||
|
dt = toc - tic
|
||||||
|
for tid in task_ids:
|
||||||
|
t = self._timings.get(tid)
|
||||||
|
if t is None:
|
||||||
|
continue
|
||||||
|
if phase == "prefill":
|
||||||
|
t.prefill_start_time = tic
|
||||||
|
t.first_token_time = toc
|
||||||
|
elif phase == "decode":
|
||||||
|
t._decode_steps += 1
|
||||||
|
t._decode_total_s += dt
|
||||||
|
|
||||||
|
# aggregate stats
|
||||||
|
|
||||||
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
|
stats: Dict[str, Any] = {}
|
||||||
|
if self._ttft_ms_count > 0:
|
||||||
|
stats["avg_ttft_ms"] = round(self._ttft_ms_sum / self._ttft_ms_count, 2)
|
||||||
|
if self._decode_tps_count > 0:
|
||||||
|
stats["avg_decode_tps"] = round(
|
||||||
|
self._decode_tps_sum / self._decode_tps_count, 2
|
||||||
|
)
|
||||||
|
if self._e2e_ms_count > 0:
|
||||||
|
stats["avg_e2e_latency_ms"] = round(
|
||||||
|
self._e2e_ms_sum / self._e2e_ms_count, 2
|
||||||
|
)
|
||||||
|
if self._completed:
|
||||||
|
stats["recent_tasks"] = [t.to_dict() for t in self._completed]
|
||||||
|
return stats
|
||||||
|
|
||||||
|
# internal
|
||||||
|
|
||||||
|
def _accumulate(self, t: TaskTiming):
|
||||||
|
if t.ttft_ms is not None:
|
||||||
|
self._ttft_ms_sum += t.ttft_ms
|
||||||
|
self._ttft_ms_count += 1
|
||||||
|
if t.decode_tps is not None:
|
||||||
|
self._decode_tps_sum += t.decode_tps
|
||||||
|
self._decode_tps_count += 1
|
||||||
|
if t.e2e_latency_ms is not None:
|
||||||
|
self._e2e_ms_sum += t.e2e_latency_ms
|
||||||
|
self._e2e_ms_count += 1
|
||||||
@@ -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.available() and head_dim in CudaBackend.HEAD_DIMS
|
||||||
|
)
|
||||||
|
self._workspace = InferenceWorkspace(
|
||||||
|
max_batch_size=kv_cache.max_batch_size,
|
||||||
|
max_seq_len=kv_cache.max_seq_len,
|
||||||
|
max_q_heads=max_q_heads,
|
||||||
|
head_dim=head_dim,
|
||||||
|
device=self.device,
|
||||||
|
dtype=self.dtype,
|
||||||
|
)
|
||||||
|
|
||||||
|
# CUDA-graph capture: one graph per (batch_size,) key.
|
||||||
|
# Enabled at init-time via _warmup_cuda_graphs for CudaBackend
|
||||||
|
# on supported head_dims; left disabled otherwise.
|
||||||
|
self._graph_ctx = CudaGraphContext()
|
||||||
|
if enable_cuda_graph:
|
||||||
|
self._try_enable_cuda_graph()
|
||||||
|
|
||||||
|
def _try_enable_cuda_graph(self):
|
||||||
|
if not self._graph_supported:
|
||||||
|
return
|
||||||
|
|
||||||
|
self._graph_ctx.set_enabled(True)
|
||||||
|
_warmup_cuda_graphs(
|
||||||
|
self.model,
|
||||||
|
self.kv_cache,
|
||||||
|
self.task_cache,
|
||||||
|
self._workspace,
|
||||||
|
self._graph_ctx,
|
||||||
|
max_batch_size=self.kv_cache.max_batch_size,
|
||||||
|
device=self.device,
|
||||||
|
)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def cuda_graph_enabled(self) -> bool:
|
||||||
|
return self._graph_ctx.enabled and self._graph_supported
|
||||||
|
|
||||||
|
def _sample_logits(
|
||||||
|
self,
|
||||||
|
logits: Tensor,
|
||||||
|
tasks: List[Task],
|
||||||
|
return_logprobs: bool = False,
|
||||||
|
info: Optional[SamplingBatchInfo] = None,
|
||||||
|
):
|
||||||
|
info = info or _build_sampling_batch_info(tasks, self.device)
|
||||||
|
if info.has_freq:
|
||||||
|
history_lists = [
|
||||||
|
t.prompt_ids[-t.rep_window :] + t.output_ids for t in tasks
|
||||||
|
]
|
||||||
|
history_lens = [len(ids) for ids in history_lists]
|
||||||
|
max_len = max(history_lens, default=0)
|
||||||
|
padded_ids = torch.zeros(
|
||||||
|
len(tasks), max_len, dtype=torch.long, device=self.device
|
||||||
|
)
|
||||||
|
padded_mask = torch.zeros(
|
||||||
|
len(tasks), max_len, dtype=torch.bool, device=self.device
|
||||||
|
)
|
||||||
|
for i, ids in enumerate(history_lists):
|
||||||
|
length = len(ids)
|
||||||
|
padded_ids[i, :length] = torch.as_tensor(
|
||||||
|
ids, dtype=torch.long, device=self.device
|
||||||
|
)
|
||||||
|
padded_mask[i, :length] = True
|
||||||
|
else:
|
||||||
|
padded_ids = None
|
||||||
|
padded_mask = None
|
||||||
|
|
||||||
|
result = sample(
|
||||||
|
logits,
|
||||||
|
temperature=info.temperatures,
|
||||||
|
top_k=info.top_ks,
|
||||||
|
top_p=info.top_ps,
|
||||||
|
frequency_penalty=info.freq_penalties,
|
||||||
|
input_ids=padded_ids,
|
||||||
|
input_mask=padded_mask,
|
||||||
|
return_logprobs=return_logprobs,
|
||||||
|
)
|
||||||
|
if not return_logprobs:
|
||||||
|
return result.tolist()
|
||||||
|
|
||||||
|
tokens, logprobs = result
|
||||||
|
tokens_list = tokens.tolist()
|
||||||
|
logprobs_list = logprobs.tolist()
|
||||||
|
for task, logprob in zip(tasks, logprobs_list):
|
||||||
|
task.output_logprobs.append(float(logprob))
|
||||||
|
return list(zip(tokens_list, logprobs_list))
|
||||||
|
|
||||||
|
def execute_prefill(
|
||||||
|
self,
|
||||||
|
tasks: List[Task],
|
||||||
|
prompt_len: int,
|
||||||
|
start_pos: int = 0,
|
||||||
|
return_logprobs: bool = False,
|
||||||
|
):
|
||||||
|
if start_pos >= prompt_len:
|
||||||
|
return []
|
||||||
|
|
||||||
|
tasks = sorted(tasks, key=lambda t: t.task_id)
|
||||||
|
batch_sz = len(tasks)
|
||||||
|
|
||||||
|
input_ids = torch.tensor(
|
||||||
|
[token for t in tasks for token in t.prompt_ids[start_pos:prompt_len]],
|
||||||
|
dtype=torch.long,
|
||||||
|
device=self.device,
|
||||||
|
)
|
||||||
|
|
||||||
|
task_ids = [t.task_id for t in tasks]
|
||||||
|
position_ids = torch.arange(
|
||||||
|
start_pos, prompt_len, dtype=torch.long, device=self.device
|
||||||
|
).repeat(batch_sz)
|
||||||
|
|
||||||
|
with (
|
||||||
|
torch.inference_mode(),
|
||||||
|
timed(f"execute_prefill b={batch_sz} prompt_len={prompt_len}", logger),
|
||||||
|
):
|
||||||
|
outputs = self.model(
|
||||||
|
input_ids,
|
||||||
|
position_ids=position_ids,
|
||||||
|
kv_cache=self.task_cache.bind(
|
||||||
|
task_ids,
|
||||||
|
self._workspace,
|
||||||
|
start_pos=start_pos,
|
||||||
|
),
|
||||||
|
fwd="prefill",
|
||||||
|
)
|
||||||
|
q_len = prompt_len - start_pos
|
||||||
|
logits = outputs["logits"][
|
||||||
|
torch.arange(1, batch_sz + 1, device=self.device) * q_len - 1
|
||||||
|
]
|
||||||
|
|
||||||
|
return tasks, self._sample_logits(logits, tasks, return_logprobs)
|
||||||
|
|
||||||
|
def execute_decode(
|
||||||
|
self, tasks: List[Task], return_logprobs: bool = False
|
||||||
|
) -> List[int]:
|
||||||
|
"""Decode next token for each task.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
return_logprobs: When ``True``, also record (and return)
|
||||||
|
the log-probability of each sampled token under the
|
||||||
|
post-strategy sampling distribution. The logprob is
|
||||||
|
appended to ``task.output_logprobs`` and the return
|
||||||
|
list becomes ``List[Tuple[int, float]]``.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
``List[int]`` of sampled token IDs, or
|
||||||
|
``List[Tuple[int, float]]`` of ``(token_id, logprob)`` when
|
||||||
|
``return_logprobs`` is ``True``.
|
||||||
|
"""
|
||||||
|
if not tasks:
|
||||||
|
return []
|
||||||
|
|
||||||
|
b = len(tasks)
|
||||||
|
ws = self._workspace
|
||||||
|
|
||||||
|
# ---- pre-replay: update input buffers in-place ----
|
||||||
|
|
||||||
|
input_ids = ws.fill_input_ids(
|
||||||
|
[t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks]
|
||||||
|
)
|
||||||
|
|
||||||
|
task_ids = [t.task_id for t in tasks]
|
||||||
|
cur_positions = [t.next_pos for t in tasks]
|
||||||
|
|
||||||
|
kv_cache = self.task_cache.bind(task_ids, ws)
|
||||||
|
|
||||||
|
task_sig = tuple(task_ids)
|
||||||
|
reuse_decode_state = (
|
||||||
|
self.task_cache.bind_was_steady
|
||||||
|
and self._decode_cache is not None
|
||||||
|
and self._decode_cache.task_sig == task_sig
|
||||||
|
)
|
||||||
|
if reuse_decode_state:
|
||||||
|
info = self._decode_cache.sampling_info
|
||||||
|
ws.position_ids[:b] += 1
|
||||||
|
else:
|
||||||
|
info = _build_sampling_batch_info(tasks, self.device)
|
||||||
|
ws.position_ids[:b].copy_(
|
||||||
|
torch.tensor(cur_positions, dtype=torch.long, device=self.device)
|
||||||
|
)
|
||||||
|
self._decode_cache = DecodeSteadyState(task_sig, cur_positions, info)
|
||||||
|
|
||||||
|
# ---- forward (graph replay or live run + capture) ----
|
||||||
|
|
||||||
|
use_graph = (
|
||||||
|
self._graph_ctx.enabled
|
||||||
|
and self._graph_supported
|
||||||
|
and get_backend().supports_graph()
|
||||||
|
)
|
||||||
|
key = (b,)
|
||||||
|
with (
|
||||||
|
torch.inference_mode(),
|
||||||
|
timed(f"execute_decode forward b={b}", logger),
|
||||||
|
):
|
||||||
|
if use_graph:
|
||||||
|
outputs = self._graph_ctx.forward(
|
||||||
|
self.model,
|
||||||
|
key=key,
|
||||||
|
input_ids=input_ids,
|
||||||
|
kv_cache=kv_cache,
|
||||||
|
position_ids=ws.position_ids[:b],
|
||||||
|
fwd="decode",
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
outputs = self.model(
|
||||||
|
input_ids,
|
||||||
|
kv_cache=kv_cache,
|
||||||
|
position_ids=ws.position_ids[:b],
|
||||||
|
fwd="decode",
|
||||||
|
)
|
||||||
|
logits = outputs["logits"]
|
||||||
|
|
||||||
|
return self._sample_logits(logits, tasks, return_logprobs, info=info)
|
||||||
@@ -0,0 +1,103 @@
|
|||||||
|
"""CUDA-graph capture for the decode model-forward step.
|
||||||
|
|
||||||
|
Mirrors SGLang's cuda-graph manager: one graph per batch size. The graph
|
||||||
|
pair. The graph captures ``model.forward()`` with workspace-backed inputs
|
||||||
|
(all at fixed addresses). Before each replay the caller updates the input
|
||||||
|
buffer content in-place so the graph sees fresh data at the same tensor
|
||||||
|
addresses.
|
||||||
|
|
||||||
|
Only the model forward is captured — sampling runs outside the graph
|
||||||
|
(via ``torch.multinomial`` which consumes a mutable RNG state).
|
||||||
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
class CudaGraphContext:
|
||||||
|
"""CUDA-graph capture/replay for decode steps.
|
||||||
|
|
||||||
|
Parameters:
|
||||||
|
enabled: When ``False``, ``forward()`` always runs the live model
|
||||||
|
forward without capture/replay (graphs are cleared). Toggle at
|
||||||
|
runtime via the ``set_enabled()`` method.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
gctx = CudaGraphContext()
|
||||||
|
with torch.inference_mode():
|
||||||
|
outputs = gctx.forward(
|
||||||
|
model,
|
||||||
|
key=(batch_size,),
|
||||||
|
input_ids=workspace.input_ids[:b].unsqueeze(1),
|
||||||
|
input_mask=input_mask,
|
||||||
|
kv_cache=kv_cache,
|
||||||
|
position_ids=workspace.position_ids[:b].unsqueeze(1),
|
||||||
|
)
|
||||||
|
|
||||||
|
The first call at a given key runs *without* capture (warmup). The
|
||||||
|
second call captures the graph. Subsequent calls replay the captured
|
||||||
|
graph. A ``torch.cuda.synchronize()`` before capture drains in-flight
|
||||||
|
work so the graph trace is clean.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, enabled: bool = False):
|
||||||
|
self._enabled = enabled
|
||||||
|
self._graphs: dict[tuple, torch.cuda.CUDAGraph] = {}
|
||||||
|
self._outputs: dict[tuple, dict[str, Tensor]] = {}
|
||||||
|
self._warmed: set[tuple] = set()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def enabled(self) -> bool:
|
||||||
|
return self._enabled
|
||||||
|
|
||||||
|
def set_enabled(self, flag: bool):
|
||||||
|
"""Enable or disable CUDA-graph capture at runtime.
|
||||||
|
|
||||||
|
Disabling clears all captured graphs (frees GPU memory) and warmup
|
||||||
|
state. Re-enabling after disable starts fresh — graphs are
|
||||||
|
re-captured on the next warmup cycle.
|
||||||
|
"""
|
||||||
|
if flag == self._enabled:
|
||||||
|
return
|
||||||
|
self._enabled = flag
|
||||||
|
if not flag:
|
||||||
|
self._graphs.clear()
|
||||||
|
self._outputs.clear()
|
||||||
|
self._warmed.clear()
|
||||||
|
|
||||||
|
def forward(self, model, *, key, **kwargs) -> dict[str, Tensor]:
|
||||||
|
"""Run ``model(**kwargs)`` via graph replay or live forward.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model: callable, e.g. ``self.model.forward``.
|
||||||
|
key: ``(batch_size,)`` — the dispatch key (one graph per batch size).
|
||||||
|
**kwargs: arguments forwarded to ``model``. All tensor arguments
|
||||||
|
must reside at stable addresses (workspace buffers).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The dict produced by ``model(**kwargs)``, e.g.
|
||||||
|
``{"logits": ..., "h0": ...}``.
|
||||||
|
"""
|
||||||
|
if not self._enabled:
|
||||||
|
self._outputs[key] = model(**kwargs)
|
||||||
|
return self._outputs[key]
|
||||||
|
|
||||||
|
if key in self._graphs:
|
||||||
|
self._graphs[key].replay()
|
||||||
|
elif key in self._warmed:
|
||||||
|
cap_output = model(**kwargs)
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
graph = torch.cuda.CUDAGraph()
|
||||||
|
with torch.cuda.graph(graph):
|
||||||
|
self._outputs[key] = model(**kwargs)
|
||||||
|
self._graphs[key] = graph
|
||||||
|
self._warmed.discard(key)
|
||||||
|
return cap_output
|
||||||
|
else:
|
||||||
|
self._warmed.add(key)
|
||||||
|
self._outputs[key] = model(**kwargs)
|
||||||
|
return self._outputs[key]
|
||||||
|
|
||||||
|
def has_graph(self, key: tuple) -> bool:
|
||||||
|
return key in self._graphs
|
||||||
@@ -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,
|
||||||
|
)
|
||||||
@@ -0,0 +1,400 @@
|
|||||||
|
import logging
|
||||||
|
import threading
|
||||||
|
import uuid
|
||||||
|
from contextlib import nullcontext
|
||||||
|
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.extension import (
|
||||||
|
ATTN_BACKEND,
|
||||||
|
AttentionBackend,
|
||||||
|
attn_backend,
|
||||||
|
get_backend,
|
||||||
|
)
|
||||||
|
from astrai.inference.cache import PagePool, TaskCacheManager
|
||||||
|
from astrai.inference.metrics import MetricsCollector
|
||||||
|
from astrai.inference.runtime.executor import Executor
|
||||||
|
from astrai.inference.task import STOP, Task, TaskManager, TaskStatus
|
||||||
|
from astrai.model.automodel import AutoModel
|
||||||
|
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class InferenceScheduler:
|
||||||
|
"""Continuous batching loop: cleanup -> refill -> prefill -> decode (all groups)."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: AutoModel,
|
||||||
|
tokenizer: AutoTokenizer,
|
||||||
|
max_batch_size: int = 16,
|
||||||
|
max_seq_len: Optional[int] = None,
|
||||||
|
device: Optional[str] = None,
|
||||||
|
dtype: Optional[torch.dtype] = None,
|
||||||
|
cache: Optional[PagePool] = None,
|
||||||
|
enable_cuda_graph: bool = True,
|
||||||
|
backend: Optional[Union[str, ATTN_BACKEND, AttentionBackend, type]] = None,
|
||||||
|
):
|
||||||
|
config = model.config
|
||||||
|
|
||||||
|
if max_seq_len is not None:
|
||||||
|
self.max_seq_len = max_seq_len
|
||||||
|
elif config.max_position_embeddings is not None:
|
||||||
|
self.max_seq_len = config.max_position_embeddings
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
"max_seq_len must be provided either as argument "
|
||||||
|
"or in model config (config.max_position_embeddings)"
|
||||||
|
)
|
||||||
|
self.device = device or next(model.parameters()).device
|
||||||
|
self.dtype = dtype or next(model.parameters()).dtype
|
||||||
|
|
||||||
|
head_dim = config.hidden_size // config.num_attention_heads
|
||||||
|
|
||||||
|
if cache is not None:
|
||||||
|
self._cache = cache
|
||||||
|
else:
|
||||||
|
self._cache = PagePool(
|
||||||
|
n_layers=config.num_hidden_layers,
|
||||||
|
n_kv_heads=config.num_key_value_heads,
|
||||||
|
head_dim=head_dim,
|
||||||
|
max_batch_size=max_batch_size,
|
||||||
|
max_seq_len=self.max_seq_len,
|
||||||
|
device=self.device,
|
||||||
|
dtype=self.dtype,
|
||||||
|
)
|
||||||
|
|
||||||
|
self._metrics = MetricsCollector()
|
||||||
|
|
||||||
|
self._task_cache = TaskCacheManager(self._cache)
|
||||||
|
|
||||||
|
self._task_mgr = TaskManager(
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
max_batch_size=max_batch_size,
|
||||||
|
max_seq_len=self.max_seq_len,
|
||||||
|
metrics=self._metrics,
|
||||||
|
)
|
||||||
|
|
||||||
|
if backend is None:
|
||||||
|
self._backend = None
|
||||||
|
active_backend = get_backend()
|
||||||
|
else:
|
||||||
|
active_backend = backend
|
||||||
|
with attn_backend(active_backend):
|
||||||
|
if backend is not None:
|
||||||
|
self._backend = get_backend()
|
||||||
|
self._backend_name = type(get_backend()).__name__
|
||||||
|
self._executor = Executor(
|
||||||
|
model=model,
|
||||||
|
kv_cache=self._cache,
|
||||||
|
task_cache=self._task_cache,
|
||||||
|
device=self.device,
|
||||||
|
dtype=self.dtype,
|
||||||
|
enable_cuda_graph=enable_cuda_graph,
|
||||||
|
)
|
||||||
|
|
||||||
|
self._stop_event = threading.Event()
|
||||||
|
self._loop_thread: Optional[threading.Thread] = None
|
||||||
|
|
||||||
|
def add_task(self, prompt: str, **kwargs) -> str:
|
||||||
|
return self._task_mgr.add_task(prompt, **kwargs)
|
||||||
|
|
||||||
|
def remove_task(self, task_id: str):
|
||||||
|
for task in self._task_mgr.remove_task(task_id):
|
||||||
|
self._task_cache.task_free(task.task_id)
|
||||||
|
|
||||||
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
|
return self._task_mgr.get_stats()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def backend_name(self) -> str:
|
||||||
|
return self._backend_name
|
||||||
|
|
||||||
|
@property
|
||||||
|
def cuda_graph_enabled(self) -> bool:
|
||||||
|
return self._executor.cuda_graph_enabled
|
||||||
|
|
||||||
|
def _backend_context(self):
|
||||||
|
if self._backend is None:
|
||||||
|
return nullcontext()
|
||||||
|
return attn_backend(self._backend)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _task_backend_groups(tasks: List[Task]):
|
||||||
|
groups = {}
|
||||||
|
for task in tasks:
|
||||||
|
groups.setdefault(task.backend, (task.backend, []))[1].append(task)
|
||||||
|
return groups.values()
|
||||||
|
|
||||||
|
def _step(
|
||||||
|
self, tasks: List[Task], return_logprobs: bool = False
|
||||||
|
) -> Tuple[List[Task], List[Task]]:
|
||||||
|
"""Advance every active task by one token (prefill + decode).
|
||||||
|
|
||||||
|
Single shared primitive for both the continuous-batching loop and
|
||||||
|
the synchronous ``run_batch`` path, so the two cannot drift.
|
||||||
|
|
||||||
|
Tasks must already be allocated in the KV cache. Tasks without output
|
||||||
|
are prefilled first and sample their first token from the final prompt
|
||||||
|
position. Tasks with output extend the cache by one position and decode
|
||||||
|
from their latest generated token.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
tasks: Active tasks to advance by one token.
|
||||||
|
return_logprobs: Forwarded to ``execute_decode``; per-token
|
||||||
|
logprobs are recorded on each task's ``output_logprobs``.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
``(decoded, aborted)``: tasks that produced a new token (its ID
|
||||||
|
already appended to ``output_ids``) and tasks that hit the
|
||||||
|
sequence cap and were marked ``ABORTED``.
|
||||||
|
"""
|
||||||
|
to_prefill = [t for t in tasks if not t.prefill_done and t.prompt_ids]
|
||||||
|
prefilled_ids = set()
|
||||||
|
produced: List[Task] = []
|
||||||
|
if to_prefill:
|
||||||
|
for t in to_prefill:
|
||||||
|
t.input_tokens = len(t.prompt_ids)
|
||||||
|
|
||||||
|
groups: Dict[Tuple[int, int, Optional[AttentionBackend]], List[Task]] = {}
|
||||||
|
for t in to_prefill:
|
||||||
|
start_pos = min(
|
||||||
|
self._task_cache.task_cached(t.task_id), len(t.prompt_ids) - 1
|
||||||
|
)
|
||||||
|
groups.setdefault((len(t.prompt_ids), start_pos, t.backend), []).append(
|
||||||
|
t
|
||||||
|
)
|
||||||
|
|
||||||
|
for (prompt_len, start_pos, _), group in groups.items():
|
||||||
|
backend = group[0].backend
|
||||||
|
backend_context = (
|
||||||
|
attn_backend(backend) if backend is not None else nullcontext()
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
backend_context,
|
||||||
|
self._metrics.record([t.task_id for t in group], "prefill"),
|
||||||
|
):
|
||||||
|
prefilled, step_out = self._executor.execute_prefill(
|
||||||
|
group, prompt_len, start_pos, return_logprobs=return_logprobs
|
||||||
|
)
|
||||||
|
|
||||||
|
for t, out in zip(prefilled, step_out):
|
||||||
|
t.output_ids.append(out[0] if return_logprobs else out)
|
||||||
|
t.output_tokens += 1
|
||||||
|
t.mark_prefill_done()
|
||||||
|
prefilled_ids.add(t.task_id)
|
||||||
|
produced.append(t)
|
||||||
|
|
||||||
|
start_logical_page = start_pos // self._cache.page_size
|
||||||
|
for t in group:
|
||||||
|
self._task_cache.task_record_hashes(
|
||||||
|
t.task_id, t.prompt_ids, start_logical_page
|
||||||
|
)
|
||||||
|
|
||||||
|
decoded: List[Task] = []
|
||||||
|
aborted: List[Task] = []
|
||||||
|
for t in tasks:
|
||||||
|
if t.task_id in prefilled_ids:
|
||||||
|
continue
|
||||||
|
if self._task_cache.task_extend(t.task_id, t.next_pos):
|
||||||
|
decoded.append(t)
|
||||||
|
else:
|
||||||
|
t.status = TaskStatus.ABORTED
|
||||||
|
aborted.append(t)
|
||||||
|
|
||||||
|
for backend, group in self._task_backend_groups(decoded):
|
||||||
|
backend_context = (
|
||||||
|
attn_backend(backend) if backend is not None else nullcontext()
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
backend_context,
|
||||||
|
self._metrics.record([t.task_id for t in group], "decode"),
|
||||||
|
):
|
||||||
|
step_out = self._executor.execute_decode(
|
||||||
|
group, return_logprobs=return_logprobs
|
||||||
|
)
|
||||||
|
for t, out in zip(group, step_out):
|
||||||
|
t.output_ids.append(out[0] if return_logprobs else out)
|
||||||
|
t.output_tokens += 1
|
||||||
|
t.advance_kv()
|
||||||
|
produced.append(t)
|
||||||
|
|
||||||
|
return produced, aborted
|
||||||
|
|
||||||
|
def _run_generation_loop(self):
|
||||||
|
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||||
|
try:
|
||||||
|
with self._backend_context():
|
||||||
|
while not self._stop_event.is_set():
|
||||||
|
finished = self._task_mgr.remove_finished_tasks(stop_ids)
|
||||||
|
for task in finished:
|
||||||
|
if task.status == TaskStatus.FINISHED:
|
||||||
|
self._task_cache.task_record_hashes(
|
||||||
|
task.task_id,
|
||||||
|
self._task_cache.task_cacheable_ids(
|
||||||
|
task.task_id, task.prompt_ids, task.output_ids
|
||||||
|
),
|
||||||
|
)
|
||||||
|
self._task_cache.task_free(task.task_id)
|
||||||
|
|
||||||
|
active = self._task_mgr.get_active_tasks()
|
||||||
|
available = self._task_mgr.max_batch_size - len(active)
|
||||||
|
if available > 0:
|
||||||
|
candidates = self._task_mgr.pull_candidates(available)
|
||||||
|
failed = []
|
||||||
|
for task in candidates:
|
||||||
|
if self._task_cache.task_alloc(
|
||||||
|
task.task_id, task.prompt_ids
|
||||||
|
):
|
||||||
|
self._task_mgr.activate(task)
|
||||||
|
else:
|
||||||
|
failed.append(task)
|
||||||
|
if failed:
|
||||||
|
self._task_mgr.return_to_waiting(failed)
|
||||||
|
|
||||||
|
if not self._task_mgr.has_work():
|
||||||
|
self._task_mgr.wait_for_tasks(timeout=1.0)
|
||||||
|
continue
|
||||||
|
|
||||||
|
active = self._task_mgr.get_active_tasks()
|
||||||
|
|
||||||
|
decoded, aborted = self._step(active)
|
||||||
|
|
||||||
|
for t in aborted:
|
||||||
|
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||||
|
|
||||||
|
for t in decoded:
|
||||||
|
new_text = t.decode_new_token(self._task_mgr.tokenizer)
|
||||||
|
if new_text:
|
||||||
|
self._task_mgr.invoke_callback(t.task_id, new_text)
|
||||||
|
if t.is_finished(stop_ids):
|
||||||
|
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
self._stop_event.set()
|
||||||
|
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
||||||
|
self._abort_and_clear(free_waiting=False)
|
||||||
|
|
||||||
|
def start(self):
|
||||||
|
if self._loop_thread is not None and self._loop_thread.is_alive():
|
||||||
|
return
|
||||||
|
self._stop_event.clear()
|
||||||
|
t = threading.Thread(target=self._run_generation_loop, daemon=True)
|
||||||
|
t.start()
|
||||||
|
self._loop_thread = t
|
||||||
|
|
||||||
|
def stop(self):
|
||||||
|
self._stop_event.set()
|
||||||
|
self._task_mgr.wake()
|
||||||
|
if self._loop_thread is not None:
|
||||||
|
self._loop_thread.join(timeout=2.0)
|
||||||
|
self._loop_thread = None
|
||||||
|
self._abort_and_clear(free_waiting=True)
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
def _abort_and_clear(self, free_waiting: bool):
|
||||||
|
"""Invoke STOP callbacks, release cache slots, and clear task queues."""
|
||||||
|
for task in self._task_mgr.get_active_tasks():
|
||||||
|
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||||
|
self._task_cache.task_free(task.task_id)
|
||||||
|
for task in self._task_mgr.get_waiting_tasks():
|
||||||
|
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||||
|
if free_waiting:
|
||||||
|
self._task_cache.task_free(task.task_id)
|
||||||
|
self._task_mgr.clear_queues()
|
||||||
|
|
||||||
|
def run_batch(
|
||||||
|
self,
|
||||||
|
prompt_ids_list: List[List[int]],
|
||||||
|
*,
|
||||||
|
max_tokens: Optional[int] = None,
|
||||||
|
temperature: float = 1.0,
|
||||||
|
top_p: float = 1.0,
|
||||||
|
top_k: int = 50,
|
||||||
|
frequency_penalty: float = 0.0,
|
||||||
|
rep_window: int = 64,
|
||||||
|
return_logprobs: bool = False,
|
||||||
|
) -> List[List[int]]:
|
||||||
|
"""Synchronous batch generation without the scheduler thread.
|
||||||
|
|
||||||
|
Accepts already-tokenized prompts (no string round-trip) and runs
|
||||||
|
prefill + decode to completion on the calling thread. Designed for
|
||||||
|
RL rollout, where logprobs of the behaviour policy must be collected
|
||||||
|
alongside generated tokens.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prompt_ids_list: ``B`` prompts, each a list of token IDs.
|
||||||
|
max_tokens: Maximum tokens to generate per prompt. ``None``
|
||||||
|
uses ``self.max_seq_len - len(prompt_ids)``.
|
||||||
|
temperature/top_p/top_k/frequency_penalty/rep_window: Sampling
|
||||||
|
parameters (uniform across the batch).
|
||||||
|
return_logprobs: If ``True``, return ``(token_ids, logprobs)``
|
||||||
|
tuples per prompt (logprobs aligned 1-to-1 with token_ids).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
``List[List[int]]`` of generated token IDs per prompt, or —
|
||||||
|
when ``return_logprobs`` is ``True`` —
|
||||||
|
``List[Tuple[List[int], List[float]]]``.
|
||||||
|
"""
|
||||||
|
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||||
|
seq_cap = self.max_seq_len
|
||||||
|
request_backend = get_backend(use_default=False)
|
||||||
|
|
||||||
|
tasks: List[Task] = []
|
||||||
|
for ids in prompt_ids_list:
|
||||||
|
if len(ids) >= seq_cap:
|
||||||
|
tasks.append(None)
|
||||||
|
continue
|
||||||
|
t_max = max_tokens
|
||||||
|
if t_max is None:
|
||||||
|
t_max = seq_cap - len(ids)
|
||||||
|
else:
|
||||||
|
t_max = min(t_max, seq_cap - len(ids))
|
||||||
|
if t_max <= 0:
|
||||||
|
tasks.append(None)
|
||||||
|
continue
|
||||||
|
task = Task(
|
||||||
|
task_id=f"batch_{uuid.uuid4().hex[:8]}",
|
||||||
|
prompt_ids=list(ids),
|
||||||
|
max_tokens=t_max,
|
||||||
|
temperature=temperature,
|
||||||
|
top_p=top_p,
|
||||||
|
top_k=top_k,
|
||||||
|
frequency_penalty=frequency_penalty,
|
||||||
|
rep_window=rep_window,
|
||||||
|
backend=request_backend,
|
||||||
|
)
|
||||||
|
if not self._task_cache.task_alloc(task.task_id, task.prompt_ids):
|
||||||
|
tasks.append(None)
|
||||||
|
continue
|
||||||
|
task.input_tokens = len(task.prompt_ids)
|
||||||
|
self._metrics.register(task.task_id)
|
||||||
|
tasks.append(task)
|
||||||
|
|
||||||
|
try:
|
||||||
|
live = [t for t in tasks if t is not None]
|
||||||
|
|
||||||
|
with self._backend_context():
|
||||||
|
while live:
|
||||||
|
decoded, _ = self._step(live, return_logprobs=return_logprobs)
|
||||||
|
live = [t for t in decoded if not t.is_finished(stop_ids)]
|
||||||
|
finally:
|
||||||
|
for t in tasks:
|
||||||
|
if t is not None:
|
||||||
|
self._metrics.mark_finished(
|
||||||
|
t.task_id, t.input_tokens, t.output_tokens
|
||||||
|
)
|
||||||
|
self._task_cache.task_free(t.task_id)
|
||||||
|
|
||||||
|
results: List[Any] = []
|
||||||
|
for t in tasks:
|
||||||
|
if t is None:
|
||||||
|
results.append(([], []) if return_logprobs else [])
|
||||||
|
elif return_logprobs:
|
||||||
|
results.append((list(t.output_ids), list(t.output_logprobs)))
|
||||||
|
else:
|
||||||
|
results.append(list(t.output_ids))
|
||||||
|
return results
|
||||||
@@ -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,152 @@
|
|||||||
|
"""Pre-allocated buffers for the inference decode hot path.
|
||||||
|
|
||||||
|
Mirrors FlashInfer / SGLang's global workspace pattern: all per-step tensors
|
||||||
|
are allocated eagerly at init (nothing is lazy), so the decode step
|
||||||
|
reads/writes fixed-address tensors with zero ``torch.empty`` calls during
|
||||||
|
the hot loop — a prerequisite for CUDA-graph capture.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
_MAX_SPLITS = 32
|
||||||
|
Q_TILE_ROWS = 64
|
||||||
|
|
||||||
|
|
||||||
|
class InferenceWorkspace:
|
||||||
|
"""Reusable fixed-shape per-step buffers for decode.
|
||||||
|
|
||||||
|
Families of buffers, all sized to ``max_batch_size`` / ``max_seq_len``
|
||||||
|
and sliced via views each step:
|
||||||
|
|
||||||
|
- ``decode_mask``: a ``[B, 1, total_len]`` validity mask, the RHS
|
||||||
|
``arange`` pre-computed so only a single ``torch.ge(out=)`` runs per
|
||||||
|
step.
|
||||||
|
- ``input_ids``: per-step token IDs filled from host (pinned, double-
|
||||||
|
buffered so an in-flight async H2D copy never races the next fill).
|
||||||
|
- KV-cache bind metadata (``req_pool_indices``, ``seq_lens``,
|
||||||
|
``kv_indptr``, ``inc``, ``out_cache_loc``), written by
|
||||||
|
``PagePool.bind_tasks`` when the Executor passes this workspace.
|
||||||
|
- ``decode_o_part`` / ``decode_ml_part``: split-KV partial result buffers
|
||||||
|
(mirrors FlashInfer's workspace). One global alloc, reused by every
|
||||||
|
decode step across all layers. Sliced views are passed to the CUDA
|
||||||
|
attention kernel so its internal ``torch.empty`` hot-path alloc goes
|
||||||
|
through a stable address (CUDA-graph capturable).
|
||||||
|
|
||||||
|
No re-allocation while the server's bounds are respected.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
max_batch_size: int,
|
||||||
|
max_seq_len: int,
|
||||||
|
max_q_heads: int,
|
||||||
|
head_dim: int,
|
||||||
|
device: torch.device,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
):
|
||||||
|
self.max_batch_size = max_batch_size
|
||||||
|
self.max_seq_len = max_seq_len
|
||||||
|
self.max_q_heads = max_q_heads
|
||||||
|
self.head_dim = head_dim
|
||||||
|
self.device = device
|
||||||
|
self.dtype = dtype
|
||||||
|
|
||||||
|
# ``position_ids[:, None, None] >= arange`` RHS, reused every step.
|
||||||
|
self.arange = torch.arange(max_seq_len, device=device)
|
||||||
|
# Decode validity mask: [max_batch, 1, max_seq_len] bool.
|
||||||
|
self.input_mask = torch.empty(
|
||||||
|
(max_batch_size, 1, max_seq_len), dtype=torch.bool, device=device
|
||||||
|
)
|
||||||
|
|
||||||
|
# Per-step token IDs. Values come from host Python lists every
|
||||||
|
# step, so the device buffer is pre-allocated (stable address for
|
||||||
|
# CUDA-graph capture) and filled via a host staging buffer. A
|
||||||
|
# double buffer keeps a copy in flight from being overwritten by
|
||||||
|
# the next fill.
|
||||||
|
self.input_ids = torch.empty((max_batch_size,), dtype=torch.long, device=device)
|
||||||
|
self._pin = [
|
||||||
|
torch.empty((max_batch_size,), dtype=torch.long),
|
||||||
|
torch.empty((max_batch_size,), dtype=torch.long),
|
||||||
|
]
|
||||||
|
self._pin_idx = 0
|
||||||
|
|
||||||
|
# KV-cache bind metadata (fixed shape, written by ``PagePool.bind_tasks``
|
||||||
|
# when the Executor passes this workspace). Stable addresses make the
|
||||||
|
# decode forward CUDA-graph capturable.
|
||||||
|
self.req_pool_indices = torch.empty(
|
||||||
|
(max_batch_size,), dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
self.seq_lens = torch.empty((max_batch_size,), dtype=torch.long, device=device)
|
||||||
|
self.kv_indptr = torch.empty(
|
||||||
|
(max_batch_size + 1,), dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
self.qo_indptr = torch.empty(
|
||||||
|
(max_batch_size + 1,), dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
max_q_tiles = max_batch_size * ((max_seq_len + Q_TILE_ROWS - 1) // Q_TILE_ROWS)
|
||||||
|
self.q_tile_to_batch = torch.empty(
|
||||||
|
(max_q_tiles,), dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
self.q_tile_to_index = torch.empty(
|
||||||
|
(max_q_tiles,), dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
self.inc = torch.arange(max_batch_size + 1, dtype=torch.int32, device=device)
|
||||||
|
self.out_cache_loc = torch.empty(
|
||||||
|
(max_batch_size, 1), dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
|
||||||
|
# Per-step position IDs (must be at a fixed address for CUDA-graph capture).
|
||||||
|
self.position_ids = torch.empty(
|
||||||
|
(max_batch_size,), dtype=torch.long, device=device
|
||||||
|
)
|
||||||
|
|
||||||
|
# Split-KV partial-result buffers for decode (persistent, one global
|
||||||
|
# alloc per process — mirrors FlashInfer's workspace pattern).
|
||||||
|
# Shape: [max_batch_size, max_q_heads, _MAX_SPLITS, head_dim] (o_part)
|
||||||
|
# [max_batch_size, max_q_heads, _MAX_SPLITS, 2] (ml_part)
|
||||||
|
self.decode_o_part = torch.empty(
|
||||||
|
(max_batch_size, max_q_heads, _MAX_SPLITS, head_dim),
|
||||||
|
dtype=torch.float32,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
self.decode_ml_part = torch.empty(
|
||||||
|
(max_batch_size, max_q_heads, _MAX_SPLITS, 2),
|
||||||
|
dtype=torch.float32,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Decode output buffer (graph-safe pre-alloc). Shape matches the
|
||||||
|
# decode kernel's output: [batch, q_head, head_dim].
|
||||||
|
self.decode_out = torch.empty(
|
||||||
|
(max_batch_size, max_q_heads, head_dim),
|
||||||
|
dtype=dtype,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
|
||||||
|
def fill_input_ids(self, ids: "list[int]") -> Tensor:
|
||||||
|
"""Write ``ids`` into the device buffer and return ``[B]``.
|
||||||
|
|
||||||
|
Host values are staged through the double buffer and copied into the
|
||||||
|
stable device buffer (``copy_`` without pinning is synchronous, so
|
||||||
|
the alternating buffers guard against an in-flight transfer).
|
||||||
|
"""
|
||||||
|
b = len(ids)
|
||||||
|
pin = self._pin[self._pin_idx]
|
||||||
|
self._pin_idx ^= 1
|
||||||
|
for i, v in enumerate(ids):
|
||||||
|
pin[i] = v
|
||||||
|
self.input_ids[:b].copy_(pin[:b])
|
||||||
|
return self.input_ids[:b]
|
||||||
|
|
||||||
|
def decode_mask(self, position_ids: Tensor, total_len: int) -> Tensor:
|
||||||
|
"""Return the ``[B, 1, total_len]`` validity mask for this step.
|
||||||
|
|
||||||
|
Written into the pre-allocated buffer via ``torch.ge(out=)`` — no
|
||||||
|
new tensor is allocated. ``position_ids`` is the current step's
|
||||||
|
``[B]`` positions; ``total_len`` must not exceed ``max_seq_len``.
|
||||||
|
"""
|
||||||
|
b = position_ids.size(0)
|
||||||
|
out = self.input_mask[:b, :, :total_len]
|
||||||
|
torch.ge(position_ids[:, None, None], self.arange[:total_len], out=out)
|
||||||
|
return out
|
||||||
@@ -0,0 +1,35 @@
|
|||||||
|
import logging
|
||||||
|
import os
|
||||||
|
|
||||||
|
|
||||||
|
class _DistributedContextFilter(logging.Filter):
|
||||||
|
def filter(self, record: logging.LogRecord) -> bool:
|
||||||
|
record.rank = os.environ.get("RANK", "0")
|
||||||
|
record.world_size = os.environ.get("WORLD_SIZE", "1")
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
def setup_logging(level: str = "INFO"):
|
||||||
|
"""Attach a StreamHandler to the ``astrai`` logger (idempotent).
|
||||||
|
|
||||||
|
Call once per process at the top of CLI scripts.
|
||||||
|
Set ``ASTR_LOG_LEVEL`` env var to override the default level.
|
||||||
|
|
||||||
|
Level names: ``DEBUG``, ``INFO``, ``WARNING``, ``ERROR``, ``CRITICAL``.
|
||||||
|
``DEBUG`` enables per-step prefill/decode timing logs
|
||||||
|
(:func:`astrai.inference.runtime.executor.timed`).
|
||||||
|
"""
|
||||||
|
logger = logging.getLogger("astrai")
|
||||||
|
if logger.handlers:
|
||||||
|
return
|
||||||
|
level_name = os.environ.get("ASTR_LOG_LEVEL", level).upper()
|
||||||
|
logger.setLevel(getattr(logging, level_name, logging.INFO))
|
||||||
|
handler = logging.StreamHandler()
|
||||||
|
handler.addFilter(_DistributedContextFilter())
|
||||||
|
handler.setFormatter(
|
||||||
|
logging.Formatter(
|
||||||
|
"%(asctime)s | %(levelname)-8s | rank=%(rank)2s/%(world_size)-2s | %(name)-32s | %(message)s",
|
||||||
|
datefmt="%Y-%m-%d %H:%M:%S",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
logger.addHandler(handler)
|
||||||
@@ -0,0 +1,35 @@
|
|||||||
|
from astrai.model.automodel import AutoModel
|
||||||
|
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.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
|
||||||
|
"AutoRegressiveLM",
|
||||||
|
"EmbeddingEncoder",
|
||||||
|
"AutoModel",
|
||||||
|
# LoRA
|
||||||
|
"LoRAConfig",
|
||||||
|
"inject_lora",
|
||||||
|
"merge_lora",
|
||||||
|
"save_lora",
|
||||||
|
"load_lora",
|
||||||
|
]
|
||||||
@@ -0,0 +1,130 @@
|
|||||||
|
"""
|
||||||
|
AutoModel base class for model loading and saving.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Union
|
||||||
|
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
from astrai.config.model_config import BaseModelConfig, ConfigFactory
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.serialization import (
|
||||||
|
HF_MODEL_TYPES,
|
||||||
|
adapt_config,
|
||||||
|
convert_hf_weights,
|
||||||
|
load_model_config,
|
||||||
|
load_model_weights,
|
||||||
|
looks_like_hf_state_dict,
|
||||||
|
save_model,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def _disable_random_init(enable: bool = True):
|
||||||
|
if not enable:
|
||||||
|
yield
|
||||||
|
return
|
||||||
|
|
||||||
|
names = (
|
||||||
|
"xavier_normal_",
|
||||||
|
"xavier_uniform_",
|
||||||
|
"kaiming_normal_",
|
||||||
|
"kaiming_uniform_",
|
||||||
|
"zeros_",
|
||||||
|
"ones_",
|
||||||
|
"constant_",
|
||||||
|
"normal_",
|
||||||
|
"uniform_",
|
||||||
|
)
|
||||||
|
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:
|
||||||
|
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):
|
||||||
|
"""Model base class with loading/saving and generation."""
|
||||||
|
|
||||||
|
def __init__(self, config: BaseModelConfig):
|
||||||
|
super().__init__()
|
||||||
|
self.config = config
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_pretrained(
|
||||||
|
cls,
|
||||||
|
path: Union[str, Path],
|
||||||
|
disable_random_init: bool = True,
|
||||||
|
strict: bool = True,
|
||||||
|
weights_format: str = "auto",
|
||||||
|
) -> nn.Module:
|
||||||
|
"""Load a model directory.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
path: Directory containing ``config.json`` and optionally
|
||||||
|
``model.safetensors``.
|
||||||
|
disable_random_init: Replace parameter initializers with no-ops
|
||||||
|
while building the model.
|
||||||
|
strict: Passed to ``load_state_dict``.
|
||||||
|
weights_format: ``"auto"`` detects HuggingFace checkpoints
|
||||||
|
(LLaMA-style keys and ``model_type``) and converts them;
|
||||||
|
``"astrai"`` skips conversion; ``"hf"`` forces it.
|
||||||
|
"""
|
||||||
|
if weights_format not in ("auto", "astrai", "hf"):
|
||||||
|
raise ValueError(
|
||||||
|
f"weights_format must be one of 'auto', 'astrai', 'hf', "
|
||||||
|
f"got {weights_format!r}"
|
||||||
|
)
|
||||||
|
|
||||||
|
model_path = Path(path)
|
||||||
|
|
||||||
|
config_path = model_path / "config.json"
|
||||||
|
if not config_path.exists():
|
||||||
|
raise FileNotFoundError(f"Config file not found: {config_path}")
|
||||||
|
|
||||||
|
raw = load_model_config(str(model_path))
|
||||||
|
is_hf_config = weights_format == "hf" or (
|
||||||
|
weights_format == "auto" and raw.get("model_type") in HF_MODEL_TYPES
|
||||||
|
)
|
||||||
|
if is_hf_config:
|
||||||
|
raw = adapt_config(raw)
|
||||||
|
|
||||||
|
config = ConfigFactory.load(raw)
|
||||||
|
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)
|
||||||
|
|
||||||
|
weights_path = model_path / "model.safetensors"
|
||||||
|
index_path = model_path / "model.safetensors.index.json"
|
||||||
|
if weights_path.exists() or index_path.exists():
|
||||||
|
state_dict = load_model_weights(str(model_path))
|
||||||
|
is_hf_weights = is_hf_config or (
|
||||||
|
weights_format == "auto" and looks_like_hf_state_dict(state_dict)
|
||||||
|
)
|
||||||
|
if is_hf_weights:
|
||||||
|
state_dict = convert_hf_weights(state_dict, config)
|
||||||
|
model.load_state_dict(state_dict, strict=strict)
|
||||||
|
|
||||||
|
return model
|
||||||
|
|
||||||
|
def save_pretrained(
|
||||||
|
self,
|
||||||
|
save_directory: Union[str, Path],
|
||||||
|
):
|
||||||
|
save_model(
|
||||||
|
config=self.config.to_dict(),
|
||||||
|
state_dict=self.state_dict(),
|
||||||
|
save_directory=str(save_directory),
|
||||||
|
)
|
||||||
@@ -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,172 @@
|
|||||||
|
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]
|
||||||
|
|
||||||
|
|
||||||
|
@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) -> FFNOutput:
|
||||||
|
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
|
||||||
@@ -0,0 +1,152 @@
|
|||||||
|
from typing import Any, Dict, Mapping, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
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(
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
@ModelFactory.register("autoregressive_lm")
|
||||||
|
class AutoRegressiveLM(AutoModel):
|
||||||
|
"""Autoregressive language model with paged KV cache."""
|
||||||
|
|
||||||
|
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(
|
||||||
|
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.lm_head = Linear(config.hidden_size, config.vocab_size)
|
||||||
|
|
||||||
|
if self.config.tie_word_embeddings is True:
|
||||||
|
self.lm_head.weight = self.embed_tokens.weight
|
||||||
|
|
||||||
|
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):
|
||||||
|
lm_head_key = "lm_head.weight"
|
||||||
|
embed_key = "embed_tokens.weight"
|
||||||
|
|
||||||
|
state_dict = dict(state_dict)
|
||||||
|
|
||||||
|
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]
|
||||||
|
else:
|
||||||
|
if lm_head_key not in state_dict and embed_key in state_dict:
|
||||||
|
# clone to avoid sharing gradients
|
||||||
|
state_dict[lm_head_key] = torch.clone(state_dict[embed_key])
|
||||||
|
|
||||||
|
return super().load_state_dict(state_dict, strict, assign)
|
||||||
|
|
||||||
|
def state_dict(self, destination=None, prefix="", keep_vars=False):
|
||||||
|
state_dict = super().state_dict(
|
||||||
|
destination=destination, prefix=prefix, keep_vars=keep_vars
|
||||||
|
)
|
||||||
|
|
||||||
|
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]
|
||||||
|
|
||||||
|
return state_dict
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
input_ids: Tensor,
|
||||||
|
input_mask: Optional[Tensor] = None,
|
||||||
|
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, 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:
|
||||||
|
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)
|
||||||
|
|
||||||
|
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])
|
||||||
@@ -0,0 +1,39 @@
|
|||||||
|
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,
|
||||||
|
get_world_size,
|
||||||
|
only_on_rank,
|
||||||
|
setup_parallel,
|
||||||
|
spawn_parallel_fn,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"get_world_size",
|
||||||
|
"get_rank",
|
||||||
|
"get_current_device",
|
||||||
|
"only_on_rank",
|
||||||
|
"setup_parallel",
|
||||||
|
"spawn_parallel_fn",
|
||||||
|
"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
|
||||||
@@ -0,0 +1,281 @@
|
|||||||
|
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, 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"]
|
||||||
|
|
||||||
|
|
||||||
|
def get_world_size() -> int:
|
||||||
|
if dist.is_available() and dist.is_initialized():
|
||||||
|
return dist.get_world_size()
|
||||||
|
else:
|
||||||
|
return 1
|
||||||
|
|
||||||
|
|
||||||
|
def get_rank() -> int:
|
||||||
|
if dist.is_available() and dist.is_initialized():
|
||||||
|
return dist.get_rank()
|
||||||
|
else:
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def setup_parallel(
|
||||||
|
rank: int,
|
||||||
|
world_size: int,
|
||||||
|
local_rank: int,
|
||||||
|
backend: str = "nccl",
|
||||||
|
master_addr: str = "localhost",
|
||||||
|
master_port: str = "29500",
|
||||||
|
device_type: str = "cuda",
|
||||||
|
):
|
||||||
|
|
||||||
|
if dist.is_available() and dist.is_initialized():
|
||||||
|
yield dist.group.WORLD
|
||||||
|
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, local_rank)
|
||||||
|
|
||||||
|
os.environ["MASTER_ADDR"] = master_addr
|
||||||
|
os.environ["MASTER_PORT"] = master_port
|
||||||
|
os.environ["LOCAL_RANK"] = str(local_rank)
|
||||||
|
os.environ["WORLD_SIZE"] = str(world_size)
|
||||||
|
os.environ["LOCAL_DEVICE"] = str(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():
|
||||||
|
torch.cuda.set_device(device_id)
|
||||||
|
elif backend == "ccl" and hasattr(torch, "xpu") and torch.xpu.is_available():
|
||||||
|
torch.xpu.set_device(device_id)
|
||||||
|
|
||||||
|
yield dist.group.WORLD
|
||||||
|
finally:
|
||||||
|
if dist.is_initialized():
|
||||||
|
dist.destroy_process_group()
|
||||||
|
|
||||||
|
|
||||||
|
def only_on_rank(rank, sync=False):
|
||||||
|
"""
|
||||||
|
decorator to run a function only on a specific rank.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def decorator(func):
|
||||||
|
@wraps(func)
|
||||||
|
def wrapper(*args, **kwargs):
|
||||||
|
ret_args = None
|
||||||
|
if get_rank() == rank:
|
||||||
|
ret_args = func(*args, **kwargs)
|
||||||
|
|
||||||
|
if sync and dist.is_available() and dist.is_initialized():
|
||||||
|
dist.barrier()
|
||||||
|
|
||||||
|
return ret_args
|
||||||
|
|
||||||
|
return wrapper
|
||||||
|
|
||||||
|
return decorator
|
||||||
|
|
||||||
|
|
||||||
|
def _run_single_rank(
|
||||||
|
rank: int,
|
||||||
|
world_size: int,
|
||||||
|
backend: str,
|
||||||
|
master_addr: str,
|
||||||
|
master_port: str,
|
||||||
|
device_type: str,
|
||||||
|
func: Callable,
|
||||||
|
kwargs: dict,
|
||||||
|
):
|
||||||
|
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,
|
||||||
|
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)
|
||||||
|
|
||||||
|
|
||||||
|
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 _is_external_launcher() -> bool:
|
||||||
|
"""Whether an external launcher (torchrun/elastic/manual env) started us."""
|
||||||
|
if dist.is_torchelastic_launched():
|
||||||
|
return True
|
||||||
|
if "LOCAL_WORLD_SIZE" in os.environ:
|
||||||
|
return True
|
||||||
|
if "RANK" in os.environ and "WORLD_SIZE" in os.environ:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def spawn_parallel_fn(
|
||||||
|
func: Callable,
|
||||||
|
world_size: int,
|
||||||
|
backend: str = "nccl",
|
||||||
|
master_addr: str = "localhost",
|
||||||
|
master_port: Optional[str] = None,
|
||||||
|
device_type: str = "cuda",
|
||||||
|
start_method: str = "spawn",
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
if master_port is None:
|
||||||
|
master_port = find_free_port()
|
||||||
|
if _is_external_launcher():
|
||||||
|
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): ...
|
||||||
@@ -0,0 +1,53 @@
|
|||||||
|
"""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,
|
||||||
|
)
|
||||||
|
from astrai.serialization.hf_adapter import (
|
||||||
|
HF_MODEL_TYPES,
|
||||||
|
adapt_config,
|
||||||
|
convert_hf_config,
|
||||||
|
convert_hf_weights,
|
||||||
|
looks_like_hf_state_dict,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"Checkpoint",
|
||||||
|
"HF_MODEL_TYPES",
|
||||||
|
"adapt_config",
|
||||||
|
"convert_hf_config",
|
||||||
|
"convert_hf_weights",
|
||||||
|
"looks_like_hf_state_dict",
|
||||||
|
"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,209 @@
|
|||||||
|
"""Model checkpoint serialization helpers."""
|
||||||
|
|
||||||
|
import io
|
||||||
|
import json
|
||||||
|
import time
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Callable, 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 _broadcast_load(loader: Callable[[], dict], broadcast: bool) -> dict:
|
||||||
|
"""Load on rank 0 and broadcast the object to all ranks."""
|
||||||
|
if not broadcast or not dist.is_initialized():
|
||||||
|
return loader()
|
||||||
|
rank = get_rank()
|
||||||
|
if rank == 0:
|
||||||
|
data = loader()
|
||||||
|
else:
|
||||||
|
data = {}
|
||||||
|
tmp = [data]
|
||||||
|
dist.broadcast_object_list(tmp, src=0)
|
||||||
|
return tmp[0]
|
||||||
|
|
||||||
|
|
||||||
|
def load_safetensors(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||||
|
return _broadcast_load(lambda: st.load_file(str(path)), broadcast)
|
||||||
|
|
||||||
|
|
||||||
|
def save_json(data: dict, path: Union[str, Path]):
|
||||||
|
with open(str(path), "w") as f:
|
||||||
|
json.dump(data, f, indent=2)
|
||||||
|
|
||||||
|
|
||||||
|
def load_json(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||||
|
return _broadcast_load(lambda: json.loads(Path(path).read_text()), broadcast)
|
||||||
|
|
||||||
|
|
||||||
|
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:
|
||||||
|
save_path = Path(save_directory)
|
||||||
|
weights_file = save_path / _WEIGHTS_FILE
|
||||||
|
if weights_file.exists():
|
||||||
|
return load_state_dict(weights_file)
|
||||||
|
|
||||||
|
index_path = save_path / "model.safetensors.index.json"
|
||||||
|
if index_path.exists():
|
||||||
|
index = load_json(index_path)
|
||||||
|
weight_map = index.get("weight_map", {})
|
||||||
|
state_dict = {}
|
||||||
|
for shard in sorted(set(weight_map.values())):
|
||||||
|
state_dict.update(load_state_dict(save_path / shard))
|
||||||
|
return state_dict
|
||||||
|
|
||||||
|
raise FileNotFoundError(f"No model weights found in {save_directory}")
|
||||||
|
|
||||||
|
|
||||||
|
def load_state_dict(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||||
|
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)
|
||||||
|
|
||||||
|
weights_path = save_path / _WEIGHTS_FILE
|
||||||
|
index_path = save_path / "model.safetensors.index.json"
|
||||||
|
if weights_path.exists() or index_path.exists():
|
||||||
|
state_dict = load_model_weights(save_dir)
|
||||||
|
config = {}
|
||||||
|
config_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,271 @@
|
|||||||
|
"""HuggingFace checkpoint adaptation for LLaMA-style decoder models.
|
||||||
|
|
||||||
|
AstrAI stores weights with its own key names (``layers.<i>.input_norm``,
|
||||||
|
``layers.<i>.mlp.gate``), while HuggingFace decoder-only checkpoints use
|
||||||
|
``model.layers.<i>.input_layernorm`` / ``model.layers.<i>.mlp.gate_proj``.
|
||||||
|
This module translates HF configs and state dicts so external checkpoints
|
||||||
|
can be loaded directly.
|
||||||
|
|
||||||
|
Supported families (LLaMA layout, dense and MoE):
|
||||||
|
- dense FFN: llama, mistral, qwen2, gemma, gemma2, phi3
|
||||||
|
- MoE FFN (Mixtral / Qwen2-MoE / DeepSeek-V3 layout): router
|
||||||
|
``mlp.gate``, routed experts ``mlp.experts.<j>``, shared experts
|
||||||
|
``mlp.shared_experts.<j>``
|
||||||
|
|
||||||
|
Not supported:
|
||||||
|
- MLA attention (DeepSeek-V2/V3 ``kv_a_proj_with_mqa``) uses a different
|
||||||
|
KV factorization and cannot be converted numerically.
|
||||||
|
- Attention/MLP bias (``attention_bias`` / ``mlp_bias``) — AstrAI
|
||||||
|
projections are bias-free.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import re
|
||||||
|
from typing import Any, Dict, Mapping
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.config.base import BaseConfig
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
HF_MODEL_TYPES = frozenset(
|
||||||
|
{
|
||||||
|
"llama",
|
||||||
|
"mistral",
|
||||||
|
"mixtral",
|
||||||
|
"qwen2",
|
||||||
|
"qwen2_moe",
|
||||||
|
"gemma",
|
||||||
|
"gemma2",
|
||||||
|
"phi3",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
_EMBED = re.compile(r"^model\.embed_tokens\.weight$")
|
||||||
|
_ATTN = re.compile(r"^model\.layers\.(\d+)\.self_attn\.(q|k|v|o)_proj\.(weight|bias)$")
|
||||||
|
_Q_NORM = re.compile(r"^model\.layers\.(\d+)\.self_attn\.q_norm\.weight$")
|
||||||
|
_K_NORM = re.compile(r"^model\.layers\.(\d+)\.self_attn\.k_norm\.weight$")
|
||||||
|
_INPUT_NORM = re.compile(r"^model\.layers\.(\d+)\.input_layernorm\.weight$")
|
||||||
|
_POST_NORM = re.compile(r"^model\.layers\.(\d+)\.post_attention_layernorm\.weight$")
|
||||||
|
_FINAL_NORM = re.compile(r"^model\.norm\.weight$")
|
||||||
|
_LM_HEAD = re.compile(r"^lm_head\.weight$")
|
||||||
|
_DENSE_MLP = re.compile(
|
||||||
|
r"^model\.layers\.(\d+)\.mlp\.(gate|up|down)_proj\.(weight|bias)$"
|
||||||
|
)
|
||||||
|
_MOE_ROUTER = re.compile(r"^model\.layers\.(\d+)\.mlp\.gate\.weight$")
|
||||||
|
_MOE_EXPERTS = re.compile(
|
||||||
|
r"^model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.(gate|up|down)_proj\.(weight|bias)$"
|
||||||
|
)
|
||||||
|
_MOE_SHARED = re.compile(
|
||||||
|
r"^model\.layers\.(\d+)\.mlp\.shared_expert(?:s)?\.(\d+)\."
|
||||||
|
r"(gate|up|down)_proj\.(weight|bias)$"
|
||||||
|
)
|
||||||
|
|
||||||
|
_ASTR_PREFIXES = ("embed_tokens.", "layers.", "norm.", "lm_head.")
|
||||||
|
|
||||||
|
|
||||||
|
def looks_like_hf_state_dict(state_dict: Mapping[str, Any]) -> bool:
|
||||||
|
"""Return True if *state_dict* uses HuggingFace key names."""
|
||||||
|
return any(
|
||||||
|
key.startswith("model.")
|
||||||
|
or "self_attn." in key
|
||||||
|
or "input_layernorm" in key
|
||||||
|
or "mlp.experts." in key
|
||||||
|
for key in state_dict
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_dense_mlp_layer(config: BaseConfig, layer_id: int) -> bool:
|
||||||
|
"""Return whether a layer uses dense MLP instead of routed experts."""
|
||||||
|
if getattr(config, "ffn_type", "mlp") != "moe":
|
||||||
|
return True
|
||||||
|
mlp_only = getattr(config, "mlp_only_layers", None) or []
|
||||||
|
if layer_id in mlp_only:
|
||||||
|
return True
|
||||||
|
step = getattr(config, "decoder_sparse_step", 1) or 1
|
||||||
|
return step > 1 and (layer_id + 1) % step != 0
|
||||||
|
|
||||||
|
|
||||||
|
def adapt_config(raw: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
"""Translate *raw* for AstrAI if it looks like an HF model config."""
|
||||||
|
if raw.get("model_type") in HF_MODEL_TYPES:
|
||||||
|
return convert_hf_config(raw)
|
||||||
|
return raw
|
||||||
|
|
||||||
|
|
||||||
|
def convert_hf_config(raw: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
"""Convert an HF LLaMA-style config dict to AstrAI field names."""
|
||||||
|
if raw.get("attention_bias") or raw.get("mlp_bias"):
|
||||||
|
raise NotImplementedError(
|
||||||
|
"attention_bias / mlp_bias checkpoints are not supported; "
|
||||||
|
"AstrAI projections are bias-free"
|
||||||
|
)
|
||||||
|
|
||||||
|
cfg: Dict[str, Any] = {}
|
||||||
|
for key in (
|
||||||
|
"vocab_size",
|
||||||
|
"hidden_size",
|
||||||
|
"num_hidden_layers",
|
||||||
|
"intermediate_size",
|
||||||
|
"rms_norm_eps",
|
||||||
|
"tie_word_embeddings",
|
||||||
|
"max_position_embeddings",
|
||||||
|
"rope_theta",
|
||||||
|
"rope_scaling",
|
||||||
|
"num_attention_heads",
|
||||||
|
"num_key_value_heads",
|
||||||
|
"use_qk_norm",
|
||||||
|
"use_gated_attention",
|
||||||
|
"kv_lora_rank",
|
||||||
|
"qk_nope_head_dim",
|
||||||
|
"qk_rope_head_dim",
|
||||||
|
"moe_intermediate_size",
|
||||||
|
"shared_expert_intermediate_size",
|
||||||
|
"topk_method",
|
||||||
|
"norm_topk_prob",
|
||||||
|
"moe_aux_loss_coef",
|
||||||
|
"decoder_sparse_step",
|
||||||
|
"mlp_only_layers",
|
||||||
|
"neftune_alpha",
|
||||||
|
):
|
||||||
|
if key in raw:
|
||||||
|
cfg[key] = raw[key]
|
||||||
|
|
||||||
|
if "qk_norm" in raw and "use_qk_norm" not in cfg:
|
||||||
|
cfg["use_qk_norm"] = raw["qk_norm"]
|
||||||
|
if (
|
||||||
|
raw.get("model_type") in ("gemma", "gemma2")
|
||||||
|
and "use_qk_norm" not in cfg
|
||||||
|
and "qk_norm" not in raw
|
||||||
|
):
|
||||||
|
# Gemma/Gemma2 always apply RMSNorm to Q and K before attention.
|
||||||
|
cfg["use_qk_norm"] = True
|
||||||
|
|
||||||
|
n_heads = raw.get("num_attention_heads")
|
||||||
|
if cfg.get("num_key_value_heads") is None and n_heads is not None:
|
||||||
|
cfg["num_key_value_heads"] = n_heads
|
||||||
|
|
||||||
|
if raw.get("head_dim") is not None and n_heads and raw.get("hidden_size"):
|
||||||
|
expected = raw["hidden_size"] // n_heads
|
||||||
|
if raw["head_dim"] != expected:
|
||||||
|
raise NotImplementedError(
|
||||||
|
f"HF head_dim={raw['head_dim']} differs from the computed "
|
||||||
|
f"head dim {expected}; AstrAI derives head_dim from "
|
||||||
|
"hidden_size / num_attention_heads"
|
||||||
|
)
|
||||||
|
|
||||||
|
if "kv_lora_rank" in raw:
|
||||||
|
cfg["attn_type"] = "mla"
|
||||||
|
|
||||||
|
n_experts = raw.get("num_local_experts") or raw.get("n_routed_experts")
|
||||||
|
if n_experts:
|
||||||
|
cfg["ffn_type"] = "moe"
|
||||||
|
cfg["n_routed_experts"] = n_experts
|
||||||
|
if "num_experts_per_tok" in raw:
|
||||||
|
cfg["n_activated_experts"] = raw["num_experts_per_tok"]
|
||||||
|
if "n_activated_experts" in raw:
|
||||||
|
cfg["n_activated_experts"] = raw["n_activated_experts"]
|
||||||
|
if "n_shared_experts" in raw:
|
||||||
|
cfg["n_shared_experts"] = raw["n_shared_experts"]
|
||||||
|
else:
|
||||||
|
# Mixtral has no shared experts; AstrAI defaults to one.
|
||||||
|
cfg["n_shared_experts"] = 0
|
||||||
|
if cfg.get("moe_intermediate_size") is None and "intermediate_size" in raw:
|
||||||
|
# MoE configs store the per-expert FFN size in intermediate_size.
|
||||||
|
cfg["moe_intermediate_size"] = raw["intermediate_size"]
|
||||||
|
first_k_dense = raw.get("first_k_dense_replace")
|
||||||
|
if isinstance(first_k_dense, int) and first_k_dense > 0:
|
||||||
|
cfg["mlp_only_layers"] = list(range(first_k_dense))
|
||||||
|
cfg["decoder_sparse_step"] = 1
|
||||||
|
|
||||||
|
cfg["model_type"] = "autoregressive_lm"
|
||||||
|
return cfg
|
||||||
|
|
||||||
|
|
||||||
|
def convert_hf_weights(
|
||||||
|
state_dict: Mapping[str, Any],
|
||||||
|
config: BaseConfig,
|
||||||
|
) -> Dict[str, torch.Tensor]:
|
||||||
|
"""Rename HF state dict keys to AstrAI names.
|
||||||
|
|
||||||
|
Keys that are already AstrAI-style pass through unchanged; unmapped
|
||||||
|
HF keys are dropped with a warning. Use with ``strict=True`` to fail
|
||||||
|
loudly when the checkpoint does not match the config.
|
||||||
|
"""
|
||||||
|
if getattr(config, "attn_type", "gqa") == "mla":
|
||||||
|
if any("kv_a_proj_with_mqa" in key for key in state_dict):
|
||||||
|
raise NotImplementedError(
|
||||||
|
"MLA attention (DeepSeek-V2/V3 kv_a_proj_with_mqa) uses a "
|
||||||
|
"different KV factorization and cannot be converted"
|
||||||
|
)
|
||||||
|
|
||||||
|
ffn_type = getattr(config, "ffn_type", "mlp")
|
||||||
|
converted: Dict[str, torch.Tensor] = {}
|
||||||
|
skipped: list[str] = []
|
||||||
|
for key, tensor in state_dict.items():
|
||||||
|
if key.startswith(_ASTR_PREFIXES):
|
||||||
|
converted[key] = tensor
|
||||||
|
continue
|
||||||
|
|
||||||
|
new_key = None
|
||||||
|
if ffn_type == "moe":
|
||||||
|
m = _MOE_ROUTER.match(key)
|
||||||
|
if m:
|
||||||
|
new_key = f"layers.{m.group(1)}.mlp.router.weight"
|
||||||
|
else:
|
||||||
|
m = _MOE_EXPERTS.match(key)
|
||||||
|
if m:
|
||||||
|
new_key = (
|
||||||
|
f"layers.{m.group(1)}.mlp.routed_experts.{m.group(2)}."
|
||||||
|
f"{m.group(3)}.{m.group(4)}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
m = _MOE_SHARED.match(key)
|
||||||
|
if m:
|
||||||
|
new_key = (
|
||||||
|
f"layers.{m.group(1)}.mlp.shared_experts.{m.group(2)}."
|
||||||
|
f"{m.group(3)}.{m.group(4)}"
|
||||||
|
)
|
||||||
|
if new_key is None:
|
||||||
|
m = _DENSE_MLP.match(key)
|
||||||
|
if m and _is_dense_mlp_layer(config, int(m.group(1))):
|
||||||
|
new_key = f"layers.{m.group(1)}.mlp.{m.group(2)}.{m.group(3)}"
|
||||||
|
else:
|
||||||
|
m = _DENSE_MLP.match(key)
|
||||||
|
if m:
|
||||||
|
new_key = f"layers.{m.group(1)}.mlp.{m.group(2)}.{m.group(3)}"
|
||||||
|
|
||||||
|
if new_key is None:
|
||||||
|
m = _ATTN.match(key)
|
||||||
|
if m:
|
||||||
|
new_key = (
|
||||||
|
f"layers.{m.group(1)}.attention.{m.group(2)}_proj.{m.group(3)}"
|
||||||
|
)
|
||||||
|
elif (m := _Q_NORM.match(key)) is not None:
|
||||||
|
new_key = f"layers.{m.group(1)}.attention.q_norm.weight"
|
||||||
|
elif (m := _K_NORM.match(key)) is not None:
|
||||||
|
new_key = f"layers.{m.group(1)}.attention.k_norm.weight"
|
||||||
|
elif (m := _INPUT_NORM.match(key)) is not None:
|
||||||
|
new_key = f"layers.{m.group(1)}.input_norm.weight"
|
||||||
|
elif (m := _POST_NORM.match(key)) is not None:
|
||||||
|
new_key = f"layers.{m.group(1)}.post_attention_norm.weight"
|
||||||
|
elif (m := _EMBED.match(key)) is not None:
|
||||||
|
new_key = "embed_tokens.weight"
|
||||||
|
elif (m := _FINAL_NORM.match(key)) is not None:
|
||||||
|
new_key = "norm.weight"
|
||||||
|
elif (m := _LM_HEAD.match(key)) is not None:
|
||||||
|
new_key = "lm_head.weight"
|
||||||
|
|
||||||
|
if new_key is None:
|
||||||
|
skipped.append(key)
|
||||||
|
else:
|
||||||
|
converted[new_key] = tensor
|
||||||
|
|
||||||
|
if skipped:
|
||||||
|
logger.warning(
|
||||||
|
"Dropped %d unmapped HuggingFace weight key(s): %s",
|
||||||
|
len(skipped),
|
||||||
|
", ".join(sorted(skipped)[:10]),
|
||||||
|
)
|
||||||
|
return converted
|
||||||
@@ -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()
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user