Compare commits
330
Commits
v1.3.5
...
925cbedc93
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
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 |
+3
-1
@@ -4,6 +4,8 @@
|
|||||||
# Allow necessary files
|
# Allow necessary files
|
||||||
!astrai/
|
!astrai/
|
||||||
!scripts/
|
!scripts/
|
||||||
!assets/
|
!docs/
|
||||||
|
!csrc/
|
||||||
|
!setup.py
|
||||||
!pyproject.toml
|
!pyproject.toml
|
||||||
!README.md
|
!README.md
|
||||||
|
|||||||
@@ -2,7 +2,7 @@
|
|||||||
name: Bug report
|
name: Bug report
|
||||||
about: Create a report to help us improve
|
about: Create a report to help us improve
|
||||||
title: "[BUG]"
|
title: "[BUG]"
|
||||||
labels: enhancement
|
labels: bug
|
||||||
assignees: ''
|
assignees: ''
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|||||||
@@ -16,9 +16,9 @@ Please delete options that are not relevant.
|
|||||||
Please describe the tests that you ran to verify your changes. Provide instructions so we can reproduce.
|
Please describe the tests that you ran to verify your changes. Provide instructions so we can reproduce.
|
||||||
|
|
||||||
## Checklist:
|
## Checklist:
|
||||||
- [ ] My code follows the style guidelines of this project (run `ruff format .` and `ruff check --fix .`)
|
- [ ] My code follows the style guidelines of this project (run `ruff format .` and `ruff check . --select I`)
|
||||||
- [ ] I have performed a self-review of my own code
|
- [ ] I have performed a self-review of my own code
|
||||||
- [ ] I have commented my code, particularly in hard-to-understand areas
|
- [ ] Code is self-documenting (no unnecessary comments)
|
||||||
- [ ] I have made corresponding changes to the documentation
|
- [ ] I have made corresponding changes to the documentation
|
||||||
- [ ] My changes generate no new warnings
|
- [ ] My changes generate no new warnings
|
||||||
- [ ] I have added tests that prove my fix is effective or that my feature works
|
- [ ] I have added tests that prove my fix is effective or that my feature works
|
||||||
|
|||||||
@@ -0,0 +1,100 @@
|
|||||||
|
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/
|
||||||
|
|
||||||
|
- 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
|
||||||
+17
-4
@@ -5,8 +5,16 @@
|
|||||||
!*/
|
!*/
|
||||||
|
|
||||||
# Allow specific file types and root files
|
# Allow specific file types and root files
|
||||||
!*.py
|
!astrai/**/*.py
|
||||||
!*.sh
|
!scripts/**/*.py
|
||||||
|
!tests/**/*.py
|
||||||
|
!csrc/**/*.py
|
||||||
|
|
||||||
|
!csrc/**/*.cu
|
||||||
|
!csrc/**/*.h
|
||||||
|
!csrc/**/*.cuh
|
||||||
|
|
||||||
|
!scripts/**/*.sh
|
||||||
|
|
||||||
# Allow GitHub files
|
# Allow GitHub files
|
||||||
!/.github/**
|
!/.github/**
|
||||||
@@ -16,8 +24,13 @@
|
|||||||
!/.dockerignore
|
!/.dockerignore
|
||||||
!/Dockerfile
|
!/Dockerfile
|
||||||
!/docker-compose.yml
|
!/docker-compose.yml
|
||||||
!/assets/**
|
!/docs/**
|
||||||
!/CONTRIBUTING.md
|
!/CONTRIBUTING.md
|
||||||
!/LICENSE
|
!/LICENSE
|
||||||
!/pyproject.toml
|
!/pyproject.toml
|
||||||
!/README.md
|
!/README.md
|
||||||
|
# Allow extension modules (only source .py)
|
||||||
|
!/astrai/extension/**/*.py
|
||||||
|
|
||||||
|
# Allow build files
|
||||||
|
!/setup.py
|
||||||
|
|||||||
+80
-48
@@ -1,68 +1,100 @@
|
|||||||
# Contributing to AstrAI
|
# Contributing to AstrAI
|
||||||
|
|
||||||
Thank you for your interest in contributing to AstrAI! This document provides guidelines and steps for contributing.
|
Thank you for your interest in contributing! This document provides step-by-step guidelines.
|
||||||
|
|
||||||
## How to Contribute
|
## Quick Start
|
||||||
|
|
||||||
### Reporting Issues
|
```bash
|
||||||
If you encounter a bug or have a feature request, please open an issue on GitHub. Include as much detail as possible:
|
git clone https://github.com/ViperEkura/AstrAI.git
|
||||||
- A clear description of the problem or request.
|
cd AstrAI
|
||||||
- Steps to reproduce (for bugs).
|
pip install -e ".[dev]" # install with dev dependencies (pytest, ruff)
|
||||||
- Your environment (Python version, OS, etc.).
|
```
|
||||||
|
|
||||||
### Submitting Changes
|
## Before You Commit
|
||||||
1. **Fork** the repository.
|
|
||||||
2. **Clone** your fork:
|
|
||||||
```bash
|
|
||||||
git clone https://github.com/your-username/AstrAI.git
|
|
||||||
cd AstrAI
|
|
||||||
```
|
|
||||||
3. **Create a feature branch**:
|
|
||||||
```bash
|
|
||||||
git checkout -b feature/your-feature-name
|
|
||||||
```
|
|
||||||
4. **Make your changes**. Follow the code style guidelines below.
|
|
||||||
5. **Commit your changes** with a descriptive commit message:
|
|
||||||
```bash
|
|
||||||
git commit -m "Add: brief description of the change"
|
|
||||||
```
|
|
||||||
6. **Push** to your fork:
|
|
||||||
```bash
|
|
||||||
git push origin feature/your-feature-name
|
|
||||||
```
|
|
||||||
7. **Open a Pull Request** (PR) against the `main` branch of the upstream repository.
|
|
||||||
|
|
||||||
## Code Style
|
Run the following checks **in order** — CI will reject if any fail.
|
||||||
|
|
||||||
AstrAI uses [Ruff](https://docs.astral.sh/ruff/) for code formatting and linting. Please ensure your code is formatted before submitting.
|
### 1. Format
|
||||||
|
|
||||||
- Run Ruff to format and lint (requires conda environment `nlp`):
|
```bash
|
||||||
```bash
|
ruff format .
|
||||||
conda run -n nlp ruff format .
|
```
|
||||||
conda run -n nlp ruff check --fix .
|
|
||||||
```
|
|
||||||
- The project uses **double quotes** for strings and **4‑space indentation** (as configured in `pyproject.toml`).
|
|
||||||
|
|
||||||
## Testing
|
> **Note**: `ruff format` may rename parameters (e.g. `mask` → `attn_mask`).
|
||||||
|
> Always review the diff after formatting.
|
||||||
|
|
||||||
If you add or modify functionality, please include appropriate tests.
|
### 2. Import sorting
|
||||||
|
|
||||||
- Run the test suite with:
|
```bash
|
||||||
```bash
|
ruff check . --select I
|
||||||
conda run -n nlp python -u -m pytest
|
```
|
||||||
```
|
|
||||||
- Ensure all tests pass before submitting your PR.
|
If this fails, **manually fix** import ordering (ruff does not auto-fix in this project's CI):
|
||||||
|
|
||||||
|
```bash
|
||||||
|
ruff check . --select I --fix .
|
||||||
|
ruff format . # re-format after fix
|
||||||
|
```
|
||||||
|
|
||||||
|
### 3. Run tests
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python -u -m pytest tests/ -v
|
||||||
|
```
|
||||||
|
|
||||||
|
> Failed tests may leave orphan tempdirs under `%TEMP%`. Clean them manually if needed.
|
||||||
|
|
||||||
|
### 4. (Optional) Full pre-commit check
|
||||||
|
|
||||||
|
If you have Git Bash available:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
bash scripts/pre_commit.sh
|
||||||
|
```
|
||||||
|
|
||||||
|
This runs format check, import sort check, and tests in one go.
|
||||||
|
|
||||||
|
## Commit Style
|
||||||
|
|
||||||
|
```
|
||||||
|
fix/feat/chore/docs/refactor/perf/test/style/ci/build/revert : 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 hook rejects | Tests or lint failed | Fix individually, do not `--no-verify` |
|
||||||
|
| 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
|
## Code Review
|
||||||
|
|
||||||
All submissions will be reviewed. We may request changes or discuss alternatives. Please be responsive to feedback.
|
- All PRs are reviewed. We may request changes.
|
||||||
|
- CI runs `ruff format --check .` then `ruff check . --select I` (no `--fix` in CI).
|
||||||
|
- Ensure all tests pass.
|
||||||
|
|
||||||
## License
|
## License
|
||||||
|
|
||||||
By contributing, you agree that your contributions will be licensed under the same [GPL-3.0 License](LICENSE) that covers the project.
|
By contributing, you agree that your contributions will be licensed under the [GPL-3.0 License](LICENSE).
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
If you have any questions, feel free to ask in the [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions) or open an issue.
|
Questions? Ask in [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions) or open an issue.
|
||||||
|
|
||||||
Happy contributing!
|
|
||||||
|
|||||||
+17
-6
@@ -1,7 +1,15 @@
|
|||||||
# AstrAI Dockerfile - Multi-stage Build (Optimized)
|
# 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
|
# Build stage - use base image with minimal build tools
|
||||||
FROM nvidia/cuda:12.6.0-base-ubuntu24.04 AS builder
|
FROM ubuntu:24.04 AS builder
|
||||||
|
|
||||||
|
ARG CUDA_TAG=cu128
|
||||||
|
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
@@ -18,21 +26,24 @@ RUN apt-get update && DEBIAN_FRONTEND=noninteractive apt-get install -y --no-ins
|
|||||||
RUN python3.12 -m venv --copies /opt/venv
|
RUN python3.12 -m venv --copies /opt/venv
|
||||||
ENV PATH="/opt/venv/bin:$PATH"
|
ENV PATH="/opt/venv/bin:$PATH"
|
||||||
|
|
||||||
# Copy source code and install dependencies
|
# Copy source code and install (deps read from pyproject.toml)
|
||||||
COPY astrai/ ./astrai/
|
COPY astrai/ ./astrai/
|
||||||
|
COPY csrc/ ./csrc/
|
||||||
|
COPY setup.py .
|
||||||
COPY pyproject.toml .
|
COPY pyproject.toml .
|
||||||
RUN pip install --no-cache-dir --upgrade pip \
|
RUN pip install --no-cache-dir --upgrade pip \
|
||||||
&& pip install --no-cache-dir . \
|
&& pip install --no-cache-dir . \
|
||||||
--extra-index-url https://download.pytorch.org/whl/cu126
|
--extra-index-url "https://download.pytorch.org/whl/${CUDA_TAG}"
|
||||||
|
|
||||||
# Production stage
|
# Production stage
|
||||||
FROM nvidia/cuda:12.6.0-base-ubuntu24.04 AS production
|
FROM ubuntu:24.04 AS production
|
||||||
|
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
# Install Python 3.12 runtime
|
# Install Python 3.12 runtime and healthcheck dependency
|
||||||
RUN apt-get update && DEBIAN_FRONTEND=noninteractive apt-get install -y --no-install-recommends \
|
RUN apt-get update && DEBIAN_FRONTEND=noninteractive apt-get install -y --no-install-recommends \
|
||||||
python3.12 \
|
python3.12 \
|
||||||
|
curl \
|
||||||
&& rm -rf /var/lib/apt/lists/*
|
&& rm -rf /var/lib/apt/lists/*
|
||||||
|
|
||||||
# Copy virtual environment from builder
|
# Copy virtual environment from builder
|
||||||
@@ -42,7 +53,7 @@ ENV PATH="/opt/venv/bin:$PATH"
|
|||||||
# Copy application code
|
# Copy application code
|
||||||
COPY astrai/ ./astrai/
|
COPY astrai/ ./astrai/
|
||||||
COPY scripts/ ./scripts/
|
COPY scripts/ ./scripts/
|
||||||
COPY assets/ ./assets/
|
COPY docs/ ./docs/
|
||||||
COPY pyproject.toml .
|
COPY pyproject.toml .
|
||||||
COPY README.md .
|
COPY README.md .
|
||||||
|
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
<div align="center">
|
<div align="center">
|
||||||
|
|
||||||
<img src="assets/images/logo.png" width="auto" alt="Logo">
|
<img src="docs/images/logo.png" width="auto" alt="Logo">
|
||||||
<p>
|
<p>
|
||||||
<strong>A lightweight Transformer training & inference framework</strong>
|
<strong>A lightweight Transformer training & inference framework</strong>
|
||||||
</p>
|
</p>
|
||||||
@@ -9,18 +9,18 @@
|
|||||||
<div align="center">
|
<div align="center">
|
||||||
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
||||||
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license">
|
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license">
|
||||||
<img src="https://img.shields.io/github/v/release/ViperEkura/AstrAI?color=76bad9" alt="release">
|
<img src="https://img.shields.io/github/v/tag/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release">
|
||||||
<img src="https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fapi.github.com%2Frepos%2FViperEkura%2FAstrAI&query=%24.stargazers_count&label=stars&suffix=%20stars&color=76bad9" alt="stars">
|
<img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars">
|
||||||
<img src="https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fapi.github.com%2Frepos%2FViperEkura%2FAstrAI&query=%24.forks_count&label=forks&suffix=%20forks&color=76bad9" alt="forks">
|
<img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks">
|
||||||
</div>
|
</div>
|
||||||
<br>
|
<br>
|
||||||
|
|
||||||
<div align="center">
|
<div align="center">
|
||||||
<a href="#english">English</a> •
|
<a href="#english">English</a> •
|
||||||
<a href="assets/docs/README-zh-CN.md">中文</a> •
|
<a href="docs/README-zh-CN.md">中文</a> •
|
||||||
<a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> •
|
<a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> •
|
||||||
<a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> •
|
<a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> •
|
||||||
<a href="https://huggingface.co/ViperEk/">HuggingFace</a>
|
<a href="https://huggingface.co/ViperEkura">HuggingFace</a>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<br>
|
<br>
|
||||||
@@ -28,7 +28,8 @@
|
|||||||
## 📖 Table of Contents
|
## 📖 Table of Contents
|
||||||
|
|
||||||
- [Features](#features)
|
- [Features](#features)
|
||||||
- [Quick Start](#quick-start)
|
- [Getting Started](#getting-started)
|
||||||
|
- [Demo](#demo)
|
||||||
- [Documentation](#documentation)
|
- [Documentation](#documentation)
|
||||||
- [Contributing](#contributing)
|
- [Contributing](#contributing)
|
||||||
- [Community](#community)
|
- [Community](#community)
|
||||||
@@ -49,55 +50,116 @@
|
|||||||
- 🤗 **HuggingFace-Style API**: AutoModel/AutoTokenizer APIs inspired by HuggingFace for easy model and tokenizer loading.
|
- 🤗 **HuggingFace-Style API**: AutoModel/AutoTokenizer APIs inspired by HuggingFace for easy model and tokenizer loading.
|
||||||
- 🔌 **Dual API Compatibility**: Supports both OpenAI and Anthropic chat completion APIs out of the box.
|
- 🔌 **Dual API Compatibility**: Supports both OpenAI and Anthropic chat completion APIs out of the box.
|
||||||
|
|
||||||
### Quick Start
|
### Getting Started
|
||||||
|
|
||||||
#### Installation
|
End-to-end walkthrough in 5 steps:
|
||||||
|
|
||||||
|
**1. Install**
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
git clone https://github.com/ViperEkura/AstrAI.git
|
git clone https://github.com/ViperEkura/AstrAI.git
|
||||||
cd AstrAI
|
cd AstrAI
|
||||||
pip install -e .
|
pip install -e . # pure PyTorch (no CUDA kernels)
|
||||||
|
# CSRC_KERNELS=true pip install -e . --no-build-isolation # optional: fused CUDA kernels
|
||||||
|
# pip install -e ".[dev]" # dev dependencies (pytest, ruff)
|
||||||
```
|
```
|
||||||
|
|
||||||
For development dependencies:
|
**2. Download model**
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
pip install -e ".[dev]"
|
python scripts/demo/download.py # downloads 1B checkpoint to params/
|
||||||
```
|
```
|
||||||
|
|
||||||
#### Download Pre-trained Model
|
**3. Preprocess data**
|
||||||
|
|
||||||
Download pre-trained model weights (1B bilingual checkpoint) to `params/`:
|
Create `pretrain.json` (preprocessing config for `seq` strategy):
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"version": 1,
|
||||||
|
"input": {"sections": [{"field": "text", "action": "train"}]},
|
||||||
|
"preprocessing": {"max_seq_len": 2048},
|
||||||
|
"output": {"storage_format": "bin"}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python scripts/demo/download.py
|
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c pretrain.json
|
||||||
```
|
```
|
||||||
|
|
||||||
Or download manually from [HuggingFace](https://huggingface.co/ViperEk/KHAOSZ) into `params/`.
|
**4. Train**
|
||||||
|
|
||||||
#### Train a Model
|
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
CUDA_VISIBLE_DEVICES=0,1,2,3 python scripts/tools/train.py \
|
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||||
--train_type seq \
|
|
||||||
--data_root_path /path/to/dataset \
|
nohup python scripts/tools/train.py \
|
||||||
--param_path /path/to/model \
|
--nprocs=4 \
|
||||||
--batch_size 4 \
|
--parallel_mode=ddp \
|
||||||
--accumulation_steps 8 \
|
--train_type=seq \
|
||||||
--max_lr 3e-4 \
|
--data_root_path=/path/to/dataset \
|
||||||
--warmup_steps 1000 \
|
--param_path=/path/to/model \
|
||||||
--n_epoch 1
|
--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 &
|
||||||
```
|
```
|
||||||
|
|
||||||
Full reference at [Parameter Guide](assets/docs/params.md).
|
**5. Serve & query**
|
||||||
|
|
||||||
#### Generate Text
|
```bash
|
||||||
|
# Terminal 1: start server
|
||||||
|
python scripts/tools/server.py --param_path ./params --device cuda
|
||||||
|
|
||||||
|
# Terminal 2: query
|
||||||
|
curl http://localhost:8000/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{"messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||||
|
```
|
||||||
|
|
||||||
|
### Demo
|
||||||
|
|
||||||
|
Check out the demos in the `scripts/demo/` folder:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Download model weights (required before running demos)
|
||||||
|
python scripts/demo/download.py # model → params/
|
||||||
|
|
||||||
|
# Interactive streaming chat (multi-turn, maintains history)
|
||||||
|
python scripts/demo/stream_chat.py
|
||||||
|
# Type your message after >>, type !exit to quit
|
||||||
|
|
||||||
|
# Batch generation (5 hardcoded prompts, non-streaming)
|
||||||
|
python scripts/demo/generate_batch.py
|
||||||
|
|
||||||
|
# Single-prompt autoregressive streaming
|
||||||
|
python scripts/demo/generate_ar.py
|
||||||
|
```
|
||||||
|
|
||||||
|
All generation demos use `temperature=0.8`, `top_p=0.95`, `top_k=50`, `max_tokens=2048` by default and require `params/` to contain model weights (run `download.py` first).
|
||||||
|
|
||||||
|
Watch a video walkthrough on [bilibili](https://www.bilibili.com/video/BV1fuLB6yEj6).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
See [Documentation](#documentation) for full references beyond the examples above.
|
||||||
|
|
||||||
|
#### Text Generation
|
||||||
|
|
||||||
|
Batch generation from a JSONL file:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python scripts/tools/generate.py \
|
python scripts/tools/generate.py \
|
||||||
--param_path /path/to/model \
|
--param_path ./params \
|
||||||
--input_json_file /path/to/input.json \
|
--input_json_file input.jsonl \
|
||||||
--output_json_file /path/to/output.json
|
--output_json_file output.jsonl
|
||||||
```
|
```
|
||||||
|
|
||||||
#### Docker
|
#### Docker
|
||||||
@@ -111,9 +173,6 @@ docker build -t astrai:latest .
|
|||||||
# Run with GPU support
|
# Run with GPU support
|
||||||
docker run --gpus all -it astrai:latest
|
docker run --gpus all -it astrai:latest
|
||||||
|
|
||||||
# Run with specific GPUs
|
|
||||||
docker run --gpus '"device=0,1"' -it astrai:latest
|
|
||||||
|
|
||||||
# Run inference server
|
# Run inference server
|
||||||
docker run --gpus all -p 8000:8000 astrai:latest \
|
docker run --gpus all -p 8000:8000 astrai:latest \
|
||||||
python -m scripts.tools.server --port 8000 --device cuda
|
python -m scripts.tools.server --port 8000 --device cuda
|
||||||
@@ -130,87 +189,47 @@ docker compose --profile cpu up -d
|
|||||||
|
|
||||||
> **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`.
|
> **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`.
|
||||||
|
|
||||||
#### Start HTTP Server
|
#### HTTP API Examples
|
||||||
|
|
||||||
Start the inference server with OpenAI and Anthropic-compatible HTTP API:
|
Additional request examples beyond the [Getting Started](#getting-started) flow:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python -m scripts.tools.server --port 8000 --device cuda
|
|
||||||
```
|
|
||||||
|
|
||||||
Make requests:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
# OpenAI-compatible
|
|
||||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
|
||||||
-H "Content-Type: application/json" \
|
|
||||||
-d '{
|
|
||||||
"messages": [{"role": "user", "content": "Hello"}],
|
|
||||||
"max_tokens": 512
|
|
||||||
}'
|
|
||||||
|
|
||||||
# OpenAI-compatible streaming
|
# OpenAI-compatible streaming
|
||||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||||
-H "Content-Type: application/json" \
|
-H "Content-Type: application/json" \
|
||||||
-d '{
|
-d '{"messages":[{"role":"user","content":"Tell a story"}],"stream":true,"max_tokens":500}'
|
||||||
"messages": [{"role": "user", "content": "Tell a story"}],
|
|
||||||
"stream": true,
|
|
||||||
"max_tokens": 500
|
|
||||||
}'
|
|
||||||
|
|
||||||
# Anthropic-compatible
|
# Anthropic-compatible
|
||||||
curl -X POST http://localhost:8000/v1/messages \
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
-H "Content-Type: application/json" \
|
-H "Content-Type: application/json" \
|
||||||
-d '{
|
-d '{"model":"astrai","system":"You are a helpful assistant.","messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||||
"model": "astrai",
|
|
||||||
"system": "You are a helpful assistant.",
|
|
||||||
"messages": [{"role": "user", "content": "Hello"}],
|
|
||||||
"max_tokens": 512
|
|
||||||
}'
|
|
||||||
|
|
||||||
# Anthropic-compatible streaming with stop sequences
|
# Anthropic-compatible streaming with stop sequences
|
||||||
curl -X POST http://localhost:8000/v1/messages \
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
-H "Content-Type: application/json" \
|
-H "Content-Type: application/json" \
|
||||||
-d '{
|
-d '{"model":"astrai","messages":[{"role":"user","content":"Write a story"}],"max_tokens":500,"stream":true,"stop_sequences":["The end"]}'
|
||||||
"model": "astrai",
|
|
||||||
"messages": [{"role": "user", "content": "Write a story"}],
|
|
||||||
"max_tokens": 500,
|
|
||||||
"stream": true,
|
|
||||||
"stop_sequences": ["The end"]
|
|
||||||
}'
|
|
||||||
|
|
||||||
# Health check
|
# Health check
|
||||||
curl http://localhost:8000/health
|
curl http://localhost:8000/health
|
||||||
```
|
```
|
||||||
|
|
||||||
#### Demo
|
See [Inference Guide](docs/guides/inference.md) for SSE streaming format, error codes, and stats endpoint.
|
||||||
|
|
||||||
Check out the demos in the `scripts/demo/` folder:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
# Download pre‑processed data (required before running demos)
|
|
||||||
python scripts/demo/download.py
|
|
||||||
|
|
||||||
# Interactive streaming chat
|
|
||||||
python scripts/demo/stream_chat.py
|
|
||||||
|
|
||||||
# Batch generation
|
|
||||||
python scripts/demo/generate_batch.py
|
|
||||||
|
|
||||||
# Auto‑regressive generation
|
|
||||||
python scripts/demo/generate_ar.py
|
|
||||||
```
|
|
||||||
|
|
||||||
Watch a video walkthrough on [bilibili](https://www.bilibili.com/video/BV1z5RPYHEkd).
|
|
||||||
|
|
||||||
### Documentation
|
### Documentation
|
||||||
|
|
||||||
| Document | Description |
|
| Document | Description |
|
||||||
|----------|-------------|
|
|----------|-------------|
|
||||||
| [Parameter Guide](./assets/docs/params.md) | Training & inference parameters |
|
| [Get Started](./docs/get-started.md) | Installation and quickstart |
|
||||||
| [Design Document](./assets/docs/design.md) | Framework architecture & module design |
|
| [CLI Reference](./docs/guides/params.md) | Parameters for all CLI tools (train, server, generate, preprocess) |
|
||||||
| [Data Flow](./assets/docs/dataflow.md) | Data processing pipeline details |
|
| [Preprocessing](./docs/guides/preprocessing.md) | Declarative JSON-driven data preprocessing |
|
||||||
| [Model Introduction](./assets/docs/introduction.md) | Model architecture & technical details |
|
| [Training](./docs/guides/training.md) | Training loop, strategies & formulas |
|
||||||
|
| [Inference](./docs/guides/inference.md) | KVCache, continuous batching, sampling & HTTP API |
|
||||||
|
| [Evaluation](./docs/guides/evaluation.md) | HumanEval, MMLU, PPL, ROUGE, IFD, IFEval |
|
||||||
|
| [Distributed](./docs/guides/distributed.md) | Multi-GPU DDP / FSDP training |
|
||||||
|
| [Architecture](./docs/developer/architecture.md) | System architecture, class diagram & design patterns |
|
||||||
|
| [Data Flow](./docs/developer/dataflow.md) | Data pipeline, storage backends & dataset architecture |
|
||||||
|
| [Internals](./docs/developer/internals.md) | Training internals: loss formulas, callback lifecycle, KV cache |
|
||||||
|
| [CUDA Kernels](./docs/developer/cuda_kernels.md) | Custom CUDA attention kernels & benchmarks |
|
||||||
|
|
||||||
### Contributing
|
### Contributing
|
||||||
|
|
||||||
@@ -227,7 +246,7 @@ For major changes, please open an issue first to discuss what you would like to
|
|||||||
|
|
||||||
- **GitHub Issues**: [Issue Tracker](https://github.com/ViperEkura/AstrAI/issues)
|
- **GitHub Issues**: [Issue Tracker](https://github.com/ViperEkura/AstrAI/issues)
|
||||||
- **Discussions**: [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions)
|
- **Discussions**: [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions)
|
||||||
- **HuggingFace**: [Model Hub](https://huggingface.co/ViperEk)
|
- **HuggingFace**: [Model Hub](https://huggingface.co/ViperEkura)
|
||||||
|
|
||||||
### License
|
### License
|
||||||
|
|
||||||
|
|||||||
@@ -1,246 +0,0 @@
|
|||||||
<div align="center">
|
|
||||||
|
|
||||||
<img src="../images/logo.png" width="auto" alt="Logo">
|
|
||||||
|
|
||||||
<div>
|
|
||||||
<a href="../../README.md">English</a> •
|
|
||||||
<a href="#chinese">中文</a>
|
|
||||||
</div>
|
|
||||||
|
|
||||||
<p>
|
|
||||||
<strong>轻量级 Transformer 训练与推理框架</strong>
|
|
||||||
</p>
|
|
||||||
</div>
|
|
||||||
|
|
||||||
<div align="center">
|
|
||||||
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
|
||||||
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license">
|
|
||||||
<img src="https://img.shields.io/github/v/release/ViperEkura/AstrAI?color=76bad9" alt="release">
|
|
||||||
<img src="https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fapi.github.com%2Frepos%2FViperEkura%2FAstrAI&query=%24.stargazers_count&label=stars&suffix=%20stars&color=76bad9" alt="stars">
|
|
||||||
<img src="https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fapi.github.com%2Frepos%2FViperEkura%2FAstrAI&query=%24.forks_count&label=forks&suffix=%20forks&color=76bad9" alt="forks">
|
|
||||||
</div>
|
|
||||||
|
|
||||||
<br>
|
|
||||||
|
|
||||||
<div align="center">
|
|
||||||
<a href="../../README.md">English</a> •
|
|
||||||
<a href="#chinese">中文</a> •
|
|
||||||
<a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> •
|
|
||||||
<a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> •
|
|
||||||
<a href="https://huggingface.co/ViperEk">HuggingFace</a>
|
|
||||||
</div>
|
|
||||||
<br>
|
|
||||||
|
|
||||||
## 📖 目录
|
|
||||||
|
|
||||||
- [特性](#特性)
|
|
||||||
- [快速开始](#快速开始)
|
|
||||||
- [文档](#文档)
|
|
||||||
- [贡献](#贡献)
|
|
||||||
- [社区](#社区)
|
|
||||||
- [许可证](#许可证)
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
<a id="chinese"></a>
|
|
||||||
## 中文
|
|
||||||
|
|
||||||
### 特性
|
|
||||||
|
|
||||||
- 🚀 **高性能**: 训练与推理双向优化,高效并行。
|
|
||||||
- 🔧 **灵活**: 支持 seq/sft/dpo/grpo 多种训练方式,可定制模型架构。
|
|
||||||
- 💡 **易用**: 简洁的 API 与丰富的示例、演示。
|
|
||||||
- 📦 **轻量**: 依赖少,部署简单。
|
|
||||||
- 🔬 **研究友好**: 模块化设计,便于实验新想法。
|
|
||||||
- 🤗 **HuggingFace 风格 API**: 类 HuggingFace 的 AutoModel/AutoTokenizer 接口,方便加载模型和分词器。
|
|
||||||
- 🔌 **双 API 兼容**: 同时支持 OpenAI 和 Anthropic 聊天补全 API,开箱即用。
|
|
||||||
|
|
||||||
### 快速开始
|
|
||||||
|
|
||||||
#### 安装
|
|
||||||
|
|
||||||
```bash
|
|
||||||
git clone https://github.com/ViperEkura/AstrAI.git
|
|
||||||
cd AstrAI
|
|
||||||
pip install -e .
|
|
||||||
```
|
|
||||||
|
|
||||||
安装开发依赖:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
pip install -e ".[dev]"
|
|
||||||
```
|
|
||||||
|
|
||||||
#### 下载预训练模型
|
|
||||||
|
|
||||||
下载预训练模型权重(1B 双语检查点)到 `params/` 目录:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python scripts/demo/download.py
|
|
||||||
```
|
|
||||||
|
|
||||||
或从 [HuggingFace](https://huggingface.co/ViperEk/KHAOSZ) 手动下载放入 `params/`。
|
|
||||||
|
|
||||||
#### 训练模型
|
|
||||||
|
|
||||||
```bash
|
|
||||||
CUDA_VISIBLE_DEVICES=0,1,2,3 python scripts/tools/train.py \
|
|
||||||
--train_type seq \
|
|
||||||
--data_root_path /path/to/dataset \
|
|
||||||
--param_path /path/to/model \
|
|
||||||
--batch_size 4 \
|
|
||||||
--accumulation_steps 8 \
|
|
||||||
--max_lr 3e-4 \
|
|
||||||
--warmup_steps 1000 \
|
|
||||||
--n_epoch 1
|
|
||||||
```
|
|
||||||
|
|
||||||
完整参数列表见[参数说明](./params.md)。
|
|
||||||
|
|
||||||
#### 文本生成
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python scripts/tools/generate.py \
|
|
||||||
--param_path /path/to/model \
|
|
||||||
--input_json_file /path/to/input.json \
|
|
||||||
--output_json_file /path/to/output.json
|
|
||||||
```
|
|
||||||
|
|
||||||
#### Docker
|
|
||||||
|
|
||||||
使用 Docker 构建和运行(推荐用于 GPU 环境):
|
|
||||||
|
|
||||||
```bash
|
|
||||||
# 构建镜像
|
|
||||||
docker build -t astrai:latest .
|
|
||||||
|
|
||||||
# 启用 GPU 运行
|
|
||||||
docker run --gpus all -it astrai:latest
|
|
||||||
|
|
||||||
# 指定特定 GPU
|
|
||||||
docker run --gpus '"device=0,1"' -it astrai:latest
|
|
||||||
|
|
||||||
# 运行推理服务
|
|
||||||
docker run --gpus all -p 8000:8000 astrai:latest \
|
|
||||||
python -m scripts.tools.server --port 8000 --device cuda
|
|
||||||
|
|
||||||
# 挂载数据卷
|
|
||||||
docker run --gpus all -v /path/to/data:/data -it astrai:latest
|
|
||||||
|
|
||||||
# Docker Compose(GPU,默认)
|
|
||||||
docker compose up -d
|
|
||||||
|
|
||||||
# Docker Compose(仅 CPU)
|
|
||||||
docker compose --profile cpu up -d
|
|
||||||
```
|
|
||||||
|
|
||||||
> **注意**: 必须使用 `--gpus all` 才能启用 CUDA 支持,否则 `torch.cuda.is_available()` 将返回 `False`。
|
|
||||||
|
|
||||||
#### 启动 HTTP 服务
|
|
||||||
|
|
||||||
启动推理服务器,支持 OpenAI 和 Anthropic 兼容的 HTTP API:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python -m scripts.tools.server --port 8000 --device cuda
|
|
||||||
```
|
|
||||||
|
|
||||||
发起请求:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
# OpenAI 兼容
|
|
||||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
|
||||||
-H "Content-Type: application/json" \
|
|
||||||
-d '{
|
|
||||||
"messages": [{"role": "user", "content": "你好"}],
|
|
||||||
"max_tokens": 512
|
|
||||||
}'
|
|
||||||
|
|
||||||
# OpenAI 兼容流式
|
|
||||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
|
||||||
-H "Content-Type: application/json" \
|
|
||||||
-d '{
|
|
||||||
"messages": [{"role": "user", "content": "讲个故事"}],
|
|
||||||
"stream": true,
|
|
||||||
"max_tokens": 500
|
|
||||||
}'
|
|
||||||
|
|
||||||
# Anthropic 兼容
|
|
||||||
curl -X POST http://localhost:8000/v1/messages \
|
|
||||||
-H "Content-Type: application/json" \
|
|
||||||
-d '{
|
|
||||||
"model": "astrai",
|
|
||||||
"system": "你是一个乐于助人的助手。",
|
|
||||||
"messages": [{"role": "user", "content": "你好"}],
|
|
||||||
"max_tokens": 512
|
|
||||||
}'
|
|
||||||
|
|
||||||
# Anthropic 兼容流式并设置停止序列
|
|
||||||
curl -X POST http://localhost:8000/v1/messages \
|
|
||||||
-H "Content-Type: application/json" \
|
|
||||||
-d '{
|
|
||||||
"model": "astrai",
|
|
||||||
"messages": [{"role": "user", "content": "写个故事"}],
|
|
||||||
"max_tokens": 500,
|
|
||||||
"stream": true,
|
|
||||||
"stop_sequences": ["结束"]
|
|
||||||
}'
|
|
||||||
|
|
||||||
# 健康检查
|
|
||||||
curl http://localhost:8000/health
|
|
||||||
```
|
|
||||||
|
|
||||||
#### 演示
|
|
||||||
|
|
||||||
查看 `scripts/demo/` 文件夹中的演示:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
# 下载预处理数据(运行演示前必需)
|
|
||||||
python scripts/demo/download.py
|
|
||||||
|
|
||||||
# 交互式流式聊天
|
|
||||||
python scripts/demo/stream_chat.py
|
|
||||||
|
|
||||||
# 批量生成
|
|
||||||
python scripts/demo/generate_batch.py
|
|
||||||
|
|
||||||
# 自回归生成
|
|
||||||
python scripts/demo/generate_ar.py
|
|
||||||
```
|
|
||||||
|
|
||||||
观看 [bilibili](https://www.bilibili.com/video/BV1z5RPYHEkd) 上的视频演示。
|
|
||||||
|
|
||||||
### 文档
|
|
||||||
|
|
||||||
| 文档 | 说明 |
|
|
||||||
|------|------|
|
|
||||||
| [参数说明](./params.md) | 训练与推理参数配置 |
|
|
||||||
| [设计文档](./design.md) | 系统架构与模块设计 |
|
|
||||||
| [数据流程](./dataflow.md) | 数据处理管道详解 |
|
|
||||||
| [模型介绍](./introduction.md) | 模型架构与技术细节 |
|
|
||||||
|
|
||||||
### 贡献
|
|
||||||
|
|
||||||
我们欢迎贡献!请参阅[贡献指南](../../CONTRIBUTING.md)了解详情。
|
|
||||||
|
|
||||||
1. Fork 本仓库。
|
|
||||||
2. 创建功能分支。
|
|
||||||
3. 提交更改。
|
|
||||||
4. 发起 Pull Request。
|
|
||||||
|
|
||||||
重大更改请先开 issue 讨论。
|
|
||||||
|
|
||||||
### 社区
|
|
||||||
|
|
||||||
- **GitHub Issues**: [问题追踪](https://github.com/ViperEkura/AstrAI/issues)
|
|
||||||
- **Discussions**: [GitHub 讨论区](https://github.com/ViperEkura/AstrAI/discussions)
|
|
||||||
- **HuggingFace**: [模型中心](https://huggingface.co/ViperEk)
|
|
||||||
|
|
||||||
### 许可证
|
|
||||||
|
|
||||||
本项目采用 [GPL-3.0 许可证](../../LICENSE)。
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
<div align="center">
|
|
||||||
<em>专为高性能与易用性设计的轻量级 Transformer 框架。</em>
|
|
||||||
</div>
|
|
||||||
@@ -1,237 +0,0 @@
|
|||||||
# AstrAI Data Flow Documentation
|
|
||||||
|
|
||||||
This document describes the data flow of the AstrAI project (a training and inference framework for autoregressive Transformer language models). It covers the complete flow from raw data to model training and inference.
|
|
||||||
|
|
||||||
## Overview
|
|
||||||
|
|
||||||
AstrAI adopts a modular design with the following main components:
|
|
||||||
- **Dataset Module** (`astrai/dataset/`): Dataset, sampler, serialization tools
|
|
||||||
- **Model Module** (`astrai/model/`): AutoModel, Transformer model and its submodules
|
|
||||||
- **Training Module** (`astrai/trainer/`): Trainer, training context, strategies, schedulers, callbacks, metric utilities
|
|
||||||
- **Inference Module** (`astrai/inference/`): Inference engine with continuous batching, streaming generation
|
|
||||||
- **Config Module** (`astrai/config/`): ModelConfig, TrainConfig
|
|
||||||
- **Factory Module** (`astrai/factory/`): Registry, BaseFactory for component registration
|
|
||||||
- **Parallel Module** (`astrai/parallel/`): Distributed training support
|
|
||||||
- **Serialization** (`astrai/serialization.py`): Checkpoint management with safetensors
|
|
||||||
|
|
||||||
## Data Flow Diagram
|
|
||||||
|
|
||||||
```mermaid
|
|
||||||
flowchart LR
|
|
||||||
subgraph A[Data Preparation]
|
|
||||||
direction TB
|
|
||||||
A1[Raw Text] --> A2[AutoTokenizer]
|
|
||||||
A2 --> A3[Tokenized .h5 files]
|
|
||||||
A3 --> A4[BaseDataset]
|
|
||||||
A4 --> A5[ResumableDistributedSampler]
|
|
||||||
A5 --> A6[DataLoader]
|
|
||||||
end
|
|
||||||
|
|
||||||
subgraph B[Training]
|
|
||||||
direction TB
|
|
||||||
B1[DataLoader] --> B2[BaseStrategy]
|
|
||||||
B2 --> B3[Transformer Forward]
|
|
||||||
B3 --> B4[Loss + Backward]
|
|
||||||
B4 --> B5[Gradient Accumulation]
|
|
||||||
B5 -->|every accum_steps| B6[Optimizer Step]
|
|
||||||
B6 --> B7[LR Scheduler]
|
|
||||||
B7 -->|next batch| B2
|
|
||||||
B6 --> B8[CheckpointCallback]
|
|
||||||
end
|
|
||||||
|
|
||||||
subgraph C[Inference]
|
|
||||||
direction TB
|
|
||||||
C1[Checkpoint] --> C2[AutoModel]
|
|
||||||
C1 --> C3[AutoTokenizer]
|
|
||||||
C2 --> C4[InferenceEngine]
|
|
||||||
C3 --> C4
|
|
||||||
C4 --> C5[InferenceScheduler]
|
|
||||||
C5 --> C6[Transformer Forward]
|
|
||||||
C6 --> C7[sample]
|
|
||||||
C7 --> C8{End?}
|
|
||||||
C8 -->|No| C6
|
|
||||||
C8 -->|Yes| C9[Generated Text]
|
|
||||||
end
|
|
||||||
|
|
||||||
A --> B
|
|
||||||
B --> C
|
|
||||||
```
|
|
||||||
|
|
||||||
## Detailed Module Descriptions
|
|
||||||
|
|
||||||
### 1. Data Serialization (`astrai/dataset/storage.py` & `astrai/serialization.py`)
|
|
||||||
|
|
||||||
- **`save_h5`**: Saves tensors by groups as HDF5 files (`.h5`), each key maps to a list of tensors
|
|
||||||
- **`load_h5`**: Loads `.h5` files, returns `Dict[str, List[Tensor]]`, supports shared memory
|
|
||||||
- **`Checkpoint`**: Encapsulates model state dict + epoch + iteration; uses safetensors
|
|
||||||
|
|
||||||
### 2. Dataset Module
|
|
||||||
|
|
||||||
#### 2.1 Dataset (`dataset.py`)
|
|
||||||
- **`BaseDataset`**: Abstract base class for windowed sequence sampling
|
|
||||||
- **`BaseSegmentFetcher` / `MultiSegmentFetcher`**: Fetch tensor segments by index range
|
|
||||||
- **`DatasetFactory`**: Creates dataset instances by `train_type` (`seq`, `sft`, `dpo`, `grpo`)
|
|
||||||
- Data keys: `"sequence"` (SEQ), `"loss_mask"` (SFT), `"chosen_mask"/"rejected_mask"` (DPO), `"masks"` (GRPO)
|
|
||||||
|
|
||||||
#### 2.2 Sampler (`sampler.py`)
|
|
||||||
- **`ResumableDistributedSampler`**: Tracks `epoch` and `iter` for breakpoint resume; supports shuffle and drop_last
|
|
||||||
|
|
||||||
### 3. Model Module
|
|
||||||
|
|
||||||
#### 3.1 Transformer / AutoModel
|
|
||||||
- **`AutoModel`**: Base class with `from_pretrained()` / `save_pretrained()`
|
|
||||||
- **`Transformer`**: Decoder-only architecture, registered via `@AutoModel.register('transformer')`
|
|
||||||
- Embedding → N×DecoderBlock → RMSNorm → Linear lm_head
|
|
||||||
- RoPE position encoding, optional weight tying
|
|
||||||
|
|
||||||
#### 3.2 Submodules (`module.py`)
|
|
||||||
- **`DecoderBlock`**: GQA attention + residual + MLP + RMSNorm
|
|
||||||
- **`GQA`**: Grouped Query Attention (also `MLA` for multi-latent attention)
|
|
||||||
- **`MLP`**: `SiLU(gate(x)) * up(x)` → down projection
|
|
||||||
- **`RotaryEmbedding`**: RoPE complex cache (freqs_cis)
|
|
||||||
- **`RMSNorm`**: Layer normalization
|
|
||||||
|
|
||||||
### 4. Training Module
|
|
||||||
|
|
||||||
#### 4.1 Training Context (`train_context.py`)
|
|
||||||
- **`TrainContext`**: Dataclass holding model, optimizer, dataloader, strategy, scheduler, checkpoint state
|
|
||||||
- **`TrainContextBuilder`**: Builder pattern — takes checkpoint for resume, builds all components
|
|
||||||
|
|
||||||
#### 4.2 Trainer (`trainer.py`)
|
|
||||||
|
|
||||||
The training loop is nested: **epoch** → **batch** (with step phase interspersed):
|
|
||||||
|
|
||||||
```
|
|
||||||
on_train_begin
|
|
||||||
on_epoch_begin
|
|
||||||
for each accumulation window of batches: ← step phase
|
|
||||||
on_step_begin
|
|
||||||
for each batch in window: ← batch phase
|
|
||||||
on_batch_begin → strategy(batch) → loss → backward → on_batch_end
|
|
||||||
iteration += 1
|
|
||||||
on_step_end
|
|
||||||
optimizer.step() → zero_grad
|
|
||||||
|
|
||||||
on_epoch_end
|
|
||||||
on_train_end
|
|
||||||
```
|
|
||||||
|
|
||||||
Key points:
|
|
||||||
- `on_step_*` fires every `accumulation_steps` batches, wrapping optimizer step AFTER the hook
|
|
||||||
- `on_batch_*` fires every batch, wrapping loss computation
|
|
||||||
- `GradientClippingCallback` fires on `on_step_end`
|
|
||||||
- LR scheduler steps inline (no `SchedulerCallback` class)
|
|
||||||
|
|
||||||
#### 4.3 Strategy (`strategy.py`)
|
|
||||||
- **`SEQStrategy`**: Next-token prediction, cross-entropy with label smoothing
|
|
||||||
- **`SFTStrategy`**: Supervised fine-tuning with loss masking
|
|
||||||
- **`DPOStrategy`**: Direct Preference Optimization with reference model
|
|
||||||
- **`GRPOStrategy`**: Group Relative Policy Optimization with clipped ratio
|
|
||||||
|
|
||||||
#### 4.4 Scheduler (`schedule.py`)
|
|
||||||
- **`CosineScheduler`**: Cosine decay + linear warmup
|
|
||||||
- **`SGDRScheduler`**: Cosine annealing with warm restarts
|
|
||||||
- Created by `SchedulerFactory` and bound to optimizer
|
|
||||||
|
|
||||||
#### 4.5 Callbacks
|
|
||||||
- **`CheckpointCallback`**: Saves safetensors at `ckpt_interval` iterations
|
|
||||||
- **`ProgressBarCallback`**: tqdm progress display
|
|
||||||
- **`MetricLoggerCallback`**: Writes JSONL metrics to `{ckpt_dir}/logs/`
|
|
||||||
- **`GradientClippingCallback`**: `clip_grad_norm_` on `on_step_end`
|
|
||||||
|
|
||||||
### 5. Inference Module
|
|
||||||
|
|
||||||
#### 5.1 Inference Engine (`engine.py`)
|
|
||||||
- **`InferenceEngine`**: Facade over scheduler; provides `generate()`, `generate_with_request()`, `generate_async()`
|
|
||||||
- Accepts `prompt: str | List[str]`, returns generator (stream) or string (non-stream)
|
|
||||||
|
|
||||||
#### 5.2 Scheduler 4-Phase Loop (`scheduler.py`)
|
|
||||||
|
|
||||||
Background thread runs continuously:
|
|
||||||
|
|
||||||
```
|
|
||||||
1. Cleanup → Remove finished tasks, free KV cache pages
|
|
||||||
2. Refill → Pop from waiting_queue, alloc pages, add to active
|
|
||||||
3. Prefill → Group active tasks by prompt_len, run full forward pass
|
|
||||||
4. Decode → Pick largest same-position group, run single-token forward
|
|
||||||
```
|
|
||||||
|
|
||||||
- **`Task`**: Tracks prompt_ids, output_ids, status (PENDING/RUNNING/FINISHED/ABORTED)
|
|
||||||
- **`KVCache`**: Facade over `Allocator` + `PrefixCache` + `PagePool` + `Storage` for paged KV cache
|
|
||||||
- **`KvcacheView`**: Batch view bundling cache + page table for attention layers
|
|
||||||
- **`sample()`**: Temperature → top-k → top-p → multinomial
|
|
||||||
|
|
||||||
#### 5.3 Server (`server.py`)
|
|
||||||
- FastAPI with OpenAI `/v1/chat/completions` and Anthropic `/v1/messages` endpoints
|
|
||||||
- Streaming via SSE, health check at `/health`, stats at `/stats`
|
|
||||||
|
|
||||||
### 6. Tokenizer Module
|
|
||||||
|
|
||||||
- **`AutoTokenizer`**: Wraps HuggingFace tokenizers (BBPE); `encode`/`decode`/`apply_chat_template`
|
|
||||||
- **`ChatTemplate`**: Jinja2-based template rendering for multi-turn chat
|
|
||||||
|
|
||||||
### 7. Factory & Parallel
|
|
||||||
|
|
||||||
- **`Registry` / `BaseFactory`**: Decorator-based component registration
|
|
||||||
- **`spawn_parallel_fn`**: Multi-process DDP launcher with NCCL backend
|
|
||||||
- **`ParallelModel` / `ColumnParallelLinear` / `RowParallelLinear`**: Tensor model parallelism
|
|
||||||
|
|
||||||
## Training Data Flow — Detailed Steps
|
|
||||||
|
|
||||||
1. **Data Preparation**
|
|
||||||
- Raw text → token IDs via `AutoTokenizer.encode()`
|
|
||||||
- Save as `.h5` files (groups of tensor lists per data key)
|
|
||||||
|
|
||||||
2. **Dataset Loading**
|
|
||||||
- `BaseDataset.load()` calls `load_h5()`, builds `MultiSegmentFetcher`
|
|
||||||
- Sliding window of `window_size` with `stride` determines sample boundaries
|
|
||||||
|
|
||||||
3. **Sampling & Batching**
|
|
||||||
- `ResumableDistributedSampler` produces shuffled index sequences
|
|
||||||
- `DataLoader` fetches `[batch_size, window_size]` tensors via `__getitem__`
|
|
||||||
|
|
||||||
4. **Strategy Forward**
|
|
||||||
- Strategy receives batch, calls `Transformer.forward()` for logits
|
|
||||||
- Computes task-specific loss (cross-entropy, DPO, GRPO)
|
|
||||||
|
|
||||||
5. **Backward & Accumulation**
|
|
||||||
- `loss = raw_loss / accumulation_steps`
|
|
||||||
- `loss.backward()` accumulates gradients
|
|
||||||
- Every `accumulation_steps` batches: `optimizer.step()` → `zero_grad()`
|
|
||||||
- Every batch: `scheduler.step()` updates learning rate
|
|
||||||
|
|
||||||
6. **Checkpoint**
|
|
||||||
- `CheckpointCallback` saves `model.state_dict()` + metadata to safetensors at `ckpt_interval` iterations
|
|
||||||
- Does NOT save optimizer/scheduler state (resume resets those)
|
|
||||||
|
|
||||||
## Inference Data Flow — Detailed Steps
|
|
||||||
|
|
||||||
1. **Model Loading**
|
|
||||||
- `AutoModel.from_pretrained(path)` loads weights from safetensors
|
|
||||||
- `torch.inference_mode()` wraps generation
|
|
||||||
|
|
||||||
2. **Prompt Construction**
|
|
||||||
- Messages → `apply_chat_template(messages, tokenize=False)` → prompt string
|
|
||||||
- `tokenizer.encode(prompt)` → token IDs (truncated to `max_prompt_len`)
|
|
||||||
|
|
||||||
3. **Continuous Batching Loop**
|
|
||||||
- **Cleanup**: Finished tasks → `stream_callback(STOP)`, free KV pages
|
|
||||||
- **Refill**: Pop from waiting queue, `PagePool.task_alloc()` for prompt pages
|
|
||||||
- **Prefill**: Group by prompt length, run full forward with `start_pos=0`
|
|
||||||
- **Decode**: Pick position group with most tasks, single-token forward:
|
|
||||||
- Model forward → `logits` → `sample()` → next token ID
|
|
||||||
- Append to `output_ids`, update `output_tokens`
|
|
||||||
- `PagePool.task_alloc()` allocates pages as needed
|
|
||||||
- `stream_callback(token)` for streaming clients
|
|
||||||
|
|
||||||
4. **Output**
|
|
||||||
- `tokenizer.decode(output_ids)` → text
|
|
||||||
- Return to caller (streaming: token-by-token; non-streaming: complete string)
|
|
||||||
|
|
||||||
## Checkpoint & Serialization
|
|
||||||
|
|
||||||
- **Training Checkpoint**: safetensors weights + epoch/iteration metadata. Optimizer/scheduler state is NOT persisted.
|
|
||||||
- **Inference Loading**: `AutoModel.from_pretrained()` loads from the same safetensors format.
|
|
||||||
- **Dataset Serialization**: HDF5 with shared memory support for large-scale pre-training data.
|
|
||||||
|
|
||||||
> Document Update Time: 2026-05-14
|
|
||||||
@@ -1,779 +0,0 @@
|
|||||||
## 1. Why I Created This Project
|
|
||||||
|
|
||||||
There are many large language models on the market today, such as GPT, LLaMA, and others, with tens of billions or even hundreds of billions of parameters. But honestly, these models have extremely high hardware requirements, making them inaccessible for ordinary developers. I thought: **Can we create a model that is both useful and can run on ordinary computers?** This is also what most people currently hope for - a locally deployable AI project that achieves complete privatization while maintaining some level of intelligence.
|
|
||||||
|
|
||||||
Thus, the AstrAI project was born - 1B parameters, Chinese-English bilingual, supporting dialogue, text generation, and the training code is open source!
|
|
||||||
|
|
||||||
## 2. System Architecture
|
|
||||||
|
|
||||||
```mermaid
|
|
||||||
classDiagram
|
|
||||||
namespace config {
|
|
||||||
class ModelConfig {
|
|
||||||
+int vocab_size
|
|
||||||
+int dim
|
|
||||||
+int n_layers
|
|
||||||
+float norm_eps
|
|
||||||
+int dim_ffn
|
|
||||||
+bool tie_weight
|
|
||||||
+int max_len
|
|
||||||
+float rope_theta
|
|
||||||
+int n_heads
|
|
||||||
+int n_kv_heads
|
|
||||||
+bool use_qk_norm
|
|
||||||
+bool use_gated_attention
|
|
||||||
+load(config_path) ModelConfig
|
|
||||||
+save(config_path)
|
|
||||||
}
|
|
||||||
|
|
||||||
class TrainConfig {
|
|
||||||
+nn.Module model
|
|
||||||
+str strategy
|
|
||||||
+Dataset dataset
|
|
||||||
+Callable optimizer_fn
|
|
||||||
+Callable scheduler_fn
|
|
||||||
+int n_epoch
|
|
||||||
+int batch_size
|
|
||||||
+int accumulation_steps
|
|
||||||
+float max_grad_norm
|
|
||||||
+int start_epoch
|
|
||||||
+int start_batch
|
|
||||||
+str ckpt_dir
|
|
||||||
+int ckpt_interval
|
|
||||||
+int random_seed
|
|
||||||
+int num_workers
|
|
||||||
+int prefetch_factor
|
|
||||||
+bool pin_memory
|
|
||||||
+int nprocs
|
|
||||||
+str backend
|
|
||||||
+str master_addr
|
|
||||||
+str master_port
|
|
||||||
+Callable parallel_wrapper
|
|
||||||
+Callable state_dict_fn
|
|
||||||
+str device_type
|
|
||||||
+dict extra_kwargs
|
|
||||||
+validate()
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
namespace dataset {
|
|
||||||
class BaseDataset {
|
|
||||||
+int window_size
|
|
||||||
+int stride
|
|
||||||
+BaseStorage storage
|
|
||||||
+load(load_path, storage_type, tokenizer)
|
|
||||||
+__getitem__(index)
|
|
||||||
+__len__()
|
|
||||||
}
|
|
||||||
|
|
||||||
class SEQDataset {
|
|
||||||
+__getitem__(index) Dict
|
|
||||||
}
|
|
||||||
|
|
||||||
class SFTDataset {
|
|
||||||
+__getitem__(index) Dict
|
|
||||||
}
|
|
||||||
|
|
||||||
class DPODataset {
|
|
||||||
+__getitem__(index) Dict
|
|
||||||
}
|
|
||||||
|
|
||||||
class GRPODataset {
|
|
||||||
+__getitem__(index) Dict
|
|
||||||
}
|
|
||||||
|
|
||||||
class BaseSegmentFetcher {
|
|
||||||
+List[Tensor] segments
|
|
||||||
+List[int] cum_lengths
|
|
||||||
+int total_length
|
|
||||||
+fetch_data(begin_idx, end_idx) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class BaseStorage {
|
|
||||||
+MultiSegmentFetcher _fetcher
|
|
||||||
+keys (property)
|
|
||||||
+load(load_path, tokenizer)
|
|
||||||
+fetch(begin, end, keys)
|
|
||||||
+__len__()
|
|
||||||
}
|
|
||||||
|
|
||||||
class H5Storage {
|
|
||||||
+load(load_path, tokenizer)
|
|
||||||
+fetch(begin, end, keys) Dict
|
|
||||||
+keys() List
|
|
||||||
}
|
|
||||||
|
|
||||||
class JSONStorage {
|
|
||||||
+load(load_path, tokenizer)
|
|
||||||
+fetch(begin, end, keys) Dict
|
|
||||||
+keys() List
|
|
||||||
}
|
|
||||||
|
|
||||||
class MultiSegmentFetcher {
|
|
||||||
+Dict multi_fetchers
|
|
||||||
+List multi_keys
|
|
||||||
+key_fetch(begin_idx, end_idx, keys) Dict
|
|
||||||
+fetch_data(begin_idx, end_idx) Dict
|
|
||||||
}
|
|
||||||
|
|
||||||
class ResumableDistributedSampler {
|
|
||||||
+int epoch
|
|
||||||
+int iter
|
|
||||||
}
|
|
||||||
|
|
||||||
class DatasetFactory {
|
|
||||||
+Registry _registry
|
|
||||||
+register(name) decorator
|
|
||||||
+create(train_type, window_size, stride) BaseDataset
|
|
||||||
+load(train_type, load_path, window_size, stride) BaseDataset
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
namespace serialization {
|
|
||||||
class Checkpoint {
|
|
||||||
+dict state_dict
|
|
||||||
+int epoch
|
|
||||||
+int iteration
|
|
||||||
+save(save_dir)
|
|
||||||
+load(save_dir) Checkpoint
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
namespace model {
|
|
||||||
class AutoModel {
|
|
||||||
+ModelConfig config
|
|
||||||
+Registry _registry
|
|
||||||
+register(model_type) decorator
|
|
||||||
+get_component_class(model_type) Type
|
|
||||||
+from_pretrained(path, disable_random_init) nn.Module
|
|
||||||
+save_pretrained(save_directory)
|
|
||||||
+to(*args, **kwargs) Self
|
|
||||||
}
|
|
||||||
|
|
||||||
class Transformer {
|
|
||||||
+ModelConfig config
|
|
||||||
+RotaryEmbedding rotary_embedding
|
|
||||||
+Embedding embed_tokens
|
|
||||||
+ModuleList layers
|
|
||||||
+RMSNorm norm
|
|
||||||
+Linear lm_head
|
|
||||||
+forward(input_ids, input_mask, paged_cache, position_ids) Tensor
|
|
||||||
+load_state_dict(state_dict)
|
|
||||||
+state_dict()
|
|
||||||
}
|
|
||||||
|
|
||||||
class DecoderBlock {
|
|
||||||
+GQA attention
|
|
||||||
+RMSNorm input_norm
|
|
||||||
+MLP mlp
|
|
||||||
+RMSNorm post_attention_norm
|
|
||||||
+forward(x, rotary_emb, attention_mask, paged_cache) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class GQA {
|
|
||||||
+int n_heads
|
|
||||||
+int n_kv_heads
|
|
||||||
+int head_dim
|
|
||||||
+Linear q_proj, k_proj, v_proj, o_proj
|
|
||||||
+RMSNorm q_norm, k_norm
|
|
||||||
+forward(x, rotary_emb, attn_mask, paged_cache) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class MLA {
|
|
||||||
+int n_heads
|
|
||||||
+int n_kv_heads
|
|
||||||
+int head_dim
|
|
||||||
+int kv_lora_rank
|
|
||||||
+int qk_nope_head_dim
|
|
||||||
+int qk_rope_head_dim
|
|
||||||
+Linear q_proj, kv_a_proj, kv_b_proj
|
|
||||||
+Linear o_proj
|
|
||||||
+RMSNorm kv_norm
|
|
||||||
+forward(x, rotary_emb, attn_mask, paged_cache) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class MLP {
|
|
||||||
+Linear up, gate, down
|
|
||||||
+forward(x) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class RMSNorm {
|
|
||||||
+Parameter weight
|
|
||||||
+float norm_eps
|
|
||||||
+forward(x) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class Linear {
|
|
||||||
+Parameter weight
|
|
||||||
+Parameter bias
|
|
||||||
+forward(x) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class RotaryEmbedding {
|
|
||||||
+int dim
|
|
||||||
+int max_len
|
|
||||||
+float base
|
|
||||||
+forward(x, position_ids=None) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class Embedding {
|
|
||||||
+Parameter weight
|
|
||||||
+forward(x) Tensor
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
namespace tokenize {
|
|
||||||
class AutoTokenizer {
|
|
||||||
+vocab_size int
|
|
||||||
+encode(tokens, out_ids, add_special_tokens) List[int]
|
|
||||||
+decode(tokens, skip_special_tokens) str
|
|
||||||
+__getattr__(name) Any (bos_id, eos_id, pad_id, stop_ids)
|
|
||||||
+apply_chat_template(messages, tokenize) Union[str, List[int]]
|
|
||||||
+set_chat_template(template)
|
|
||||||
+load(path)
|
|
||||||
+from_pretrained(path) AutoTokenizer
|
|
||||||
+save_pretrained(save_path)
|
|
||||||
}
|
|
||||||
|
|
||||||
class ChatTemplate {
|
|
||||||
+String template_str
|
|
||||||
+render(messages, system_prompt, **extra_variables) str
|
|
||||||
+from_string(template) ChatTemplate
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
namespace factory {
|
|
||||||
class Registry {
|
|
||||||
+Dict _entries
|
|
||||||
+register(name, component_cls, category, priority)
|
|
||||||
+get(name) Type
|
|
||||||
+list_names() List[str]
|
|
||||||
}
|
|
||||||
|
|
||||||
class BaseFactory {
|
|
||||||
+Registry _registry
|
|
||||||
+register(name, category, priority) decorator
|
|
||||||
+create(name, *args, **kwargs) T
|
|
||||||
+list_registered() list
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
namespace trainer {
|
|
||||||
class Trainer {
|
|
||||||
+TrainConfig train_config
|
|
||||||
+List[TrainCallback] callbacks
|
|
||||||
+train(checkpoint)
|
|
||||||
+_build_context(checkpoint) TrainContext
|
|
||||||
+_get_default_callbacks() List[TrainCallback]
|
|
||||||
}
|
|
||||||
|
|
||||||
class TrainContext {
|
|
||||||
+nn.Module model
|
|
||||||
+BaseStrategy strategy
|
|
||||||
+DataLoader dataloader
|
|
||||||
+Optimizer optimizer
|
|
||||||
+LRScheduler scheduler
|
|
||||||
+Checkpoint checkpoint
|
|
||||||
+int epoch
|
|
||||||
+int iteration
|
|
||||||
+float loss
|
|
||||||
+int world_size
|
|
||||||
+int rank
|
|
||||||
}
|
|
||||||
|
|
||||||
class TrainContextBuilder {
|
|
||||||
+TrainConfig config
|
|
||||||
+with_checkpoint(checkpoint) TrainContextBuilder
|
|
||||||
+build() TrainContext
|
|
||||||
}
|
|
||||||
|
|
||||||
class BaseStrategy {
|
|
||||||
+nn.Module model
|
|
||||||
+str device
|
|
||||||
+compute_loss(batch) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class StrategyFactory {
|
|
||||||
+Registry _registry
|
|
||||||
+register(name) decorator
|
|
||||||
+create(model, train_type, device, **kwargs) BaseStrategy
|
|
||||||
}
|
|
||||||
|
|
||||||
class SEQStrategy {
|
|
||||||
+float label_smoothing
|
|
||||||
+compute_loss(batch) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class SFTStrategy {
|
|
||||||
+float label_smoothing
|
|
||||||
+compute_loss(batch) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class DPOStrategy {
|
|
||||||
+nn.Module ref_model
|
|
||||||
+float beta
|
|
||||||
+str reduction
|
|
||||||
+compute_loss(batch) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class GRPOStrategy {
|
|
||||||
+nn.Module ref_model
|
|
||||||
+float clip_eps
|
|
||||||
+float kl_coef
|
|
||||||
+int group_size
|
|
||||||
+str reduction
|
|
||||||
+int sync_interval
|
|
||||||
+compute_loss(batch) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class BaseScheduler {
|
|
||||||
+get_lr() List[float]
|
|
||||||
+step()
|
|
||||||
}
|
|
||||||
|
|
||||||
class SchedulerFactory {
|
|
||||||
+Registry _registry
|
|
||||||
+register(name) decorator
|
|
||||||
+create(optimizer, schedule_type, **kwargs) BaseScheduler
|
|
||||||
}
|
|
||||||
|
|
||||||
class CosineScheduler {
|
|
||||||
+int warmup_steps
|
|
||||||
+int lr_decay_steps
|
|
||||||
+float min_rate
|
|
||||||
}
|
|
||||||
|
|
||||||
class SGDRScheduler {
|
|
||||||
+int warmup_steps
|
|
||||||
+int cycle_length
|
|
||||||
+float min_rate
|
|
||||||
+int t_mult
|
|
||||||
}
|
|
||||||
|
|
||||||
class TrainCallback {
|
|
||||||
+on_train_begin(context)
|
|
||||||
+on_train_end(context)
|
|
||||||
+on_epoch_begin(context)
|
|
||||||
+on_epoch_end(context)
|
|
||||||
+on_step_begin(context)
|
|
||||||
+on_step_end(context)
|
|
||||||
+on_batch_begin(context)
|
|
||||||
+on_batch_end(context)
|
|
||||||
+on_error(context)
|
|
||||||
}
|
|
||||||
|
|
||||||
class GradientClippingCallback {
|
|
||||||
+float max_grad_norm
|
|
||||||
+on_step_begin(context)
|
|
||||||
}
|
|
||||||
|
|
||||||
class CheckpointCallback {
|
|
||||||
+str save_dir
|
|
||||||
+int interval
|
|
||||||
+_save_checkpoint(context)
|
|
||||||
+on_batch_end(context)
|
|
||||||
+on_train_end(context)
|
|
||||||
+on_error(context)
|
|
||||||
}
|
|
||||||
|
|
||||||
class ProgressBarCallback {
|
|
||||||
+int num_epoch
|
|
||||||
+on_epoch_begin(context)
|
|
||||||
+on_batch_end(context)
|
|
||||||
+on_epoch_end(context)
|
|
||||||
}
|
|
||||||
|
|
||||||
class MetricLoggerCallback {
|
|
||||||
+str log_dir
|
|
||||||
+int save_interval
|
|
||||||
+on_batch_end(context)
|
|
||||||
+on_train_end(context)
|
|
||||||
}
|
|
||||||
|
|
||||||
class CallbackFactory {
|
|
||||||
+Registry _registry
|
|
||||||
+register(name) decorator
|
|
||||||
+create(name, **kwargs) TrainCallback
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
namespace inference {
|
|
||||||
class InferenceEngine {
|
|
||||||
+nn.Module model
|
|
||||||
+AutoTokenizer tokenizer
|
|
||||||
+InferenceScheduler scheduler
|
|
||||||
+generate(prompt, stream, max_tokens, temperature, top_p, top_k) Union[Generator, str, List[str]]
|
|
||||||
+generate_with_request(request) Union[Generator, str, List[str]]
|
|
||||||
+generate_async(prompt, max_tokens, temperature, top_p, top_k) AsyncGenerator
|
|
||||||
+get_stats() Dict
|
|
||||||
+shutdown()
|
|
||||||
}
|
|
||||||
|
|
||||||
class InferenceScheduler {
|
|
||||||
+nn.Module model
|
|
||||||
+AutoTokenizer tokenizer
|
|
||||||
+KVCache _page_cache
|
|
||||||
+int max_batch_size
|
|
||||||
+int max_seq_len
|
|
||||||
+int max_prompt_len
|
|
||||||
+int page_size
|
|
||||||
+TaskManager _task_mgr
|
|
||||||
+add_task(prompt, max_tokens, temperature, top_p, top_k, stream_callback) str
|
|
||||||
+remove_task(task_id)
|
|
||||||
+start()
|
|
||||||
+stop()
|
|
||||||
+get_stats() Dict
|
|
||||||
}
|
|
||||||
|
|
||||||
class Allocator {
|
|
||||||
+int _free_mask
|
|
||||||
+int refs_count
|
|
||||||
+LRU _lru
|
|
||||||
+alloc() int
|
|
||||||
+free(idx, keep_cached)
|
|
||||||
+inc_ref(idx)
|
|
||||||
+touch(idx)
|
|
||||||
+ref_count(idx) int
|
|
||||||
}
|
|
||||||
|
|
||||||
class PrefixCache {
|
|
||||||
+int _page_size
|
|
||||||
+evict(page_idx)
|
|
||||||
+has_page(idx) bool
|
|
||||||
+lookup(token_ids) List[int]
|
|
||||||
+record(page_idx, token_ids, logical_page_idx)
|
|
||||||
}
|
|
||||||
|
|
||||||
class PagePool {
|
|
||||||
-Allocator _alloc
|
|
||||||
-PrefixCache _prefix
|
|
||||||
+alloc() int
|
|
||||||
+free(idx)
|
|
||||||
+inc_ref(idx)
|
|
||||||
+lookup(token_ids) List[int]
|
|
||||||
+record(page_idx, token_ids, logical_page_idx)
|
|
||||||
}
|
|
||||||
|
|
||||||
class Storage {
|
|
||||||
+int n_layers
|
|
||||||
+int page_size
|
|
||||||
+int head_dim
|
|
||||||
+int n_kv_heads
|
|
||||||
+Tensor k_cache
|
|
||||||
+Tensor v_cache
|
|
||||||
+write(layer_id, page_table, start_pos, k, v)
|
|
||||||
+gather(layer_id, page_table, total_len) Tuple[Tensor, Tensor]
|
|
||||||
}
|
|
||||||
|
|
||||||
class KVCache {
|
|
||||||
-PagePool _pool
|
|
||||||
-Storage _storage
|
|
||||||
-TaskTable _table
|
|
||||||
+int page_size
|
|
||||||
+task_alloc(task_id, prompt_ids) bool
|
|
||||||
+task_free(task_id)
|
|
||||||
+task_extend(task_id, pos) bool
|
|
||||||
+task_cached(task_id) int
|
|
||||||
+task_record_hashes(task_id, prompt_ids, start_logical_page)
|
|
||||||
+make_table_tensor(task_ids, device) Tensor
|
|
||||||
+bind(page_table, total_len) KvcacheView
|
|
||||||
}
|
|
||||||
|
|
||||||
class KvcacheView {
|
|
||||||
-Storage _storage
|
|
||||||
+Tensor _page_table
|
|
||||||
+int _total_len
|
|
||||||
+write(layer_id, k, v)
|
|
||||||
+gather(layer_id) Tuple[Tensor, Tensor]
|
|
||||||
}
|
|
||||||
|
|
||||||
class TaskTable {
|
|
||||||
+set(task_id, page_table, cached)
|
|
||||||
+get(task_id) List[int]
|
|
||||||
+get_cached(task_id) int
|
|
||||||
+get_ref(task_id) List[int]
|
|
||||||
+pop(task_id) Tuple[List[int], int]
|
|
||||||
+table_tensor(task_ids, device) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class Task {
|
|
||||||
+str task_id
|
|
||||||
+List prompt_ids
|
|
||||||
+int max_tokens
|
|
||||||
+float temperature
|
|
||||||
+float top_p
|
|
||||||
+int top_k
|
|
||||||
+TaskStatus status
|
|
||||||
+List output_ids
|
|
||||||
+int input_tokens
|
|
||||||
+int output_tokens
|
|
||||||
+float arrival_time
|
|
||||||
+float finish_time
|
|
||||||
+Callable stream_callback
|
|
||||||
+int next_pos
|
|
||||||
+is_finished(stop_ids) bool
|
|
||||||
}
|
|
||||||
|
|
||||||
class TaskStatus {
|
|
||||||
<<enumeration>>
|
|
||||||
PENDING
|
|
||||||
RUNNING
|
|
||||||
FINISHED
|
|
||||||
ABORTED
|
|
||||||
}
|
|
||||||
|
|
||||||
class GenerationRequest {
|
|
||||||
+List[Dict] messages
|
|
||||||
+int top_k
|
|
||||||
+float top_p
|
|
||||||
+float temperature
|
|
||||||
+Optional[int] max_tokens
|
|
||||||
+bool stream
|
|
||||||
}
|
|
||||||
|
|
||||||
class BaseSamplingStrategy {
|
|
||||||
<<abstract>>
|
|
||||||
+apply(logits, filter_value) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class TemperatureStrategy {
|
|
||||||
+float temperature
|
|
||||||
+apply(logits, filter_value) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class TopKStrategy {
|
|
||||||
+int top_k
|
|
||||||
+apply(logits, filter_value) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class TopPStrategy {
|
|
||||||
+float top_p
|
|
||||||
+apply(logits, filter_value) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class SamplingPipeline {
|
|
||||||
+List strategies
|
|
||||||
+apply(logits, filter_value) Tensor
|
|
||||||
+sample(logits, filter_value) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class GenerateResult {
|
|
||||||
+List[Tuple[int, str]] tokens
|
|
||||||
+List[str] results
|
|
||||||
+List[bool] _done
|
|
||||||
+append(token, idx)
|
|
||||||
+get_results() List[str]
|
|
||||||
+pop_all() List[str]
|
|
||||||
+wait(timeout) bool
|
|
||||||
+wait_completion()
|
|
||||||
}
|
|
||||||
|
|
||||||
class ChatMessage {
|
|
||||||
+str role
|
|
||||||
+str content
|
|
||||||
}
|
|
||||||
|
|
||||||
class ChatCompletionRequest {
|
|
||||||
+List[ChatMessage] messages
|
|
||||||
+float temperature
|
|
||||||
+float top_p
|
|
||||||
+int top_k
|
|
||||||
+int max_tokens
|
|
||||||
+bool stream
|
|
||||||
+Optional[str] stop
|
|
||||||
+Optional[int] n
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
namespace parallel {
|
|
||||||
class Functions {
|
|
||||||
+spawn_parallel_fn(fn, nprocs)
|
|
||||||
+setup_parallel(rank, world_size, backend, master_addr, master_port, device_type)
|
|
||||||
+get_current_device() str
|
|
||||||
+get_world_size() int
|
|
||||||
+get_rank() int
|
|
||||||
}
|
|
||||||
|
|
||||||
class ParallelModel {
|
|
||||||
+dist.ProcessGroup process_group
|
|
||||||
+int rank
|
|
||||||
+int world_size
|
|
||||||
}
|
|
||||||
|
|
||||||
class ColumnParallelLinear {
|
|
||||||
+forward(x) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class RowParallelLinear {
|
|
||||||
+forward(x) Tensor
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
%% Relationships
|
|
||||||
TrainConfig --> BaseDataset : uses
|
|
||||||
TrainConfig ..> BaseStrategy : selects
|
|
||||||
StrategyFactory ..> BaseStrategy : creates
|
|
||||||
BaseStrategy <|-- SEQStrategy
|
|
||||||
BaseStrategy <|-- SFTStrategy
|
|
||||||
BaseStrategy <|-- DPOStrategy
|
|
||||||
BaseStrategy <|-- GRPOStrategy
|
|
||||||
DPOStrategy --> Transformer : uses
|
|
||||||
GRPOStrategy --> Transformer : uses
|
|
||||||
Trainer --> TrainConfig : uses
|
|
||||||
Trainer --> TrainContextBuilder : uses
|
|
||||||
Trainer --> TrainCallback : manages
|
|
||||||
TrainContextBuilder --> TrainContext : creates
|
|
||||||
TrainContextBuilder --> StrategyFactory : uses
|
|
||||||
Checkpoint ..> Checkpoint : serializes
|
|
||||||
TrainContext --> Checkpoint : manages
|
|
||||||
TrainContext --> BaseStrategy : uses
|
|
||||||
TrainContext --> BaseScheduler : uses
|
|
||||||
SchedulerFactory ..> BaseScheduler : creates
|
|
||||||
BaseScheduler <|-- CosineScheduler
|
|
||||||
BaseScheduler <|-- SGDRScheduler
|
|
||||||
CallbackFactory ..> TrainCallback : creates
|
|
||||||
TrainCallback <|-- GradientClippingCallback
|
|
||||||
TrainCallback <|-- CheckpointCallback
|
|
||||||
TrainCallback <|-- ProgressBarCallback
|
|
||||||
TrainCallback <|-- MetricLoggerCallback
|
|
||||||
PagePool --> Allocator : composes
|
|
||||||
PagePool --> PrefixCache : composes
|
|
||||||
KVCache --> PagePool : composes
|
|
||||||
KVCache --> Storage : composes
|
|
||||||
KVCache --> TaskTable : composes
|
|
||||||
KvcacheView --> Storage : wraps
|
|
||||||
InferenceEngine --> InferenceScheduler : uses
|
|
||||||
InferenceEngine --> GenerationRequest : uses
|
|
||||||
InferenceEngine --> GenerateResult : creates
|
|
||||||
InferenceScheduler --> Task : manages
|
|
||||||
InferenceScheduler --> TaskStatus : uses
|
|
||||||
InferenceScheduler --> KVCache : uses
|
|
||||||
InferenceScheduler --> Transformer : uses
|
|
||||||
Task --> TaskStatus : uses
|
|
||||||
InferenceEngine --> Transformer : uses
|
|
||||||
BaseSamplingStrategy <|-- TemperatureStrategy
|
|
||||||
BaseSamplingStrategy <|-- TopKStrategy
|
|
||||||
BaseSamplingStrategy <|-- TopPStrategy
|
|
||||||
SamplingPipeline --> BaseSamplingStrategy : composes
|
|
||||||
BaseDataset <|-- SEQDataset
|
|
||||||
BaseDataset <|-- SFTDataset
|
|
||||||
BaseDataset <|-- DPODataset
|
|
||||||
BaseDataset <|-- GRPODataset
|
|
||||||
DatasetFactory ..> BaseDataset : creates
|
|
||||||
BaseStorage <|-- H5Storage
|
|
||||||
BaseStorage <|-- JSONStorage
|
|
||||||
BaseDataset --> BaseStorage : uses
|
|
||||||
MultiSegmentFetcher --> BaseSegmentFetcher : uses
|
|
||||||
AutoModel <|-- Transformer
|
|
||||||
AutoModel --> ModelConfig : contains
|
|
||||||
Transformer --> DecoderBlock : uses
|
|
||||||
Transformer --> RotaryEmbedding : uses
|
|
||||||
Transformer --> Embedding : uses
|
|
||||||
DecoderBlock --> GQA : uses
|
|
||||||
DecoderBlock --> MLP : uses
|
|
||||||
DecoderBlock --> RMSNorm : uses
|
|
||||||
TrainContextBuilder --> ResumableDistributedSampler : creates
|
|
||||||
ResumableDistributedSampler --> BaseDataset : samples
|
|
||||||
ParallelModel <|-- RowParallelLinear
|
|
||||||
ParallelModel <|-- ColumnParallelLinear
|
|
||||||
AutoTokenizer --> ChatTemplate : uses
|
|
||||||
BaseFactory <|-- AutoModel
|
|
||||||
BaseFactory <|-- DatasetFactory
|
|
||||||
BaseFactory <|-- StrategyFactory
|
|
||||||
BaseFactory <|-- SchedulerFactory
|
|
||||||
BaseFactory <|-- CallbackFactory
|
|
||||||
```
|
|
||||||
|
|
||||||
### Module Overview
|
|
||||||
|
|
||||||
| Module | Components | Description |
|
|
||||||
|--------|------------|-------------|
|
|
||||||
| **astrai.config** | ModelConfig, TrainConfig | Configuration management |
|
|
||||||
| **astrai.dataset** | BaseDataset, SEQDataset, SFTDataset, DPODataset, GRPODataset, BaseStorage, H5Storage, JSONStorage, BaseSegmentFetcher, MultiSegmentFetcher, ResumableDistributedSampler, DatasetFactory, save_h5, load_h5 | Dataset loading and management |
|
|
||||||
| **astrai.serialization** | Checkpoint | Model serialization and checkpoint management |
|
|
||||||
| **astrai.model** | AutoModel, Transformer, DecoderBlock, GQA, MLA, MLP, RMSNorm, Linear, RotaryEmbedding, Embedding | Neural network model |
|
|
||||||
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
|
||||||
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy, StrategyFactory, BaseScheduler, SchedulerFactory, TrainCallback, CallbackFactory | Training workflow management |
|
|
||||||
| **astrai.inference** | InferenceEngine, InferenceScheduler, KVCache, KvcacheView, Allocator, PrefixCache, PagePool, Storage, TaskTable, Task, TaskStatus, GenerationRequest, BaseSamplingStrategy, TemperatureStrategy, TopKStrategy, TopPStrategy, SamplingPipeline, ChatMessage, ChatCompletionRequest | Inference service with continuous batching and paged KV cache |
|
|
||||||
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank, get_world_size, get_current_device, ParallelModel, ColumnParallelLinear, RowParallelLinear | Distributed parallel |
|
|
||||||
| **astrai.factory** | Registry, BaseFactory | Generic component registration |
|
|
||||||
|
|
||||||
### Design Patterns
|
|
||||||
|
|
||||||
| Pattern | Classes | Purpose |
|
|
||||||
|---------|---------|---------|
|
|
||||||
| **Strategy** | `BaseStrategy`, `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy`, `StrategyFactory` | Flexible training strategy switching, supports SEQ/SFT/DPO/GRPO |
|
|
||||||
| **Builder** | `TrainContextBuilder` | Chain-building training context, step-by-step initialization of components |
|
|
||||||
| **Factory** | `StrategyFactory`, `SchedulerFactory`, `DatasetFactory`, `CallbackFactory`, `BaseFactory` | Decorator registration mechanism, dynamically create training strategies, schedulers, datasets, and callbacks |
|
|
||||||
| **Observer** | `TrainCallback`, `CallbackFactory` | Callback mechanism for training process monitoring (checkpoint, early stopping, metrics) |
|
|
||||||
| **Context** | `TrainContext` | Training process state container with model, optimizer, scheduler and checkpoint |
|
|
||||||
| **Registry** | `BaseFactory`, `Registry` | Generic component registration with category and priority support |
|
|
||||||
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with O(1) alloc/free via bitmask + LRU eviction |
|
|
||||||
| **Strategy (Sampling)** | `BaseSamplingStrategy`, `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations with temperature, top-k, top-p |
|
|
||||||
| **Producer-Consumer** | `InferenceScheduler`, `Task`, `waiting_queue`, `active_tasks` | Continuous batching with dynamic task queue management |
|
|
||||||
| **Event-Driven** | `threading.Event`, `_task_event` | Non-blocking wait mechanism for task scheduling using Python's `threading` module |
|
|
||||||
| **AutoModel Registry** | `AutoModel`, `Transformer` | Model type registration and dynamic loading via decorator pattern |
|
|
||||||
| **Generator Pattern** | `GenerateResult`, `GenerationRequest` | Event-based result notification for streaming/non-streaming generation |
|
|
||||||
|
|
||||||
### Core Relationships
|
|
||||||
|
|
||||||
1. **Configuration → Training**: `TrainConfig` holds model, dataset, optimizer_fn, scheduler_fn and other training configuration references
|
|
||||||
2. **Training Flow**: `Trainer` → `TrainContextBuilder` → `TrainContext`, uses `BaseStrategy` to compute loss
|
|
||||||
3. **Strategy Selection**: `StrategyFactory` creates corresponding strategy instance based on `train_type`
|
|
||||||
4. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `Transformer`, uses `KVCache` (backed by `Allocator` + `PrefixCache` + `PagePool` + `Storage`) for paged KV cache management and `SamplingPipeline` for efficient continuous batching with streaming/non-streaming
|
|
||||||
5. **Distributed Support**: `spawn_parallel_fn` and `setup_parallel` provide multi-process training capability for `Trainer`
|
|
||||||
6. **Dataset Loading**: `DatasetFactory` creates datasets (SEQDataset, SFTDataset, DPODataset, GRPODataset), supports HDF5 loading via `BaseSegmentFetcher` and `MultiSegmentFetcher`
|
|
||||||
7. **Checkpoint Management**: `Checkpoint` handles model state serialization/deserialization with safetensors
|
|
||||||
8. **Scheduler Support**: `SchedulerFactory` creates learning rate schedulers (CosineScheduler, SGDRScheduler)
|
|
||||||
9. **AutoModel Loading**: `AutoModel.from_pretrained()` dynamically loads model based on `config.json` model_type, uses `Registry` pattern for model type registration
|
|
||||||
|
|
||||||
## 3. Training Process
|
|
||||||
|
|
||||||
The common training process for large language models (LLM) typically includes three stages: **Pre-training (SEQ)**, **Supervised Fine-Tuning (SFT)**, and **Reinforcement Learning from Human Feedback (DPO/GRPO)**. This system is designed to support seamless end-to-end flow, achieving efficient switching and state management of different training stages through modular strategies.
|
|
||||||
|
|
||||||
### Core Formulas
|
|
||||||
|
|
||||||
**Pre-training (SEQ):**
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{PT}} = - \sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta)
|
|
||||||
$$
|
|
||||||
|
|
||||||
**SFT:**
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{SFT}} = - \sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
|
|
||||||
$$
|
|
||||||
|
|
||||||
**DPO:**
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{DPO}} = -\mathbb{E}_{(x, y_w, y_l) \sim D} \left[ \log \sigma\left( \beta \log \frac{\pi_\theta(y_w \mid x)}{\pi_{\text{ref}}(y_w \mid x)} - \beta \log \frac{\pi_\theta(y_l \mid x)}{\pi_{\text{ref}}(y_l \mid x)} \right) \right]
|
|
||||||
$$
|
|
||||||
|
|
||||||
**GRPO:**
|
|
||||||
|
|
||||||
GRPO (Group Relative Policy Optimization) computes advantages from multiple responses to the same prompt, then optimizes using a PPO-style clipped objective:
|
|
||||||
|
|
||||||
$$
|
|
||||||
\text{Advantage}_i = \frac{r_i - \mu}{\sigma + \epsilon}
|
|
||||||
$$
|
|
||||||
|
|
||||||
Where $r_i$ is the reward for the $i$-th response, $\mu$ and $\sigma$ are the mean and standard deviation of group rewards.
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{GRPO}} = -\mathbb{E} \left[ \min\left( \frac{\pi_\theta(a|s)}{\pi_{\text{ref}}(a|s)} \cdot A, \text{clip}\left(\frac{\pi_\theta(a|s)}{\pi_{\text{ref}}(a|s)}, 1-\epsilon, 1+\epsilon\right) \cdot A \right) \right] + \lambda \cdot D_{KL}
|
|
||||||
$$
|
|
||||||
|
|
||||||
The KL divergence term uses mean squared error approximation:
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{KL} = \lambda \cdot \mathbb{E} \left[ (\log \pi_\theta - \log \pi_{\text{ref}})^2 \right]
|
|
||||||
$$
|
|
||||||
|
|
||||||
The final loss is the sum of both: $L = L_{\text{policy}} + L_{KL}$
|
|
||||||
|
|
||||||
Through the above three-stage progressive training, the model completes its evolution from a general language foundation to a specialized, highly-aligned dialogue intelligence.
|
|
||||||
|
|
||||||
> Document Update Time: 2026-05-14
|
|
||||||
@@ -1,334 +0,0 @@
|
|||||||
## Model Introduction
|
|
||||||
|
|
||||||
### 1. Model Architecture
|
|
||||||
|
|
||||||
This model uses the Transformer architecture with GQA mechanism (q_head=24, kv_head=4), which saves KV cache memory compared to traditional MHA. The model is built by stacking multiple layers of Transformer blocks, with 1.0 billion parameters. Transformer is an autoregressive model that calculates the relationship between all previous tokens to obtain the probability distribution of the next token.
|
|
||||||
|
|
||||||
The model now uses the **AutoModel** base class for flexible loading and saving:
|
|
||||||
|
|
||||||
```python
|
|
||||||
from astrai.model import AutoModel
|
|
||||||
|
|
||||||
# Load model from checkpoint
|
|
||||||
model = AutoModel.from_pretrained("path/to/model")
|
|
||||||
|
|
||||||
# Save model to new directory
|
|
||||||
model.save_pretrained("path/to/save")
|
|
||||||
```
|
|
||||||
|
|
||||||
The Transformer model is registered via `@AutoModel.register('transformer')` decorator, allowing easy extension for new model types.
|
|
||||||
|
|
||||||
```mermaid
|
|
||||||
flowchart TB
|
|
||||||
subgraph Layers["Transformer Layers"]
|
|
||||||
direction TB
|
|
||||||
A[Input Embedding] --> B[Transformer Block\nLayer 1]
|
|
||||||
B --> C[Transformer Block\nLayer ...]
|
|
||||||
C --> D[Transformer Block\nLayer ...]
|
|
||||||
D --> E[RMSNorm]
|
|
||||||
E --> F[Linear]
|
|
||||||
F --> G[SoftMax]
|
|
||||||
end
|
|
||||||
|
|
||||||
subgraph TransformerBlock["Transformer Block"]
|
|
||||||
direction TB
|
|
||||||
H[x] --> I[RMSNorm]
|
|
||||||
I --> J[Linear → Q/K/V]
|
|
||||||
J --> K[Q]
|
|
||||||
J --> L[K]
|
|
||||||
J --> M[V]
|
|
||||||
K --> N[RoPE]
|
|
||||||
L --> O[RoPE]
|
|
||||||
N --> P["Q @ K^T / sqrt(d)"]
|
|
||||||
O --> P
|
|
||||||
P --> Q[Masked SoftMax]
|
|
||||||
Q --> R[S @ V]
|
|
||||||
M --> R
|
|
||||||
R --> S[Linear]
|
|
||||||
S --> T[+]
|
|
||||||
H --> T
|
|
||||||
T --> U[RMSNorm]
|
|
||||||
U --> V["Linear (gate)"]
|
|
||||||
U --> W["Linear (up)"]
|
|
||||||
V --> X[SiLU]
|
|
||||||
X --> Y[×]
|
|
||||||
W --> Y
|
|
||||||
Y --> Z["Linear (down)"]
|
|
||||||
Z --> AA[+]
|
|
||||||
T --> AA
|
|
||||||
AA --> BB[x']
|
|
||||||
end
|
|
||||||
|
|
||||||
classDef main fill:#e6f3ff,stroke:#0066cc;
|
|
||||||
classDef block fill:#fff2e6,stroke:#cc6600;
|
|
||||||
class Layers main;
|
|
||||||
class TransformerBlock block;
|
|
||||||
```
|
|
||||||
|
|
||||||
What is an autoregressive model? After splitting a sentence into tokens, the model predicts the probability distribution of the next token. This means the model calculates the probability of the next possible token and its corresponding probability based on the given context (the sequence of tokens that have already appeared).
|
|
||||||
|
|
||||||
#### 1. Autoregression
|
|
||||||
|
|
||||||
In autoregressive modeling, when a sentence is tokenized into a sequence of tokens, the model learns to predict what comes next. Given a sequence of tokens as input, the model calculates a probability distribution over all possible next tokens. This distribution tells us how likely each potential next token is, given the current context.
|
|
||||||
|
|
||||||
For instance, if the input sequence contains tokens representing a question, the model might predict that certain response tokens have higher probabilities than others. The sampling process then selects one token from this distribution—controlled by parameters like top_k, top_p, and temperature—to serve as the next token in the sequence.
|
|
||||||
|
|
||||||
Once a token is selected, it is appended to the input sequence, and the model repeats this process. The updated sequence is then fed back into the model to predict the next token. This iterative process continues until either a special end-of-sequence token is generated, or the maximum sequence length is reached. These control tokens are essential because without them, the model would continue generating tokens indefinitely, eventually exhausting available memory.
|
|
||||||
|
|
||||||
#### 2. Causal Mask
|
|
||||||
|
|
||||||
Transformers use attention mechanism. The input shape is generally [bsz, seq_len], and the output is [bsz, seq_len, n_dim]. To predict the next token, the model's input and output must be offset by one position. The target predicted by the model must be offset by one position, and during training we also use the offset-by-one method:
|
|
||||||
|
|
||||||
```
|
|
||||||
sequence : [[1, 2, 3, 4, 5, 6]]
|
|
||||||
input_ids: [[1, 2, 3, 4, 5]]
|
|
||||||
target_ids: [[2, 3, 4, 5, 6]]
|
|
||||||
```
|
|
||||||
|
|
||||||
The attention score calculation formula is:
|
|
||||||
|
|
||||||
$$ s_{ij} = softmax(\frac{q_i^Tk_j}{\sqrt{d_k}}) $$
|
|
||||||
$$ s_{ij} := s_{ij} + mask_{ij} $$
|
|
||||||
|
|
||||||
Here, the attention score represents the degree to which the model attends to the similarity between two tokens.
|
|
||||||
|
|
||||||
For decoder-only structure models, to prevent the model from "stealing" information from future positions, a mask needs to be added during attention calculation. We need to apply a mask before attention score calculation. This mask is typically a lower triangular matrix, and for a sequence of length n, its shape is [n, n]. Below is an example of how to create such a causal mask matrix for a sequence of length 5:
|
|
||||||
|
|
||||||
```
|
|
||||||
[[0, -inf, -inf, -inf, -inf],
|
|
||||||
[0, 0, -inf, -inf, -inf],
|
|
||||||
[0, 0, 0, -inf, -inf],
|
|
||||||
[0, 0, 0, 0, -inf],
|
|
||||||
[0, 0, 0, 0, 0]]
|
|
||||||
```
|
|
||||||
|
|
||||||
In this matrix, 0 represents positions that can be attended to, while -inf represents positions that should be masked (i.e., should not be attended to). Because this matrix ensures that after the softmax, the parts of the attention scores where $j > i$ change from `inf` to 0, meaning the model cannot see future information.
|
|
||||||
|
|
||||||
#### 3. Rotary Position Embedding
|
|
||||||
|
|
||||||
Rotary Position Embedding (RoPE) is a position encoding method designed to solve the problem of lacking direct modeling of sequence position information in Transformer models. Unlike traditional position encodings (such as sine and cosine function position encodings), RoPE embeds position information directly into the Query (Q) and Key (K) vectors, allowing the model to more naturally handle relative position relationships in sequences.
|
|
||||||
|
|
||||||
$$ q_i = R_i W_q x_i $$
|
|
||||||
$$ k_j = R_j W_k x_j $$
|
|
||||||
$$ q_i^T k_j = (R_i W_q x_i)^T( R_j W_k x_j) = x_i^T W_q^T R_{i-j} W_k x_j $$
|
|
||||||
|
|
||||||
The $R_{i-j}$ controls the attenuation of attention for different tokens at different relative distances. When the absolute value of $i - j$ is larger, the degree of attenuation is stronger. This approach allows the model to learn relative position relationships, enabling the model to scale and adapt to longer sequences.
|
|
||||||
|
|
||||||
## KV Cache Implementation
|
|
||||||
|
|
||||||
According to the attention calculation formula:
|
|
||||||
|
|
||||||
$$
|
|
||||||
\begin{align*}
|
|
||||||
o_i &= \sum_j s_{ij} v_{j} \newline
|
|
||||||
s_{ij} &= \text{softmax}\left( \frac{q_{i} k_{j}}{\sqrt{d_k}} \right)
|
|
||||||
\end{align*}
|
|
||||||
$$
|
|
||||||
|
|
||||||
Since the model is an autoregressive model, we only need to calculate for the last part of the sequence, meaning the index $i$ is fixed as the last element of the sequence, and we compute $o_{n}$:
|
|
||||||
|
|
||||||
$$
|
|
||||||
\begin{align*}
|
|
||||||
o_n &= \sum_j s_{j}v_{j} \newline
|
|
||||||
s_j &= \text{softmax}\left(\frac{q_n k_{j}}{\sqrt{d_k}} \right)
|
|
||||||
\end{align*}
|
|
||||||
$$
|
|
||||||
|
|
||||||
If we expand the expression:
|
|
||||||
|
|
||||||
$$
|
|
||||||
o_n = \sum_j \text{softmax}\left(\frac{q_n k_{j}}{\sqrt{d_k}}\right)v_{j}
|
|
||||||
$$
|
|
||||||
|
|
||||||
In the above expression, only k and v have length indices, while $q$ does not. Therefore, during the calculation process, the input of $q$ is fixed as the last token from the previous input, while $k$ and $v$ need to be cached for parts of different lengths. Also, when caching, note that position encoding calculation should be performed before KV cache computation, otherwise there will be position encoding calculation errors.
|
|
||||||
|
|
||||||
### 4. AutoModel Loading
|
|
||||||
|
|
||||||
The project now uses the **AutoModel** base class for flexible model loading and saving:
|
|
||||||
|
|
||||||
```python
|
|
||||||
from astrai.model import AutoModel
|
|
||||||
|
|
||||||
# Load model from checkpoint
|
|
||||||
model = AutoModel.from_pretrained("path/to/model")
|
|
||||||
|
|
||||||
# Save model to new directory
|
|
||||||
model.save_pretrained("path/to/save")
|
|
||||||
```
|
|
||||||
|
|
||||||
The Transformer model is registered via `@AutoModel.register('transformer')` decorator, allowing easy extension for new model types. The `from_pretrained` method automatically loads the `config.json` to determine the model type and uses safetensors format for weights.
|
|
||||||
|
|
||||||
### 5. Continuous Batching Inference
|
|
||||||
|
|
||||||
The inference engine supports **continuous batching** for efficient batch processing:
|
|
||||||
|
|
||||||
```python
|
|
||||||
from astrai.inference import InferenceEngine, GenerationRequest
|
|
||||||
|
|
||||||
# Create inference engine with continuous batching
|
|
||||||
engine = InferenceEngine(
|
|
||||||
model=model,
|
|
||||||
tokenizer=tokenizer,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Use GenerationRequest with messages format
|
|
||||||
request = GenerationRequest(
|
|
||||||
messages=[
|
|
||||||
{"role": "system", "content": "You are a helpful assistant."},
|
|
||||||
{"role": "user", "content": "Hello"},
|
|
||||||
],
|
|
||||||
temperature=0.8,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50,
|
|
||||||
max_tokens=None,
|
|
||||||
stream=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Generate with streaming
|
|
||||||
for token in engine.generate_with_request(request):
|
|
||||||
print(token, end="", flush=True)
|
|
||||||
```
|
|
||||||
|
|
||||||
The continuous batching feature allows dynamic batch composition where new requests can join at any time and completed requests are released immediately.
|
|
||||||
|
|
||||||
## HTTP API Usage
|
|
||||||
|
|
||||||
The inference server provides HTTP endpoints for remote inference. Start the server first:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python -m scripts.tools.server --port 8000
|
|
||||||
```
|
|
||||||
|
|
||||||
### OpenAI-Compatible Endpoint
|
|
||||||
|
|
||||||
The server provides an OpenAI-compatible chat completion endpoint at `/v1/chat/completions`:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
|
||||||
-H "Content-Type: application/json" \
|
|
||||||
-d '{
|
|
||||||
"messages": [
|
|
||||||
{"role": "system", "content": "You are a helpful assistant."},
|
|
||||||
{"role": "user", "content": "Hello, how are you?"}
|
|
||||||
],
|
|
||||||
"temperature": 0.8,
|
|
||||||
"max_tokens": 2048,
|
|
||||||
"stream": false
|
|
||||||
}'
|
|
||||||
```
|
|
||||||
|
|
||||||
**Request Parameters:**
|
|
||||||
| Parameter | Type | Default | Description |
|
|
||||||
|-----------|------|---------|-------------|
|
|
||||||
| `messages` | List[dict] | Required | Chat messages with role and content |
|
|
||||||
| `temperature` | float | 1.0 | Sampling temperature (0.0-2.0) |
|
|
||||||
| `top_p` | float | 1.0 | Nucleus sampling threshold |
|
|
||||||
| `top_k` | int | 50 | Top-k sampling parameter |
|
|
||||||
| `max_tokens` | int | 1024 | Maximum tokens to generate |
|
|
||||||
| `stream` | bool | false | Enable streaming response |
|
|
||||||
|
|
||||||
**Response (non-streaming):**
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"id": "chatcmpl-1234567890",
|
|
||||||
"object": "chat.completion",
|
|
||||||
"created": 1234567890,
|
|
||||||
"model": "astrai",
|
|
||||||
"choices": [
|
|
||||||
{
|
|
||||||
"index": 0,
|
|
||||||
"message": {"role": "assistant", "content": "Hello! I'm doing well..."},
|
|
||||||
"finish_reason": "stop"
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"usage": {
|
|
||||||
"prompt_tokens": 20,
|
|
||||||
"completion_tokens": 15,
|
|
||||||
"total_tokens": 35
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
### Streaming Response
|
|
||||||
|
|
||||||
Enable streaming for real-time token-by-token output:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
|
||||||
-H "Content-Type: application/json" \
|
|
||||||
-d '{
|
|
||||||
"messages": [{"role": "user", "content": "Write a story"}],
|
|
||||||
"stream": true,
|
|
||||||
"max_tokens": 500
|
|
||||||
}'
|
|
||||||
```
|
|
||||||
|
|
||||||
The server uses Server-Sent Events (SSE) with content type `text/event-stream`.
|
|
||||||
|
|
||||||
### Anthropic-Compatible Endpoint
|
|
||||||
|
|
||||||
The server also provides an Anthropic-compatible endpoint at `/v1/messages`:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
curl -X POST http://localhost:8000/v1/messages \
|
|
||||||
-H "Content-Type: application/json" \
|
|
||||||
-d '{
|
|
||||||
"model": "astrai",
|
|
||||||
"system": "You are a helpful assistant.",
|
|
||||||
"messages": [{"role": "user", "content": "Hello, how are you?"}],
|
|
||||||
"max_tokens": 2048
|
|
||||||
}'
|
|
||||||
```
|
|
||||||
|
|
||||||
Response:
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"id": "msg_abc123...",
|
|
||||||
"type": "message",
|
|
||||||
"role": "assistant",
|
|
||||||
"model": "astrai",
|
|
||||||
"content": [{"type": "text", "text": "Hello! I am doing well..."}],
|
|
||||||
"stop_reason": "end_turn",
|
|
||||||
"stop_sequence": null,
|
|
||||||
"usage": {"input_tokens": 20, "output_tokens": 15}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
Streaming:
|
|
||||||
```bash
|
|
||||||
curl -X POST http://localhost:8000/v1/messages \
|
|
||||||
-H "Content-Type: application/json" \
|
|
||||||
-d '{
|
|
||||||
"model": "astrai",
|
|
||||||
"system": "You are a helpful assistant.",
|
|
||||||
"messages": [{"role": "user", "content": "Write a short poem"}],
|
|
||||||
"max_tokens": 500,
|
|
||||||
"stream": true
|
|
||||||
}'
|
|
||||||
```
|
|
||||||
|
|
||||||
Supports `stop_sequences` for early termination:
|
|
||||||
```bash
|
|
||||||
curl -X POST http://localhost:8000/v1/messages \
|
|
||||||
-H "Content-Type: application/json" \
|
|
||||||
-d '{
|
|
||||||
"model": "astrai",
|
|
||||||
"messages": [{"role": "user", "content": "Write a story"}],
|
|
||||||
"max_tokens": 500,
|
|
||||||
"stop_sequences": ["The end", "THE END"]
|
|
||||||
}'
|
|
||||||
```
|
|
||||||
|
|
||||||
### Health Check
|
|
||||||
|
|
||||||
Monitor server and model status:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
curl http://localhost:8000/health
|
|
||||||
# {"status": "ok", "model_loaded": true}
|
|
||||||
|
|
||||||
curl http://localhost:8000/stats
|
|
||||||
# {"total_tasks": 10, "total_tokens": 5000, "active_tasks": 1, "waiting_queue": 0}
|
|
||||||
```
|
|
||||||
|
|
||||||
> Document Update Time: 2026-05-14
|
|
||||||
@@ -1,158 +0,0 @@
|
|||||||
# Parameter Documentation
|
|
||||||
|
|
||||||
## Training Parameters
|
|
||||||
|
|
||||||
### Basic Parameters
|
|
||||||
|
|
||||||
| Parameter | Description | Default |
|
|
||||||
|-----------|-------------|---------|
|
|
||||||
| `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`) | required |
|
|
||||||
| `--data_root_path` | Dataset root directory | required |
|
|
||||||
| `--param_path` | Model parameters or checkpoint path | required |
|
|
||||||
| `--n_epoch` | Total training epochs | 1 |
|
|
||||||
| `--batch_size` | Batch size | 1 |
|
|
||||||
| `--accumulation_steps` | Gradient accumulation steps between optimizer steps | 1 |
|
|
||||||
|
|
||||||
### Learning Rate Scheduling
|
|
||||||
|
|
||||||
| Parameter | Description | Default |
|
|
||||||
|-----------|-------------|---------|
|
|
||||||
| `--warmup_steps` | Warmup steps | 1000 |
|
|
||||||
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
|
||||||
| `--max_grad_norm` | Maximum gradient norm for clipping | 1.0 |
|
|
||||||
|
|
||||||
### Optimizer (AdamW)
|
|
||||||
|
|
||||||
| Parameter | Description | Default |
|
|
||||||
|-----------|-------------|---------|
|
|
||||||
| `--adamw_beta1` | AdamW beta1 | 0.9 |
|
|
||||||
| `--adamw_beta2` | AdamW beta2 | 0.95 |
|
|
||||||
| `--adamw_weight_decay` | AdamW weight decay | 0.01 |
|
|
||||||
|
|
||||||
### Data Loading
|
|
||||||
|
|
||||||
| Parameter | Description | Default |
|
|
||||||
|-----------|-------------|---------|
|
|
||||||
| `--window_size` | Max input sequence length | model config `max_len` |
|
|
||||||
| `--stride` | Stride for sliding window over sequences | None |
|
|
||||||
| `--random_seed` | Random seed for reproducibility | 3407 |
|
|
||||||
| `--num_workers` | DataLoader worker processes | 4 |
|
|
||||||
| `--no_pin_memory` | Disable pin_memory (enabled by default) | (flag) |
|
|
||||||
|
|
||||||
### Checkpoint & Resume
|
|
||||||
|
|
||||||
| Parameter | Description | Default |
|
|
||||||
|-----------|-------------|---------|
|
|
||||||
| `--ckpt_interval` | Iterations between checkpoints | 5000 |
|
|
||||||
| `--ckpt_dir` | Checkpoint save directory | checkpoint |
|
|
||||||
| `--start_epoch` | Resume from epoch (0 = from scratch) | 0 |
|
|
||||||
| `--start_batch` | Resume from batch iteration | 0 |
|
|
||||||
|
|
||||||
### Distributed Training
|
|
||||||
|
|
||||||
| Parameter | Description | Default |
|
|
||||||
|-----------|-------------|---------|
|
|
||||||
| `--nprocs` | Number of GPUs / processes | 1 |
|
|
||||||
| `--device_type` | Device type | cuda |
|
|
||||||
|
|
||||||
### Strategy-specific
|
|
||||||
|
|
||||||
| Parameter | Description | Default | Used by |
|
|
||||||
|-----------|-------------|---------|---------|
|
|
||||||
| `--dpo_beta` | DPO beta value | 0.1 | `dpo` |
|
|
||||||
| `--label_smoothing` | Label smoothing for cross-entropy loss | 0.1 | `seq`, `sft` |
|
|
||||||
| `--group_size` | GRPO group size | 4 | `grpo` |
|
|
||||||
| `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo` |
|
|
||||||
| `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo` |
|
|
||||||
| `--grpo_sync_interval` | GRPO ref_model sync interval (steps) | 200 | `grpo` |
|
|
||||||
|
|
||||||
### Usage Example
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python scripts/tools/train.py \
|
|
||||||
--train_type seq \
|
|
||||||
--data_root_path /path/to/dataset \
|
|
||||||
--param_path /path/to/model \
|
|
||||||
--n_epoch 3 \
|
|
||||||
--batch_size 4 \
|
|
||||||
--accumulation_steps 8 \
|
|
||||||
--max_lr 3e-4 \
|
|
||||||
--warmup_steps 2000 \
|
|
||||||
--max_grad_norm 1.0 \
|
|
||||||
--ckpt_interval 5000 \
|
|
||||||
--ckpt_dir ./checkpoints \
|
|
||||||
--num_workers 4 \
|
|
||||||
--nprocs 1 \
|
|
||||||
--device_type cuda
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Generation Parameters
|
|
||||||
|
|
||||||
### GenerationRequest Parameters
|
|
||||||
|
|
||||||
| Parameter | Description | Default Value |
|
|
||||||
|-----------|-------------|---------------|
|
|
||||||
| `messages` | List of message dictionaries (role, content) | required |
|
|
||||||
| `temperature` | Sampling temperature (higher = more random) | 1.0 |
|
|
||||||
| `top_p` | Nucleus sampling threshold | 1.0 |
|
|
||||||
| `top_k` | Top-k sampling count | 50 |
|
|
||||||
| `max_tokens` | Maximum generation length | None (unlimited) |
|
|
||||||
| `stream` | Whether to stream output | False |
|
|
||||||
|
|
||||||
### Usage Example
|
|
||||||
|
|
||||||
```python
|
|
||||||
import torch
|
|
||||||
from astrai.model import AutoModel
|
|
||||||
from astrai.tokenize import AutoTokenizer
|
|
||||||
from astrai.inference import InferenceEngine, GenerationRequest
|
|
||||||
|
|
||||||
# Load model using AutoModel
|
|
||||||
model = AutoModel.from_pretrained("your_model_dir")
|
|
||||||
|
|
||||||
# Load tokenizer
|
|
||||||
tokenizer = AutoTokenizer.from_pretrained("your_model_dir")
|
|
||||||
|
|
||||||
# Create engine with separate model and tokenizer
|
|
||||||
engine = InferenceEngine(
|
|
||||||
model=model,
|
|
||||||
tokenizer=tokenizer,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Build request with messages format
|
|
||||||
request = GenerationRequest(
|
|
||||||
messages=[
|
|
||||||
{"role": "system", "content": "You are a helpful assistant."},
|
|
||||||
{"role": "user", "content": "Hello"},
|
|
||||||
],
|
|
||||||
temperature=0.8,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50,
|
|
||||||
max_tokens=None,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Generate (streaming)
|
|
||||||
for token in engine.generate_with_request(request):
|
|
||||||
print(token, end="", flush=True)
|
|
||||||
|
|
||||||
# Or use simple generate interface
|
|
||||||
result = engine.generate(
|
|
||||||
prompt="Hello",
|
|
||||||
stream=False,
|
|
||||||
max_tokens=1024,
|
|
||||||
temperature=0.8,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50,
|
|
||||||
)
|
|
||||||
```
|
|
||||||
|
|
||||||
### Generation Modes
|
|
||||||
|
|
||||||
| Mode | Description |
|
|
||||||
|------|-------------|
|
|
||||||
| `stream=True` | Streaming output, yields token by token |
|
|
||||||
| `stream=False` | Non-streaming output, returns complete result |
|
|
||||||
|
|
||||||
> Document Update Time: 2026-05-14
|
|
||||||
+109
-15
@@ -1,32 +1,126 @@
|
|||||||
__version__ = "1.3.5"
|
__version__ = "1.3.12"
|
||||||
__author__ = "ViperEkura"
|
__author__ = "ViperEkura"
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
|
||||||
from astrai.config import (
|
from astrai.config import (
|
||||||
ModelConfig,
|
AutoRegressiveLMConfig,
|
||||||
|
BaseModelConfig,
|
||||||
|
ConfigFactory,
|
||||||
|
EncoderConfig,
|
||||||
|
PipelineConfig,
|
||||||
TrainConfig,
|
TrainConfig,
|
||||||
)
|
)
|
||||||
from astrai.dataset import DatasetFactory
|
from astrai.dataset import (
|
||||||
|
BaseDataset,
|
||||||
|
DatasetFactory,
|
||||||
|
RDSampler,
|
||||||
|
Store,
|
||||||
|
StoreFactory,
|
||||||
|
)
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.inference import (
|
from astrai.inference import (
|
||||||
GenerationRequest,
|
GenerationRequest,
|
||||||
InferenceEngine,
|
InferenceEngine,
|
||||||
|
ProtocolHandler,
|
||||||
|
SamplingPipeline,
|
||||||
|
get_app,
|
||||||
|
run_server,
|
||||||
|
sample,
|
||||||
)
|
)
|
||||||
from astrai.model import AutoModel, Transformer
|
from astrai.model import (
|
||||||
from astrai.tokenize import AutoTokenizer
|
AutoModel,
|
||||||
from astrai.trainer import CallbackFactory, SchedulerFactory, StrategyFactory, Trainer
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def setup_logging(level: str = "INFO"):
|
||||||
|
"""Attach a handler to the ``astrai`` logger (only, not root).
|
||||||
|
|
||||||
|
Call once per process, e.g. at the top of CLI scripts.
|
||||||
|
Set ``ASTR_LOG_LEVEL`` to override the default ``INFO``.
|
||||||
|
"""
|
||||||
|
_logger = logging.getLogger("astrai")
|
||||||
|
if _logger.handlers:
|
||||||
|
return
|
||||||
|
_level = getattr(
|
||||||
|
logging, os.environ.get("ASTR_LOG_LEVEL", level).upper(), logging.INFO
|
||||||
|
)
|
||||||
|
_logger.setLevel(_level)
|
||||||
|
_handler = logging.StreamHandler()
|
||||||
|
_handler.setFormatter(
|
||||||
|
logging.Formatter(
|
||||||
|
"%(asctime)s | %(levelname)-7s | %(name)s | %(message)s",
|
||||||
|
datefmt="%Y-%m-%d %H:%M:%S",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
_logger.addHandler(_handler)
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"Transformer",
|
"AutoRegressiveLM",
|
||||||
"ModelConfig",
|
"AutoRegressiveLMConfig",
|
||||||
"TrainConfig",
|
"AutoModel",
|
||||||
"DatasetFactory",
|
|
||||||
"AutoTokenizer",
|
"AutoTokenizer",
|
||||||
|
"BaseDataset",
|
||||||
|
"BaseFactory",
|
||||||
|
"BaseModelConfig",
|
||||||
|
"BaseScheduler",
|
||||||
|
"BaseStrategy",
|
||||||
|
"CallbackFactory",
|
||||||
|
"ChatTemplate",
|
||||||
|
"Checkpoint",
|
||||||
|
"ConfigFactory",
|
||||||
|
"DatasetFactory",
|
||||||
|
"EmbeddingEncoder",
|
||||||
|
"EncoderConfig",
|
||||||
|
"ExecutorFactory",
|
||||||
"GenerationRequest",
|
"GenerationRequest",
|
||||||
"InferenceEngine",
|
"InferenceEngine",
|
||||||
"Trainer",
|
"LoRAConfig",
|
||||||
"CallbackFactory",
|
"Pipeline",
|
||||||
"StrategyFactory",
|
"PipelineConfig",
|
||||||
|
"ProtocolHandler",
|
||||||
|
"RDSampler",
|
||||||
|
"SamplingPipeline",
|
||||||
"SchedulerFactory",
|
"SchedulerFactory",
|
||||||
"BaseFactory",
|
"Store",
|
||||||
"AutoModel",
|
"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",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -1,8 +1,25 @@
|
|||||||
from astrai.config.model_config import ModelConfig
|
from astrai.config.model_config import (
|
||||||
|
AutoRegressiveLMConfig,
|
||||||
|
BaseModelConfig,
|
||||||
|
ConfigFactory,
|
||||||
|
EncoderConfig,
|
||||||
|
)
|
||||||
|
from astrai.config.preprocess_config import (
|
||||||
|
InputConfig,
|
||||||
|
OutputConfig,
|
||||||
|
PipelineConfig,
|
||||||
|
ProcessingConfig,
|
||||||
|
)
|
||||||
from astrai.config.train_config import TrainConfig
|
from astrai.config.train_config import TrainConfig
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
# Model configuration
|
"BaseModelConfig",
|
||||||
"ModelConfig",
|
"AutoRegressiveLMConfig",
|
||||||
|
"EncoderConfig",
|
||||||
|
"ConfigFactory",
|
||||||
"TrainConfig",
|
"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)
|
||||||
+149
-30
@@ -1,42 +1,161 @@
|
|||||||
import json
|
from typing import Any, Dict, Optional
|
||||||
from dataclasses import asdict, dataclass
|
|
||||||
from typing import Optional, Self
|
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
|
@dataclass
|
||||||
class ModelConfig:
|
class BaseModelConfig(BaseConfig):
|
||||||
# basic config
|
"""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
|
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.
|
||||||
|
"""
|
||||||
|
|
||||||
vocab_size: Optional[int] = None
|
vocab_size: Optional[int] = None
|
||||||
dim: Optional[int] = None
|
hidden_size: Optional[int] = None
|
||||||
|
num_hidden_layers: Optional[int] = None
|
||||||
n_layers: Optional[int] = None
|
rms_norm_eps: Optional[float] = None
|
||||||
norm_eps: Optional[float] = None
|
intermediate_size: Optional[int] = None
|
||||||
dim_ffn: Optional[int] = None
|
tie_word_embeddings: Optional[bool] = None
|
||||||
tie_weight: Optional[bool] = None
|
max_position_embeddings: Optional[int] = None
|
||||||
|
|
||||||
# RoPE
|
|
||||||
max_len: Optional[int] = None
|
|
||||||
rope_theta: Optional[float] = None
|
rope_theta: Optional[float] = None
|
||||||
|
rope_scaling: Optional[dict] = None
|
||||||
# GQA
|
attn_type: str = "gqa"
|
||||||
n_heads: Optional[int] = None
|
num_attention_heads: Optional[int] = None
|
||||||
n_kv_heads: Optional[int] = None
|
num_key_value_heads: Optional[int] = None
|
||||||
use_qk_norm: Optional[bool] = None
|
use_qk_norm: Optional[bool] = None
|
||||||
use_gated_attention: 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
|
||||||
|
|
||||||
def load(self, config_path: str) -> Self:
|
@field_validator("attn_type")
|
||||||
config = {}
|
def _validate_attn_type(cls, v: str) -> str:
|
||||||
with open(config_path, "r") as f:
|
if v not in _ATTN_TYPES:
|
||||||
config.update(json.load(f))
|
raise ValueError(
|
||||||
|
f"attn_type must be one of {sorted(_ATTN_TYPES)}, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
for key, value in config.items():
|
@field_validator("ffn_type")
|
||||||
if hasattr(self, key):
|
def _validate_ffn_type(cls, v: str) -> str:
|
||||||
setattr(self, key, value)
|
if v not in _FFN_TYPES:
|
||||||
|
raise ValueError(f"ffn_type must be one of {sorted(_FFN_TYPES)}, got {v!r}")
|
||||||
|
return v
|
||||||
|
|
||||||
return self
|
|
||||||
|
|
||||||
def save(self, config_path: str):
|
@dataclass
|
||||||
config_dict = {k: v for k, v in asdict(self).items() if v is not None}
|
@ConfigFactory.register("embedding")
|
||||||
with open(config_path, "w") as f:
|
class EncoderConfig(BaseModelConfig):
|
||||||
json.dump(config_dict, f, indent=4)
|
"""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)
|
||||||
+199
-83
@@ -1,98 +1,214 @@
|
|||||||
from dataclasses import dataclass, field
|
from dataclasses import field
|
||||||
from typing import Callable, Optional
|
from typing import Any, Callable, Dict, List, Optional
|
||||||
|
|
||||||
import torch.nn as nn
|
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 import Optimizer
|
||||||
from torch.optim.lr_scheduler import LRScheduler
|
from torch.optim.lr_scheduler import LRScheduler
|
||||||
from torch.utils.data import Dataset
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
|
from astrai.config.base import BaseConfig
|
||||||
|
from astrai.model.components.lora import LoRAConfig
|
||||||
|
|
||||||
@dataclass
|
_TRAIN_TYPES = frozenset({"seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"})
|
||||||
class TrainConfig:
|
_PARALLEL_MODES = frozenset({"none", "ddp", "fsdp"})
|
||||||
# basic setting
|
_BACKENDS = frozenset({"nccl", "gloo"})
|
||||||
model: nn.Module = field(default=None, metadata={"help": "Model for training."})
|
_START_METHODS = frozenset({"spawn", "fork", "forkserver"})
|
||||||
strategy: str = field(default=None, metadata={"help": "Training strategy."})
|
_COMPILE_MODES = frozenset({"default", "reduce-overhead", "max-autotune"})
|
||||||
dataset: Dataset = field(default=None, metadata={"help": "Dataset for training."})
|
|
||||||
optimizer_fn: Callable[[nn.Module], Optimizer] = field(
|
|
||||||
default=None, metadata={"help": "Optimizer factory for training."}
|
|
||||||
)
|
|
||||||
scheduler_fn: Callable[[Optimizer], LRScheduler] = field(
|
|
||||||
default=None, metadata={"help": "Scheduler factory for training."}
|
|
||||||
)
|
|
||||||
n_epoch: int = field(default=1, metadata={"help": "Number of epochs for training."})
|
|
||||||
batch_size: int = field(default=4, metadata={"help": "Batch size for training."})
|
|
||||||
accumulation_steps: int = field(
|
|
||||||
default=1, metadata={"help": "Number of iterations between steps."}
|
|
||||||
)
|
|
||||||
max_grad_norm: float = field(
|
|
||||||
default=1.0, metadata={"help": "Maximum gradient norm."}
|
|
||||||
)
|
|
||||||
|
|
||||||
# checkpoint setting
|
|
||||||
start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."})
|
|
||||||
start_batch: int = field(
|
|
||||||
default=0, metadata={"help": "Start batch iteration for training."}
|
|
||||||
)
|
|
||||||
ckpt_dir: str = field(
|
|
||||||
default="./checkpoint", metadata={"help": "Checkpoint directory."}
|
|
||||||
)
|
|
||||||
ckpt_interval: int = field(
|
|
||||||
default=5000, metadata={"help": "Number of iterations between checkpoints."}
|
|
||||||
)
|
|
||||||
|
|
||||||
# dataloader setting
|
@dataclass(config=ConfigDict(arbitrary_types_allowed=True))
|
||||||
random_seed: int = field(default=3407, metadata={"help": "Random seed."})
|
class TrainConfig(BaseConfig):
|
||||||
num_workers: int = field(
|
"""Training configuration.
|
||||||
default=0, metadata={"help": "Number of workers for dataloader."}
|
|
||||||
)
|
|
||||||
prefetch_factor: Optional[int] = field(
|
|
||||||
default=None, metadata={"help": "Prefetch factor for dataloader."}
|
|
||||||
)
|
|
||||||
pin_memory: bool = field(
|
|
||||||
default=False, metadata={"help": "Pin memory for dataloader."}
|
|
||||||
)
|
|
||||||
|
|
||||||
# distributed training
|
Combines hyperparameters with runtime objects (model_fn, dataset, etc.).
|
||||||
nprocs: int = field(
|
Only JSON-serializable fields are written to checkpoint meta via to_dict().
|
||||||
default=1, metadata={"help": "Number of processes for distributed training."}
|
|
||||||
)
|
|
||||||
backend: str = field(
|
|
||||||
default="nccl", metadata={"help": "Distributed training backend."}
|
|
||||||
)
|
|
||||||
master_addr: str = field(
|
|
||||||
default="localhost",
|
|
||||||
metadata={"help": "Master address for distributed training."},
|
|
||||||
)
|
|
||||||
master_port: str = field(
|
|
||||||
default="29500", metadata={"help": "Master port for distributed training."}
|
|
||||||
)
|
|
||||||
parallel_wrapper: Optional[Callable] = field(
|
|
||||||
default=None, metadata={"help": "Parallel function for training."}
|
|
||||||
)
|
|
||||||
state_dict_fn: Optional[Callable] = field(
|
|
||||||
default=None, metadata={"help": "Parallel function for state dict saving."}
|
|
||||||
)
|
|
||||||
|
|
||||||
# others
|
Args:
|
||||||
device_type: str = field(
|
model_fn (Callable[[], nn.Module]): Model factory for training.
|
||||||
default="cuda", metadata={"help": "Device type for distributed training."}
|
strategy (str): Training strategy (seq, sft, dpo, grpo, online_*).
|
||||||
)
|
dataset (Dataset): Dataset for training.
|
||||||
extra_kwargs: dict = field(
|
optimizer_fn (Callable[[nn.Module], Optimizer]): Optimizer factory for training.
|
||||||
default_factory=dict, metadata={"help": "Other arguments."}
|
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.
|
||||||
|
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.
|
||||||
|
rollout_interval (int): Number of optimizer steps between online rollouts. Defaults to 512.
|
||||||
|
rollout_temperature (float): Sampling temperature for online rollout. Defaults to 0.7.
|
||||||
|
rollout_top_k (int): Top-k filtering for online rollout, 0=disable. Defaults to 0.
|
||||||
|
rollout_top_p (float): Top-p (nucleus) filtering for online rollout. Defaults to 0.9.
|
||||||
|
rollout_max_tokens (int): Maximum generated tokens per response in rollout. Defaults to 1024.
|
||||||
|
reward_model_fn (Optional[Callable]): Factory for reward model, required for online RL strategies. Defaults to None.
|
||||||
|
executor_kwargs (Dict[str, Any]): Extra kwargs passed to ExecutorFactory.create(). Defaults to {}.
|
||||||
|
extra_kwargs (Dict[str, Any]): Other arguments. Defaults to {}.
|
||||||
|
"""
|
||||||
|
|
||||||
|
model_fn: Callable[[], nn.Module]
|
||||||
|
strategy: str
|
||||||
|
dataset: Dataset
|
||||||
|
optimizer_fn: Callable[[nn.Module], Optimizer]
|
||||||
|
scheduler_fn: Callable[[Optimizer], LRScheduler]
|
||||||
|
optimizer_name: Optional[str] = None
|
||||||
|
optimizer_hyperparameters: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
n_epoch: int = 1
|
||||||
|
batch_per_device: int = 4
|
||||||
|
grad_accum_steps: int = 1
|
||||||
|
max_grad_norm: Optional[float] = 1.0
|
||||||
|
gradient_checkpointing_modules: List[type] = field(default_factory=list)
|
||||||
|
compile_mode: Optional[str] = None
|
||||||
|
|
||||||
|
start_epoch: int = 0
|
||||||
|
start_samples: int = 0
|
||||||
|
ckpt_dir: str = "./checkpoint"
|
||||||
|
ckpt_interval: int = 5000
|
||||||
|
|
||||||
|
lora: Optional[LoRAConfig] = None
|
||||||
|
|
||||||
|
metrics: List[str] = field(default_factory=lambda: ["loss", "lr", "grad_norm"])
|
||||||
|
|
||||||
|
random_seed: int = 3407
|
||||||
|
num_workers: int = 0
|
||||||
|
prefetch_factor: Optional[int] = None
|
||||||
|
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
|
||||||
|
|
||||||
|
rollout_interval: int = 512
|
||||||
|
rollout_temperature: float = 0.7
|
||||||
|
rollout_top_k: int = 0
|
||||||
|
rollout_top_p: float = 0.9
|
||||||
|
rollout_max_tokens: int = 1024
|
||||||
|
reward_model_fn: Optional[Callable] = None
|
||||||
|
|
||||||
|
executor_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
extra_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
@field_validator("strategy")
|
||||||
|
def _validate_strategy(cls, v: str) -> str:
|
||||||
|
if v not in _TRAIN_TYPES:
|
||||||
|
raise ValueError(
|
||||||
|
f"strategy must be one of {sorted(_TRAIN_TYPES)}, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("parallel_mode")
|
||||||
|
def _validate_parallel_mode(cls, v: str) -> str:
|
||||||
|
if v not in _PARALLEL_MODES:
|
||||||
|
raise ValueError(
|
||||||
|
f"parallel_mode must be one of {sorted(_PARALLEL_MODES)}, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("backend")
|
||||||
|
def _validate_backend(cls, v: str) -> str:
|
||||||
|
if v not in _BACKENDS:
|
||||||
|
raise ValueError(f"backend must be one of {sorted(_BACKENDS)}, got {v!r}")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("start_method")
|
||||||
|
def _validate_start_method(cls, v: str) -> str:
|
||||||
|
if v not in _START_METHODS:
|
||||||
|
raise ValueError(
|
||||||
|
f"start_method must be one of {sorted(_START_METHODS)}, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator("compile_mode")
|
||||||
|
def _validate_compile_mode(cls, v: Optional[str]) -> Optional[str]:
|
||||||
|
if v is not None and v not in _COMPILE_MODES:
|
||||||
|
raise ValueError(
|
||||||
|
f"compile_mode must be one of {sorted(_COMPILE_MODES)} or None, got {v!r}"
|
||||||
|
)
|
||||||
|
return v
|
||||||
|
|
||||||
|
@field_validator(
|
||||||
|
"n_epoch",
|
||||||
|
"batch_per_device",
|
||||||
|
"grad_accum_steps",
|
||||||
|
"ckpt_interval",
|
||||||
|
"val_step",
|
||||||
|
"rollout_interval",
|
||||||
|
"rollout_max_tokens",
|
||||||
)
|
)
|
||||||
|
def _validate_positive_int(cls, v: int) -> int:
|
||||||
|
if v <= 0:
|
||||||
|
raise ValueError(f"must be positive, got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
def __post_init__(self):
|
@field_validator("rollout_temperature")
|
||||||
self.validate()
|
def _validate_positive_float(cls, v: float) -> float:
|
||||||
|
if v <= 0:
|
||||||
|
raise ValueError(f"must be positive, got {v}")
|
||||||
|
return v
|
||||||
|
|
||||||
def validate(self):
|
@field_validator("rollout_top_p")
|
||||||
required_fields = [
|
def _validate_top_p(cls, v: float) -> float:
|
||||||
"model",
|
if not 0 < v <= 1:
|
||||||
"strategy",
|
raise ValueError(f"rollout_top_p must be in (0, 1], got {v}")
|
||||||
"dataset",
|
return v
|
||||||
"optimizer_fn",
|
|
||||||
"scheduler_fn",
|
|
||||||
]
|
|
||||||
|
|
||||||
for field_name in required_fields:
|
@field_validator("rollout_top_k", "num_workers", "neftune_alpha")
|
||||||
if getattr(self, field_name) is None:
|
def _validate_non_negative(cls, v):
|
||||||
raise ValueError(f"{field_name} is required.")
|
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
|
||||||
|
|||||||
+24
-24
@@ -1,37 +1,37 @@
|
|||||||
from astrai.dataset.dataset import (
|
from astrai.dataset.dataset import (
|
||||||
BaseDataset,
|
BaseDataset,
|
||||||
DatasetFactory,
|
DatasetFactory,
|
||||||
|
dpo_collate_fn,
|
||||||
|
grpo_collate_fn,
|
||||||
)
|
)
|
||||||
from astrai.dataset.sampler import ResumableDistributedSampler
|
from astrai.dataset.sampler import RDSampler
|
||||||
from astrai.dataset.storage import (
|
from astrai.dataset.storage import (
|
||||||
BaseSegmentFetcher,
|
JsonlStore,
|
||||||
BaseStorage,
|
MmapStore,
|
||||||
H5Storage,
|
Recordable,
|
||||||
JSONStorage,
|
Store,
|
||||||
MultiSegmentFetcher,
|
StoreFactory,
|
||||||
available_storage_types,
|
Streamable,
|
||||||
create_storage,
|
|
||||||
detect_format,
|
detect_format,
|
||||||
load_h5,
|
)
|
||||||
load_json,
|
from astrai.serialization import (
|
||||||
save_h5,
|
load_bin,
|
||||||
save_json,
|
save_bin,
|
||||||
)
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"BaseDataset",
|
"BaseDataset",
|
||||||
"DatasetFactory",
|
"DatasetFactory",
|
||||||
"BaseSegmentFetcher",
|
"dpo_collate_fn",
|
||||||
"MultiSegmentFetcher",
|
"grpo_collate_fn",
|
||||||
"BaseStorage",
|
"Store",
|
||||||
"H5Storage",
|
"Streamable",
|
||||||
"JSONStorage",
|
"Recordable",
|
||||||
"create_storage",
|
"StoreFactory",
|
||||||
|
"MmapStore",
|
||||||
|
"JsonlStore",
|
||||||
"detect_format",
|
"detect_format",
|
||||||
"available_storage_types",
|
"save_bin",
|
||||||
"save_h5",
|
"load_bin",
|
||||||
"load_h5",
|
"RDSampler",
|
||||||
"save_json",
|
|
||||||
"load_json",
|
|
||||||
"ResumableDistributedSampler",
|
|
||||||
]
|
]
|
||||||
|
|||||||
+420
-192
@@ -1,278 +1,506 @@
|
|||||||
"""Dataset implementations with factory pattern for training."""
|
"""Dataset implementations for training.
|
||||||
|
|
||||||
|
Composition over inheritance — every dataset is a thin wrapper that
|
||||||
|
binds a :class:`Store` to a particular train-type's key mapping. All
|
||||||
|
sample-id → token/record indexing lives on the Store; datasets never
|
||||||
|
know about window/stride math or segment layouts.
|
||||||
|
|
||||||
|
Class hierarchy:
|
||||||
|
|
||||||
|
BaseDataset (ABC) — holds a Store, exposes __len__/keys,
|
||||||
|
overrides __getitem__
|
||||||
|
├── SEQDataset — next-token prediction (stream)
|
||||||
|
├── SFTDataset — loss-mask + position_ids (stream)
|
||||||
|
├── DPODataset — chosen/rejected pairs (record)
|
||||||
|
└── GRPODataset — prompt + response group (record)
|
||||||
|
|
||||||
|
``DatasetFactory.load(train_type, load_path, window_size, stride, …)``
|
||||||
|
builds the Store (auto-detecting format) before constructing the
|
||||||
|
matching dataset. Passing ``store=`` skips Store construction.
|
||||||
|
|
||||||
|
When a record dataset (DPO) reads from raw JSONL, a *processor*
|
||||||
|
function (pure ``record -> Dict[str, Tensor]``) is forwarded to
|
||||||
|
:class:`JsonlStore` so tokenisation happens on the fly.
|
||||||
|
"""
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import Dict, List, Optional
|
from functools import partial
|
||||||
|
from typing import Callable, Dict, List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
from torch.utils.data import Dataset
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
from astrai.dataset.storage import (
|
from astrai.dataset.storage import (
|
||||||
BaseStorage,
|
Store,
|
||||||
create_storage,
|
StoreFactory,
|
||||||
detect_format,
|
detect_format,
|
||||||
)
|
)
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
|
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):
|
class BaseDataset(Dataset, ABC):
|
||||||
"""Abstract base class for all dataset types.
|
"""Abstract base class for dataset types.
|
||||||
|
|
||||||
Implements common functionality for window-based data fetching.
|
Holds a :class:`Store`. All sample-id indexing is delegated to the
|
||||||
Uses a storage abstraction for format-agnostic data loading.
|
store — this class exposes ``__len__`` as ``len(store)`` and the
|
||||||
|
``keys`` property as ``store.keys``. Subclasses implement
|
||||||
|
``__getitem__`` with the train-type-specific key mapping and any
|
||||||
|
training-only index arithmetic (e.g. the next-token ``+1`` shift).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, window_size: int, stride: int):
|
required_keys: List[str] = []
|
||||||
|
|
||||||
|
def __init__(self, store: Store):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.window_size = window_size
|
self.store: Store = store
|
||||||
self.stride = stride
|
validate_keys(store, self.required_keys)
|
||||||
self.storage: Optional[BaseStorage] = None
|
|
||||||
|
|
||||||
def load(self, load_path: str, storage_type: Optional[str] = None, tokenizer=None):
|
def __len__(self) -> int:
|
||||||
"""Load dataset from the given path.
|
return len(self.store)
|
||||||
|
|
||||||
Auto-detects the storage format if not specified.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
load_path: Path to the data directory or file
|
|
||||||
storage_type: Force a specific storage type ("h5", "json"),
|
|
||||||
or None for auto-detection
|
|
||||||
tokenizer: Callable str -> List[int], used to tokenize raw text
|
|
||||||
in JSON files. Ignored for HDF5.
|
|
||||||
"""
|
|
||||||
if storage_type is None:
|
|
||||||
storage_type = detect_format(load_path)
|
|
||||||
self.storage = create_storage(storage_type)
|
|
||||||
self.storage.load(load_path, tokenizer=tokenizer)
|
|
||||||
|
|
||||||
def load_json(self, load_path: str, tokenizer=None):
|
|
||||||
"""Load dataset from JSON files explicitly.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
load_path: Path to the JSON data file or directory
|
|
||||||
tokenizer: Optional tokenizer callable for raw text JSON.
|
|
||||||
"""
|
|
||||||
self.load(load_path, storage_type="json", tokenizer=tokenizer)
|
|
||||||
|
|
||||||
@property
|
|
||||||
def count(self) -> int:
|
|
||||||
"""Return the total number of raw elements (tokens) in the dataset."""
|
|
||||||
if self.storage is None:
|
|
||||||
return 0
|
|
||||||
return len(self.storage)
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def keys(self) -> List[str]:
|
def keys(self) -> List[str]:
|
||||||
"""Return the available data keys."""
|
return self.store.keys
|
||||||
if self.storage is None:
|
|
||||||
return []
|
|
||||||
return self.storage.keys
|
|
||||||
|
|
||||||
def get_index(self, index: int) -> tuple:
|
@property
|
||||||
"""Calculate begin and end indices for a sample.
|
def token_count(self) -> int:
|
||||||
|
return self.store.token_count
|
||||||
Args:
|
|
||||||
index: Sample index
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tuple of (begin_idx, end_idx)
|
|
||||||
"""
|
|
||||||
if self.storage is None:
|
|
||||||
raise RuntimeError("Dataset not loaded, call load() first")
|
|
||||||
total = len(self.storage)
|
|
||||||
if total <= self.window_size:
|
|
||||||
raise ValueError(
|
|
||||||
f"Data too short: {total} tokens <= window_size {self.window_size}"
|
|
||||||
)
|
|
||||||
|
|
||||||
begin_idx = min(index * self.stride, total - 1 - self.window_size)
|
|
||||||
end_idx = min(begin_idx + self.window_size, total - 1)
|
|
||||||
|
|
||||||
return begin_idx, end_idx
|
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
"""Get a single sample by index.
|
|
||||||
|
|
||||||
Must be implemented by subclasses.
|
|
||||||
"""
|
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
def __len__(self) -> int:
|
|
||||||
if self.storage is None:
|
|
||||||
return 0
|
|
||||||
total = len(self.storage)
|
|
||||||
if total <= self.window_size:
|
|
||||||
return 0
|
|
||||||
return (total - 1 - self.window_size) // self.stride + 1
|
|
||||||
|
|
||||||
|
|
||||||
class DatasetFactory(BaseFactory["BaseDataset"]):
|
class DatasetFactory(BaseFactory["BaseDataset"]):
|
||||||
"""Factory class for creating dataset instances.
|
"""Factory for creating dataset instances by train-type.
|
||||||
|
|
||||||
Supports decorator-based registration for extensible dataset types.
|
Use :meth:`DatasetFactory.register("custom")` to register new
|
||||||
All default dataset types (seq, sft, dpo, grpo) are registered automatically
|
dataset classes; they must inherit from :class:`BaseDataset`.
|
||||||
when their classes are defined with the decorator.
|
|
||||||
|
|
||||||
Example usage:
|
|
||||||
@DatasetFactory.register("custom")
|
|
||||||
class CustomDataset(BaseDataset):
|
|
||||||
...
|
|
||||||
|
|
||||||
dataset = DatasetFactory.create("custom", window_size, stride)
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def _validate_component(cls, dataset_cls: type) -> None:
|
|
||||||
"""Validate that the dataset class inherits from BaseDataset."""
|
|
||||||
if not issubclass(dataset_cls, BaseDataset):
|
|
||||||
raise TypeError(f"{dataset_cls.__name__} must inherit from BaseDataset")
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def create(cls, train_type: str, window_size: int, stride: int) -> "BaseDataset":
|
|
||||||
"""Create a dataset instance.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
train_type: Type of training ("seq", "sft", "dpo", "grpo")
|
|
||||||
window_size: Window size for data sampling
|
|
||||||
stride: Stride between consecutive samples
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Dataset instance
|
|
||||||
"""
|
|
||||||
return super().create(train_type, window_size, stride)
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def load(
|
def load(
|
||||||
cls,
|
cls,
|
||||||
train_type: str,
|
train_type: str,
|
||||||
load_path: str,
|
load_path: Optional[str] = None,
|
||||||
window_size: int,
|
window_size: int = 0,
|
||||||
stride: Optional[int] = None,
|
stride: Optional[int] = None,
|
||||||
storage_type: Optional[str] = None,
|
storage_type: Optional[str] = None,
|
||||||
tokenizer=None,
|
tokenizer_path: Optional[str] = None,
|
||||||
|
max_len: int = 2048,
|
||||||
|
store: Optional[Store] = None,
|
||||||
|
**kwargs,
|
||||||
) -> "BaseDataset":
|
) -> "BaseDataset":
|
||||||
"""Create and load a dataset in one step.
|
"""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:
|
Args:
|
||||||
train_type: Type of training dataset
|
train_type: Registered dataset name ("seq", "sft", "dpo",
|
||||||
load_path: Path to the data file
|
"grpo", …).
|
||||||
window_size: Window size for data sampling
|
load_path: Path to the data file or directory (ignored if
|
||||||
stride: Stride between consecutive samples (default: same as window_size)
|
*store* is given).
|
||||||
storage_type: Storage type ("h5", "json") or None for auto-detection
|
window_size: Stream window length — only meaningful for
|
||||||
tokenizer: Callable str -> List[int] for raw text JSON tokenization
|
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:
|
Returns:
|
||||||
Loaded dataset instance
|
Loaded dataset instance.
|
||||||
"""
|
"""
|
||||||
|
if store is not None:
|
||||||
|
return cls.create(train_type, store=store)
|
||||||
|
|
||||||
|
if load_path is None:
|
||||||
|
raise ValueError("Either load_path or store must be provided")
|
||||||
|
|
||||||
|
if storage_type is None:
|
||||||
|
storage_type = detect_format(load_path)
|
||||||
|
|
||||||
if stride is None:
|
if stride is None:
|
||||||
stride = window_size
|
stride = window_size
|
||||||
|
|
||||||
dataset = cls.create(train_type, window_size, stride)
|
processor = cls._maybe_build_processor(
|
||||||
dataset.load(load_path, storage_type=storage_type, tokenizer=tokenizer)
|
train_type, storage_type, tokenizer_path, max_len
|
||||||
|
)
|
||||||
|
|
||||||
return dataset
|
store_window = cls._store_window_for(train_type, window_size)
|
||||||
|
store = StoreFactory.create(
|
||||||
|
storage_type,
|
||||||
|
window_size=store_window,
|
||||||
|
stride=stride if stride else store_window,
|
||||||
|
)
|
||||||
|
if processor is not None:
|
||||||
|
store.load(load_path, processor=processor, **kwargs)
|
||||||
|
else:
|
||||||
|
load_kwargs = dict(kwargs)
|
||||||
|
if (
|
||||||
|
tokenizer_path is not None
|
||||||
|
and storage_type == "jsonl"
|
||||||
|
and train_type in ("seq", "sft")
|
||||||
|
and "tokenizer_path" not in load_kwargs
|
||||||
|
):
|
||||||
|
load_kwargs["tokenizer_path"] = tokenizer_path
|
||||||
|
store.load(load_path, **load_kwargs)
|
||||||
|
|
||||||
@classmethod
|
return cls.create(train_type, store=store)
|
||||||
def available_types(cls) -> list:
|
|
||||||
"""Return list of registered dataset type names."""
|
@staticmethod
|
||||||
return cls.list_registered()
|
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")
|
@DatasetFactory.register("seq")
|
||||||
class SEQDataset(BaseDataset):
|
class SEQDataset(BaseDataset):
|
||||||
"""Dataset for sequential next-token prediction training."""
|
"""Dataset for sequential next-token prediction training.
|
||||||
|
|
||||||
def __init__(self, window_size: int, stride: int):
|
Stream mode: ``store.fetch(begin, end, "sequence")`` returns the
|
||||||
super().__init__(window_size, stride)
|
input window; the +1 shifted call returns the next-token target.
|
||||||
|
"""
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
required_keys = ["sequence"]
|
||||||
return self.storage.fetch(begin_idx, end_idx, "sequence")
|
|
||||||
|
|
||||||
def __getitem__(self, index):
|
def __getitem__(self, index: int):
|
||||||
begin_idx, end_idx = self.get_index(index)
|
begin, end = self.store.sample_window(index)
|
||||||
|
x = self.store.fetch(begin, end, "sequence")
|
||||||
x = self._fetch_data(begin_idx, end_idx).to(dtype=torch.long)
|
y = self.store.fetch(begin + 1, end + 1, "sequence")
|
||||||
y = self._fetch_data(begin_idx + 1, end_idx + 1).to(dtype=torch.long)
|
return {
|
||||||
|
"input_ids": x.to(dtype=torch.long),
|
||||||
return {"input_ids": x, "target_ids": y}
|
"target_ids": y.to(dtype=torch.long),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@DatasetFactory.register("sft")
|
@DatasetFactory.register("sft")
|
||||||
class SFTDataset(BaseDataset):
|
class SFTDataset(BaseDataset):
|
||||||
"""Dataset for supervised fine-tuning with loss masking."""
|
"""Dataset for supervised fine-tuning with loss masking.
|
||||||
|
|
||||||
def __init__(self, window_size: int, stride: int):
|
Stream mode: ``sequence``/``loss_mask``/``position_ids`` are sliced
|
||||||
super().__init__(window_size, stride)
|
to the window. ``loss_mask`` and ``target_ids`` use the +1 shifted
|
||||||
|
slice so they align with the predicted positions.
|
||||||
|
"""
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
required_keys = ["sequence", "loss_mask", "position_ids"]
|
||||||
return self.storage.fetch(begin_idx, end_idx, key)
|
|
||||||
|
|
||||||
def __getitem__(self, index):
|
def __getitem__(self, index: int):
|
||||||
begin_idx, end_idx = self.get_index(index)
|
begin, end = self.store.sample_window(index)
|
||||||
|
x = self.store.fetch(begin, end, "sequence")
|
||||||
x = self._fetch_data(begin_idx, end_idx, "sequence").to(dtype=torch.long)
|
y = self.store.fetch(begin + 1, end + 1, "sequence")
|
||||||
y = self._fetch_data(begin_idx + 1, end_idx + 1, "sequence").to(
|
position_ids = self.store.fetch(begin, end, "position_ids")
|
||||||
dtype=torch.long
|
loss_mask = self.store.fetch(begin + 1, end + 1, "loss_mask")
|
||||||
)
|
return {
|
||||||
loss_mask = self._fetch_data(begin_idx + 1, end_idx + 1, "loss_mask").to(
|
"input_ids": x.to(dtype=torch.long),
|
||||||
dtype=torch.bool
|
"target_ids": y.to(dtype=torch.long),
|
||||||
)
|
"position_ids": position_ids.to(dtype=torch.long),
|
||||||
|
"loss_mask": loss_mask.to(dtype=torch.bool),
|
||||||
return {"input_ids": x, "target_ids": y, "loss_mask": loss_mask}
|
}
|
||||||
|
|
||||||
|
|
||||||
@DatasetFactory.register("dpo")
|
@DatasetFactory.register("dpo")
|
||||||
class DPODataset(BaseDataset):
|
class DPODataset(BaseDataset):
|
||||||
"""Dataset for Direct Preference Optimization training."""
|
"""Record-structured dataset for Direct Preference Optimization.
|
||||||
|
|
||||||
def __init__(self, window_size: int, stride: int):
|
Each sample is one preference pair (chosen + rejected) and is an
|
||||||
super().__init__(window_size, stride)
|
independent training unit — no windowing, stride, or cross-record
|
||||||
|
concatenation. This keeps each sequence self-contained so attention
|
||||||
|
never leaks across preference pairs.
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
Two loading paths (handled by :class:`DatasetFactory`):
|
||||||
return self.storage.fetch(begin_idx, end_idx, key)
|
|
||||||
|
|
||||||
def __getitem__(self, index: int):
|
- **Pre-tokenized** (bin): ``store.load(path)`` reads per-record
|
||||||
begin_idx, end_idx = self.get_index(index)
|
tensors; ``__getitem__`` returns them directly.
|
||||||
|
- **Raw JSONL** (``tokenizer_path=...``): builds a lazy processor
|
||||||
|
via :func:`dpo_processor` that tokenises on the fly — no packing,
|
||||||
|
no ``position_ids``.
|
||||||
|
"""
|
||||||
|
|
||||||
chosen = self._fetch_data(begin_idx, end_idx, "chosen").to(dtype=torch.long)
|
required_keys = ["chosen", "rejected", "chosen_mask", "rejected_mask"]
|
||||||
rejected = self._fetch_data(begin_idx, end_idx, "rejected").to(dtype=torch.long)
|
|
||||||
chosen_mask = self._fetch_data(begin_idx, end_idx, "chosen_mask").to(
|
|
||||||
dtype=torch.bool
|
|
||||||
)
|
|
||||||
rejected_mask = self._fetch_data(begin_idx, end_idx, "rejected_mask").to(
|
|
||||||
dtype=torch.bool
|
|
||||||
)
|
|
||||||
|
|
||||||
|
def make_processor(self, tokenizer, max_len: int):
|
||||||
|
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
|
||||||
|
|
||||||
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
return {
|
return {
|
||||||
"chosen": chosen,
|
"chosen": self.store.fetch_record(index, "chosen").to(dtype=torch.long),
|
||||||
"rejected": rejected,
|
"rejected": self.store.fetch_record(index, "rejected").to(dtype=torch.long),
|
||||||
"chosen_mask": chosen_mask,
|
"chosen_mask": self.store.fetch_record(index, "chosen_mask").to(
|
||||||
"rejected_mask": rejected_mask,
|
dtype=torch.bool
|
||||||
|
),
|
||||||
|
"rejected_mask": self.store.fetch_record(index, "rejected_mask").to(
|
||||||
|
dtype=torch.bool
|
||||||
|
),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@DatasetFactory.register("grpo")
|
@DatasetFactory.register("grpo")
|
||||||
class GRPODataset(BaseDataset):
|
class GRPODataset(BaseDataset):
|
||||||
"""Dataset for Group Relative Policy Optimization training."""
|
"""Dataset for offline Group Relative Policy Optimization.
|
||||||
|
|
||||||
def __init__(self, window_size: int, stride: int):
|
Each sample is one prompt with its group of responses and scalar
|
||||||
super().__init__(window_size, stride)
|
rewards — an independent training unit with no windowing or stride.
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
Expected storage layout (produced by JsonlStore or pre-tokenized):
|
||||||
return self.storage.fetch(begin_idx, end_idx, key)
|
|
||||||
|
- ``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]:
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
begin_idx, end_idx = self.get_index(index)
|
prompts = self.store.fetch_record(index, "prompts")
|
||||||
|
responses = self.store.fetch_record(index, "responses")
|
||||||
prompts = self._fetch_data(begin_idx, end_idx, "prompts")
|
masks = self.store.fetch_record(index, "masks")
|
||||||
responses = self._fetch_data(begin_idx, end_idx, "responses")
|
rewards = self.store.fetch_record(index, "rewards")
|
||||||
masks = self._fetch_data(begin_idx, end_idx, "masks")
|
|
||||||
rewards = self._fetch_data(begin_idx, end_idx, "rewards")
|
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"prompts": prompts,
|
"prompts": prompts.to(dtype=torch.long),
|
||||||
"responses": responses,
|
"responses": [r.to(dtype=torch.long) for r in responses],
|
||||||
"masks": masks,
|
"masks": [m.to(dtype=torch.bool) for m in masks],
|
||||||
"rewards": rewards,
|
"rewards": rewards.to(dtype=torch.float32),
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,7 +5,15 @@ import torch.distributed as dist
|
|||||||
from torch.utils.data import Dataset, Sampler
|
from torch.utils.data import Dataset, Sampler
|
||||||
|
|
||||||
|
|
||||||
class ResumableDistributedSampler(Sampler[int]):
|
class RDSampler(Sampler[int]):
|
||||||
|
"""Resumable Distributed Sampler.
|
||||||
|
|
||||||
|
A distributed sampler that supports checkpoint-based resume: iteration
|
||||||
|
state (epoch, position) is tracked so training can continue from the
|
||||||
|
exact sample after a restart. Shards the dataset across
|
||||||
|
``dist.world_size`` replicas with optional shuffling.
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
data_source: Dataset,
|
data_source: Dataset,
|
||||||
@@ -43,6 +51,7 @@ class ResumableDistributedSampler(Sampler[int]):
|
|||||||
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
|
||||||
|
|
||||||
@@ -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
|
||||||
|
|||||||
+571
-257
@@ -1,105 +1,69 @@
|
|||||||
"""Storage backends for different data formats.
|
"""Storage backends for different data formats.
|
||||||
|
|
||||||
Each storage handles format-specific loading (HDF5, JSON, etc.) and provides
|
Architecture (composition over inheritance):
|
||||||
a uniform interface for data access and length observation via fetchers.
|
|
||||||
|
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 bisect
|
||||||
|
import glob
|
||||||
import json
|
import json
|
||||||
import os
|
import logging
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Callable, Dict, List, Optional, Union
|
from typing import Callable, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
import h5py
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.preprocessing.transform import TokenizeTransform
|
||||||
|
from astrai.serialization import (
|
||||||
|
load_bin,
|
||||||
|
load_bin_offsets,
|
||||||
|
)
|
||||||
|
|
||||||
def save_h5(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
|
logger = logging.getLogger(__name__)
|
||||||
os.makedirs(file_path, exist_ok=True)
|
|
||||||
full_file_path = os.path.join(file_path, f"{file_name}.h5")
|
|
||||||
with h5py.File(full_file_path, "w") as f:
|
|
||||||
for key, tensors in tensor_group.items():
|
|
||||||
grp = f.create_group(key)
|
|
||||||
for idx, tensor in enumerate(tensors):
|
|
||||||
arr = tensor.cpu().numpy()
|
|
||||||
grp.create_dataset(f"data_{idx}", data=arr)
|
|
||||||
|
|
||||||
|
|
||||||
def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
|
|
||||||
tensor_group: Dict[str, List[Tensor]] = {}
|
|
||||||
|
|
||||||
root_path = Path(file_path)
|
|
||||||
h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5"))
|
|
||||||
|
|
||||||
for h5_file in h5_files:
|
|
||||||
with h5py.File(h5_file, "r") as f:
|
|
||||||
for key in f.keys():
|
|
||||||
grp = f[key]
|
|
||||||
dsets = []
|
|
||||||
for dset_name in grp.keys():
|
|
||||||
dset = grp[dset_name]
|
|
||||||
tensor = torch.from_numpy(dset[:])
|
|
||||||
if share_memory:
|
|
||||||
tensor = tensor.share_memory_()
|
|
||||||
dsets.append(tensor)
|
|
||||||
|
|
||||||
if tensor_group.get(key) is None:
|
|
||||||
tensor_group[key] = []
|
|
||||||
tensor_group[key].extend(dsets)
|
|
||||||
|
|
||||||
return tensor_group
|
|
||||||
|
|
||||||
|
|
||||||
def save_json(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
|
|
||||||
os.makedirs(file_path, exist_ok=True)
|
|
||||||
full_file_path = os.path.join(file_path, f"{file_name}.json")
|
|
||||||
json_data = {}
|
|
||||||
for key, tensors in tensor_group.items():
|
|
||||||
json_data[key] = [tensor.tolist() for tensor in tensors]
|
|
||||||
with open(full_file_path, "w", encoding="utf-8") as f:
|
|
||||||
json.dump(json_data, f, ensure_ascii=False)
|
|
||||||
|
|
||||||
|
|
||||||
def load_json(
|
|
||||||
file_path: str,
|
|
||||||
share_memory: bool = True,
|
|
||||||
tokenizer: Optional[Callable[[str], List[int]]] = None,
|
|
||||||
) -> Dict[str, List[Tensor]]:
|
|
||||||
"""Load tensor data from JSON files.
|
|
||||||
|
|
||||||
Supports two modes:
|
|
||||||
- Pre-tokenized: JSON values are List[List[int]] (token IDs), loaded as-is.
|
|
||||||
- Raw text: JSON values are List[str], tokenized via ``tokenizer`` callable
|
|
||||||
at load time. A ``tokenizer`` receives a str and returns List[int].
|
|
||||||
|
|
||||||
Non-data JSON files (e.g. config.json) with scalar/object values are
|
|
||||||
silently skipped.
|
|
||||||
"""
|
|
||||||
tensor_group: Dict[str, List[Tensor]] = {}
|
|
||||||
root_path = Path(file_path)
|
|
||||||
json_files = list(root_path.rglob("*.json")) + list(root_path.rglob("*.jsonl"))
|
|
||||||
for json_file in json_files:
|
|
||||||
with open(json_file, "r", encoding="utf-8") as f:
|
|
||||||
data = json.load(f)
|
|
||||||
if not isinstance(data, dict):
|
|
||||||
continue
|
|
||||||
for key, sequences in data.items():
|
|
||||||
if not isinstance(sequences, list):
|
|
||||||
continue
|
|
||||||
tensors = []
|
|
||||||
for seq in sequences:
|
|
||||||
if tokenizer is not None and isinstance(seq, str):
|
|
||||||
seq = tokenizer(seq)
|
|
||||||
tensor = torch.tensor(seq, dtype=torch.long)
|
|
||||||
if share_memory:
|
|
||||||
tensor = tensor.share_memory_()
|
|
||||||
tensors.append(tensor)
|
|
||||||
if tensor_group.get(key) is None:
|
|
||||||
tensor_group[key] = []
|
|
||||||
tensor_group[key].extend(tensors)
|
|
||||||
return tensor_group
|
|
||||||
|
|
||||||
|
|
||||||
def detect_format(load_path: str) -> str:
|
def detect_format(load_path: str) -> str:
|
||||||
@@ -109,7 +73,7 @@ def detect_format(load_path: str) -> str:
|
|||||||
load_path: Directory or file path
|
load_path: Directory or file path
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Format string ("h5" or "json")
|
Format string ("h5", "bin", "jsonl", or "processed")
|
||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
FileNotFoundError: If no supported data files are found
|
FileNotFoundError: If no supported data files are found
|
||||||
@@ -117,196 +81,546 @@ def detect_format(load_path: str) -> str:
|
|||||||
root = Path(load_path)
|
root = Path(load_path)
|
||||||
if root.is_file():
|
if root.is_file():
|
||||||
suffix = root.suffix.lower()
|
suffix = root.suffix.lower()
|
||||||
if suffix in (".h5", ".hdf5"):
|
if suffix == ".jsonl":
|
||||||
return "h5"
|
return "jsonl"
|
||||||
if suffix in (".json", ".jsonl"):
|
|
||||||
return "json"
|
|
||||||
raise ValueError(f"Unsupported file format: {suffix}")
|
raise ValueError(f"Unsupported file format: {suffix}")
|
||||||
|
|
||||||
h5_files = list(root.rglob("*.h5")) + list(root.rglob("*.hdf5"))
|
bin_files = [Path(p) for p in glob.glob(str(root / "**" / "*.bin"), recursive=True)]
|
||||||
if h5_files:
|
if bin_files:
|
||||||
return "h5"
|
has_meta = (root / "meta.json").exists() or len(
|
||||||
json_files = list(root.rglob("*.json")) + list(root.rglob("*.jsonl"))
|
[Path(p) for p in glob.glob(str(root / "**" / "meta.json"), recursive=True)]
|
||||||
if json_files:
|
) > 0
|
||||||
return "json"
|
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}")
|
raise FileNotFoundError(f"No supported data files found at {load_path}")
|
||||||
|
|
||||||
|
|
||||||
class BaseSegmentFetcher:
|
class Store(ABC):
|
||||||
"""Fetches data segments across multiple tensor segments.
|
"""Common base for all storage backends.
|
||||||
|
|
||||||
Maintains cumulative lengths for efficient range queries across
|
A Store owns both its data layout AND its sample-id → token/record
|
||||||
multiple discontinuous segments.
|
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.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, segments: List[Tensor]):
|
segments_are_records: bool = False
|
||||||
self.segments = segments
|
|
||||||
self.cum_lengths = []
|
|
||||||
|
|
||||||
total = 0
|
def __init__(
|
||||||
for seg in segments:
|
self,
|
||||||
total += torch.numel(seg)
|
window_size: int = 0,
|
||||||
self.cum_lengths.append(total)
|
stride: Optional[int] = None,
|
||||||
|
):
|
||||||
self.total_length = total
|
self._data: Dict[str, List[Tensor]] = {}
|
||||||
|
self._cum: Dict[str, List[int]] = {}
|
||||||
def __len__(self) -> int:
|
self._offsets: Dict[str, List[int]] = {}
|
||||||
return self.total_length
|
self._length: int = 0
|
||||||
|
self._num_records: int = 0
|
||||||
def fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
self._window_size: int = int(window_size)
|
||||||
"""Fetch data in the range [begin_idx, end_idx)."""
|
self._stride: int = int(stride) if stride is not None else int(window_size)
|
||||||
if not (
|
|
||||||
0 <= begin_idx < self.total_length and 0 <= end_idx <= self.total_length
|
|
||||||
):
|
|
||||||
raise ValueError("begin_idx or end_idx out of bounds")
|
|
||||||
if begin_idx >= end_idx:
|
|
||||||
return torch.tensor([], dtype=torch.long)
|
|
||||||
|
|
||||||
seg_start_idx = bisect.bisect_right(self.cum_lengths, begin_idx)
|
|
||||||
seg_end_idx = bisect.bisect_left(self.cum_lengths, end_idx)
|
|
||||||
|
|
||||||
result_segments = []
|
|
||||||
|
|
||||||
for i in range(seg_start_idx, seg_end_idx + 1):
|
|
||||||
prev_cum = self.cum_lengths[i - 1] if i > 0 else 0
|
|
||||||
start = max(begin_idx - prev_cum, 0)
|
|
||||||
end = min(end_idx - prev_cum, len(self.segments[i]))
|
|
||||||
result_segments.append(self.segments[i][start:end])
|
|
||||||
|
|
||||||
return torch.cat(result_segments, dim=0)
|
|
||||||
|
|
||||||
|
|
||||||
class MultiSegmentFetcher:
|
|
||||||
"""Manages multiple segment fetchers for different data keys."""
|
|
||||||
|
|
||||||
def __init__(self, multi_segments: Dict):
|
|
||||||
self.multi_keys = list(multi_segments.keys())
|
|
||||||
self.multi_fetchers = {
|
|
||||||
key: BaseSegmentFetcher(segments)
|
|
||||||
for key, segments in multi_segments.items()
|
|
||||||
}
|
|
||||||
|
|
||||||
def __len__(self) -> int:
|
|
||||||
"""Returns the minimum length across all fetchers."""
|
|
||||||
if not self.multi_fetchers:
|
|
||||||
return 0
|
|
||||||
len_list = [len(seg) for seg in self.multi_fetchers.values()]
|
|
||||||
return min(len_list)
|
|
||||||
|
|
||||||
def key_fetch(
|
|
||||||
self, begin_idx: int, end_idx: int, keys: Union[str, List[str]]
|
|
||||||
) -> Dict:
|
|
||||||
"""Fetch data for specific keys."""
|
|
||||||
fetch_dict = {}
|
|
||||||
keys = [keys] if isinstance(keys, str) else keys
|
|
||||||
|
|
||||||
for key in keys:
|
|
||||||
fetcher = self.multi_fetchers[key]
|
|
||||||
fetch_tensor = fetcher.fetch_data(begin_idx, end_idx)
|
|
||||||
fetch_dict[key] = fetch_tensor
|
|
||||||
|
|
||||||
return fetch_dict if len(keys) > 1 else fetch_dict[keys[0]]
|
|
||||||
|
|
||||||
def fetch_data(self, begin_idx: int, end_idx: int) -> Dict:
|
|
||||||
"""Fetch all keys."""
|
|
||||||
return self.key_fetch(begin_idx, end_idx, self.multi_keys)
|
|
||||||
|
|
||||||
|
|
||||||
class BaseStorage(ABC):
|
|
||||||
"""Abstract storage backend for loading and dispatching data.
|
|
||||||
|
|
||||||
Storage encapsulates format-specific loading and provides a uniform
|
|
||||||
interface for data access and length observation. Subclasses handle
|
|
||||||
different data formats (HDF5, JSON, etc.) while exposing the same
|
|
||||||
fetch interface.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self):
|
|
||||||
self._fetcher: Optional[MultiSegmentFetcher] = None
|
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def load(self, load_path: str, tokenizer=None) -> None:
|
def load(self, path: str, **kwargs) -> None:
|
||||||
"""Load data from the given path into internal fetcher."""
|
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
def __len__(self) -> int:
|
|
||||||
"""Total number of raw elements (tokens) in storage."""
|
|
||||||
if self._fetcher is None:
|
|
||||||
return 0
|
|
||||||
return len(self._fetcher)
|
|
||||||
|
|
||||||
def fetch(self, begin_idx: int, end_idx: int, keys: Union[str, List[str]]):
|
|
||||||
"""Fetch data for the given keys and index range.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
begin_idx: Starting index (inclusive)
|
|
||||||
end_idx: Ending index (exclusive)
|
|
||||||
keys: Single key or list of keys to fetch
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tensor if single key, Dict[str, Tensor] if multiple keys
|
|
||||||
"""
|
|
||||||
if self._fetcher is None:
|
|
||||||
raise RuntimeError("Storage not loaded")
|
|
||||||
return self._fetcher.key_fetch(begin_idx, end_idx, keys)
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def keys(self) -> List[str]:
|
def keys(self) -> List[str]:
|
||||||
"""Return the data keys available in this storage."""
|
return list(self._data.keys())
|
||||||
if self._fetcher is None:
|
|
||||||
return []
|
@property
|
||||||
return self._fetcher.multi_keys
|
def window_size(self) -> int:
|
||||||
|
return self._window_size
|
||||||
|
|
||||||
|
@property
|
||||||
|
def stride(self) -> int:
|
||||||
|
return self._stride
|
||||||
|
|
||||||
|
@property
|
||||||
|
def token_count(self) -> int:
|
||||||
|
"""Total tokens across all stream segments.
|
||||||
|
|
||||||
|
Useful for the bounds-checked raw :meth:`fetch` and as the
|
||||||
|
legacy ``len(store)`` value.
|
||||||
|
"""
|
||||||
|
return self._length
|
||||||
|
|
||||||
|
@property
|
||||||
|
def num_records(self) -> int:
|
||||||
|
"""Number of records available via :meth:`fetch_record`.
|
||||||
|
|
||||||
|
Non-zero only when the backing layout provides per-record
|
||||||
|
indexing (JSONL segments or bin ``_offsets``).
|
||||||
|
"""
|
||||||
|
return self._num_records
|
||||||
|
|
||||||
|
@property
|
||||||
|
def num_samples(self) -> int:
|
||||||
|
"""Number of items produced by ``__getitem__``.
|
||||||
|
|
||||||
|
Stream-mode wins when ``window_size > 0`` and there are tokens
|
||||||
|
to slice; otherwise falls back to ``num_records``.
|
||||||
|
"""
|
||||||
|
if self._window_size > 0 and self._length > 0:
|
||||||
|
total = self._length
|
||||||
|
w = self._window_size
|
||||||
|
if total <= w:
|
||||||
|
return 0
|
||||||
|
return (total - 1 - w) // self._stride + 1
|
||||||
|
return self._num_records
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
return self.num_samples
|
||||||
|
|
||||||
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
|
if index < 0:
|
||||||
|
index += self.num_samples
|
||||||
|
if not 0 <= index < self.num_samples:
|
||||||
|
raise IndexError(
|
||||||
|
f"Store index out of range: {index}, num_samples={self.num_samples}"
|
||||||
|
)
|
||||||
|
if self._window_size > 0 and self._length > 0:
|
||||||
|
begin, end = self.sample_window(index)
|
||||||
|
keys = self._stream_keys()
|
||||||
|
return {k: self.fetch(begin, end, k) for k in keys}
|
||||||
|
return self.fetch_record(index, self._record_keys())
|
||||||
|
|
||||||
|
def sample_window(self, index: int) -> Tuple[int, int]:
|
||||||
|
"""Return ``(begin, end)`` token positions for stream sample *index*.
|
||||||
|
|
||||||
|
The clipped tail keeps the last reachable window inside the
|
||||||
|
token river instead of overshooting. Caller is responsible
|
||||||
|
for staying within :attr:`num_samples`: an out-of-range index
|
||||||
|
raises ``IndexError``.
|
||||||
|
"""
|
||||||
|
if self._window_size <= 0:
|
||||||
|
raise RuntimeError("sample_window() requires window_size > 0 (stream mode)")
|
||||||
|
if self._window_size <= 0 or self._length <= self._window_size:
|
||||||
|
raise IndexError(
|
||||||
|
f"Data too short for window: token_count={self._length}, "
|
||||||
|
f"window_size={self._window_size}"
|
||||||
|
)
|
||||||
|
if not 0 <= index < self.num_samples:
|
||||||
|
raise IndexError(
|
||||||
|
f"Sample index out of range: {index}, num_samples={self.num_samples}"
|
||||||
|
)
|
||||||
|
total = self._length
|
||||||
|
begin = min(index * self._stride, total - 1 - self._window_size)
|
||||||
|
end = min(begin + self._window_size, total - 1)
|
||||||
|
return begin, end
|
||||||
|
|
||||||
|
def _stream_keys(self) -> List[str]:
|
||||||
|
out: List[str] = []
|
||||||
|
for k, tensors in self._data.items():
|
||||||
|
if tensors and isinstance(tensors[0], list):
|
||||||
|
continue
|
||||||
|
out.append(k)
|
||||||
|
return out
|
||||||
|
|
||||||
|
def _record_keys(self) -> List[str]:
|
||||||
|
return list(self._data.keys())
|
||||||
|
|
||||||
|
def _normalize(
|
||||||
|
self,
|
||||||
|
raw: Dict[str, list],
|
||||||
|
offsets: Optional[Dict[str, List[int]]] = None,
|
||||||
|
):
|
||||||
|
"""Register segments and pre-compute indices for both access modes.
|
||||||
|
|
||||||
|
Stream mode: ``_cum[key]`` accumulates per-segment lengths so
|
||||||
|
``Streamable._fetch_stream_key`` can bisect across segments
|
||||||
|
without concatenation.
|
||||||
|
|
||||||
|
Record mode: if *offsets* is provided (bin layout),
|
||||||
|
``_offsets[key]`` stores cumulative per-record offsets into the
|
||||||
|
single concatenated segment. Otherwise, when
|
||||||
|
``segments_are_records`` is True (JSONL), ``_data[key]`` is
|
||||||
|
a per-record list and ``fetch_record`` indexes it directly.
|
||||||
|
|
||||||
|
Nested keys (GRPO ``responses``/``masks`` as
|
||||||
|
``List[List[Tensor]]``) are stored as-is and excluded from both
|
||||||
|
cumulative bookkeepings — they are only accessed record-by-record.
|
||||||
|
"""
|
||||||
|
flat_lengths = []
|
||||||
|
for key, tensors in raw.items():
|
||||||
|
self._data[key] = tensors
|
||||||
|
if not tensors:
|
||||||
|
self._cum[key] = []
|
||||||
|
flat_lengths.append(0)
|
||||||
|
continue
|
||||||
|
if isinstance(tensors[0], list):
|
||||||
|
self._cum[key] = []
|
||||||
|
continue
|
||||||
|
cum = []
|
||||||
|
total = 0
|
||||||
|
for t in tensors:
|
||||||
|
total += t.shape[0]
|
||||||
|
cum.append(total)
|
||||||
|
self._cum[key] = cum
|
||||||
|
flat_lengths.append(cum[-1] if cum else 0)
|
||||||
|
self._length = min(flat_lengths) if flat_lengths else 0
|
||||||
|
|
||||||
|
valid_offsets: Dict[str, List[int]] = {}
|
||||||
|
if offsets:
|
||||||
|
for key, off in offsets.items():
|
||||||
|
segs = self._data.get(key, [])
|
||||||
|
if len(segs) == 1 and len(off) > 1:
|
||||||
|
valid_offsets[key] = off
|
||||||
|
elif len(segs) > 1:
|
||||||
|
logger.warning(
|
||||||
|
"Key '%s' has %d segments with offsets — record mode "
|
||||||
|
"disabled for this key (multi-shard bin+offsets not "
|
||||||
|
"supported). Merge shards or use JSONL.",
|
||||||
|
key,
|
||||||
|
len(segs),
|
||||||
|
)
|
||||||
|
self._offsets = valid_offsets
|
||||||
|
if valid_offsets:
|
||||||
|
record_counts = [len(v) - 1 for v in valid_offsets.values()]
|
||||||
|
self._num_records = min(record_counts) if record_counts else 0
|
||||||
|
elif self.segments_are_records:
|
||||||
|
per_record_counts = []
|
||||||
|
for key, tensors in self._data.items():
|
||||||
|
if tensors and isinstance(tensors[0], list):
|
||||||
|
continue
|
||||||
|
per_record_counts.append(len(tensors))
|
||||||
|
self._num_records = min(per_record_counts) if per_record_counts else 0
|
||||||
|
else:
|
||||||
|
self._num_records = 0
|
||||||
|
|
||||||
|
|
||||||
class H5Storage(BaseStorage):
|
class Streamable:
|
||||||
"""HDF5-based storage backend (pre-tokenized data)."""
|
"""Mixin granting raw token-stream access via :meth:`fetch`.
|
||||||
|
|
||||||
def load(self, load_path: str, tokenizer=None) -> None:
|
Stateless trait relying on ``self._data``, ``self._cum``,
|
||||||
segments = load_h5(load_path)
|
``self._length`` maintained by :class:`Store`. Stream mode is
|
||||||
self._fetcher = MultiSegmentFetcher(segments)
|
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.
|
||||||
class JSONStorage(BaseStorage):
|
|
||||||
"""JSON-based storage backend.
|
|
||||||
|
|
||||||
Supports two modes:
|
|
||||||
- Pre-tokenized: JSON values are List[List[int]], loaded as-is.
|
|
||||||
- Raw text: JSON values are List[str], tokenized via ``tokenizer``
|
|
||||||
callable (str -> List[int]) at load time.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def load(self, load_path: str, tokenizer=None) -> None:
|
def fetch(
|
||||||
segments = load_json(load_path, tokenizer=tokenizer)
|
self,
|
||||||
self._fetcher = MultiSegmentFetcher(segments)
|
begin: int,
|
||||||
|
end: int,
|
||||||
|
keys: Union[str, List[str]],
|
||||||
|
):
|
||||||
|
return _stream_fetch(self, begin, end, keys)
|
||||||
|
|
||||||
|
|
||||||
_STORAGE_REGISTRY: Dict[str, type] = {
|
def _stream_fetch(self, begin: int, end: int, keys: Union[str, List[str]]):
|
||||||
"h5": H5Storage,
|
if not getattr(self, "_data", None):
|
||||||
"json": JSONStorage,
|
raise RuntimeError("Store not loaded")
|
||||||
}
|
if not (0 <= begin < self._length and 0 <= end <= self._length):
|
||||||
|
|
||||||
|
|
||||||
def create_storage(storage_type: str) -> BaseStorage:
|
|
||||||
"""Create a storage instance by type name.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
storage_type: Storage type name ("h5", "json")
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Storage instance
|
|
||||||
|
|
||||||
Raises:
|
|
||||||
ValueError: If the storage type is unknown
|
|
||||||
"""
|
|
||||||
storage_cls = _STORAGE_REGISTRY.get(storage_type)
|
|
||||||
if storage_cls is None:
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Unknown storage type: '{storage_type}'. "
|
f"Index out of bounds: begin={begin}, end={end}, length={self._length}"
|
||||||
f"Available: {sorted(_STORAGE_REGISTRY.keys())}"
|
|
||||||
)
|
)
|
||||||
return storage_cls()
|
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 available_storage_types() -> List[str]:
|
def _fetch_stream_key(self, key: str, begin: int, end: int) -> Tensor:
|
||||||
"""Return list of registered storage type names."""
|
segments = self._data[key]
|
||||||
return sorted(_STORAGE_REGISTRY.keys())
|
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.
|
||||||
|
"""
|
||||||
|
|
||||||
|
CONFIG_NAME = "dataset_config.json"
|
||||||
|
segments_are_records = True
|
||||||
|
|
||||||
|
_DEFAULT_MESSAGES_CONFIG = {
|
||||||
|
"version": 1,
|
||||||
|
"input": {
|
||||||
|
"sections": [{"field": "messages", "action": "$role", "template": True}]
|
||||||
|
},
|
||||||
|
"mask": {"system": "mask", "user": "mask", "assistant": "train"},
|
||||||
|
"mask_default": "mask",
|
||||||
|
"output": {"position_ids_mode": "doc_reset"},
|
||||||
|
}
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
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:
|
||||||
|
root = Path(path)
|
||||||
|
config_path = root / self.CONFIG_NAME if root.is_dir() else None
|
||||||
|
if config_path is not None and config_path.exists():
|
||||||
|
transform = TokenizeTransform.from_config_file(str(config_path))
|
||||||
|
else:
|
||||||
|
tokenizer_path = kwargs.get("tokenizer_path")
|
||||||
|
if not tokenizer_path:
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"JSONL dataset config not found. Expected "
|
||||||
|
f"{self.CONFIG_NAME} alongside *.jsonl files, pass an "
|
||||||
|
f"explicit transform, pass processor= for lazy "
|
||||||
|
f"on-the-fly tokenisation, or pass tokenizer_path= to "
|
||||||
|
f"use the built-in messages config."
|
||||||
|
)
|
||||||
|
config = PipelineConfig.from_dict(self._DEFAULT_MESSAGES_CONFIG)
|
||||||
|
transform = TokenizeTransform(config, tokenizer_path)
|
||||||
|
|
||||||
|
transformed = transform.apply(records)
|
||||||
|
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,51 @@
|
|||||||
|
"""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.attention_backend import (
|
||||||
|
ATTN_BACKEND,
|
||||||
|
AttentionBackend,
|
||||||
|
CudaBackend,
|
||||||
|
TorchNativeBackend,
|
||||||
|
attention,
|
||||||
|
attn_backend,
|
||||||
|
get_backend,
|
||||||
|
)
|
||||||
|
from astrai.extension.attention_ops import (
|
||||||
|
TensorLayout,
|
||||||
|
attn_decode,
|
||||||
|
attn_paged_decode,
|
||||||
|
attn_prefill,
|
||||||
|
)
|
||||||
|
from astrai.extension.loader import KERNEL_NAMES, is_available
|
||||||
|
from astrai.extension.rotary_backend import apply_rotary_emb
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"ATTN_BACKEND",
|
||||||
|
"AttentionBackend",
|
||||||
|
"CudaBackend",
|
||||||
|
"TorchNativeBackend",
|
||||||
|
"TensorLayout",
|
||||||
|
"attention",
|
||||||
|
"attn_backend",
|
||||||
|
"get_backend",
|
||||||
|
"attn_decode",
|
||||||
|
"attn_paged_decode",
|
||||||
|
"attn_prefill",
|
||||||
|
"is_available",
|
||||||
|
"KERNEL_NAMES",
|
||||||
|
"apply_rotary_emb",
|
||||||
|
]
|
||||||
@@ -0,0 +1,401 @@
|
|||||||
|
"""Attention backend abstraction with context-manager switching.
|
||||||
|
|
||||||
|
The backend encapsulates KV cache I/O and attention computation. The
|
||||||
|
attention module (GQA/MLA) keeps projections, rotary, QK-norm, gating,
|
||||||
|
and output projection; the backend handles everything from "write K/V
|
||||||
|
to cache" through "SDPA output".
|
||||||
|
|
||||||
|
Usage — mirroring ``torch.nn.attention.sdpa_kernel``:
|
||||||
|
|
||||||
|
from astrai.extension import attn_backend, ATTN_BACKEND
|
||||||
|
|
||||||
|
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
|
||||||
|
engine.generate("hello")
|
||||||
|
|
||||||
|
# or with an instance:
|
||||||
|
with attn_backend(TorchNativeBackend()):
|
||||||
|
...
|
||||||
|
|
||||||
|
# or the shorthand (instance is itself a context manager):
|
||||||
|
with TorchNativeBackend():
|
||||||
|
...
|
||||||
|
|
||||||
|
Thread-safe via ``contextvars`` — each scheduler thread gets its own
|
||||||
|
active backend. ``get_backend()`` returns the active one, falling back
|
||||||
|
to a process-wide ``TorchNativeBackend`` singleton.
|
||||||
|
|
||||||
|
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
||||||
|
(blhd). The backend returns ``[batch, seq_len, n_heads * head_dim]``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import contextvars
|
||||||
|
import enum
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from typing import Optional, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.extension.attention_ops import (
|
||||||
|
attn_paged_decode,
|
||||||
|
attn_paged_prefill,
|
||||||
|
)
|
||||||
|
from astrai.inference.core.cache import KVCache
|
||||||
|
|
||||||
|
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
|
||||||
|
"attn_backend"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ATTN_BACKEND(enum.Enum):
|
||||||
|
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
|
||||||
|
|
||||||
|
TORCH_NATIVE = "torch_native"
|
||||||
|
CUDA = "cuda"
|
||||||
|
|
||||||
|
|
||||||
|
def get_backend() -> "AttentionBackend":
|
||||||
|
"""Return the active backend for the current thread/context.
|
||||||
|
|
||||||
|
Falls back to a ``TorchNativeBackend`` singleton when no backend
|
||||||
|
has been activated via ``with``.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
return _current_backend.get()
|
||||||
|
except LookupError:
|
||||||
|
return _default_backend
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def attn_backend(backend: Union[ATTN_BACKEND, "AttentionBackend", type]):
|
||||||
|
"""Context manager to select an attention backend.
|
||||||
|
|
||||||
|
Mirrors ``torch.nn.attention.sdpa_kernel``. Accepts an
|
||||||
|
``ATTN_BACKEND`` enum value, a backend class, or a backend instance.
|
||||||
|
|
||||||
|
Examples::
|
||||||
|
|
||||||
|
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
|
||||||
|
...
|
||||||
|
with attn_backend(TorchNativeBackend):
|
||||||
|
...
|
||||||
|
with attn_backend(TorchNativeBackend()):
|
||||||
|
...
|
||||||
|
"""
|
||||||
|
if isinstance(backend, ATTN_BACKEND):
|
||||||
|
instance = _BACKEND_REGISTRY[backend]()
|
||||||
|
elif isinstance(backend, type) and issubclass(backend, AttentionBackend):
|
||||||
|
instance = backend()
|
||||||
|
elif isinstance(backend, AttentionBackend):
|
||||||
|
instance = backend
|
||||||
|
else:
|
||||||
|
raise TypeError(
|
||||||
|
f"expected ATTN_BACKEND, AttentionBackend type, or instance, "
|
||||||
|
f"got {type(backend).__name__}"
|
||||||
|
)
|
||||||
|
token = _current_backend.set(instance)
|
||||||
|
try:
|
||||||
|
yield instance
|
||||||
|
finally:
|
||||||
|
_current_backend.reset(token)
|
||||||
|
|
||||||
|
|
||||||
|
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
||||||
|
"""Expand KV heads to match Q heads for GQA."""
|
||||||
|
bs, slen, n_heads, head_dim = x.shape
|
||||||
|
if n_rep == 1:
|
||||||
|
return x
|
||||||
|
return (
|
||||||
|
x[:, :, :, None, :]
|
||||||
|
.expand(bs, slen, n_heads, n_rep, head_dim)
|
||||||
|
.reshape(bs, slen, n_heads * n_rep, head_dim)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def attention(
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache] = None,
|
||||||
|
layer_id: int = 0,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Functional attention entry point — mirrors ``F.scaled_dot_product_attention``.
|
||||||
|
|
||||||
|
Delegates to the active backend (set via ``with attn_backend(...)``).
|
||||||
|
Handles KV cache I/O, GQA head expansion, and causal masking so the
|
||||||
|
caller only needs to provide projected q/k/v.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [batch, q_len, n_heads, head_dim] (blhd)
|
||||||
|
k: [batch, q_len, n_kv_heads, head_dim] (blhd)
|
||||||
|
v: [batch, q_len, n_kv_heads, head_dim] (blhd)
|
||||||
|
kv_cache: cache dataclass, or None for training (no cache).
|
||||||
|
layer_id: transformer layer index for buffer access.
|
||||||
|
attn_mask: pre-built attention mask (SDPA-compatible).
|
||||||
|
is_causal: whether to apply causal masking.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, q_len, n_heads * head_dim]
|
||||||
|
"""
|
||||||
|
backend = get_backend()
|
||||||
|
return backend.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
|
|
||||||
|
|
||||||
|
class AttentionBackend(ABC):
|
||||||
|
"""Abstract base for attention computation strategies.
|
||||||
|
|
||||||
|
Subclasses implement ``fwd_decode`` (q_len == 1, with cache) and
|
||||||
|
``fwd_prefill`` (q_len > 1, with or without cache). The public
|
||||||
|
``forward`` method dispatches based on q_len.
|
||||||
|
|
||||||
|
Three equivalent ways to activate a backend::
|
||||||
|
|
||||||
|
with attn_backend(ATTN_BACKEND.TORCH_NATIVE): # enum
|
||||||
|
...
|
||||||
|
with attn_backend(TorchNativeBackend): # class
|
||||||
|
...
|
||||||
|
with TorchNativeBackend(): # instance
|
||||||
|
...
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __enter__(self) -> "AttentionBackend":
|
||||||
|
self._token = _current_backend.set(self)
|
||||||
|
return self
|
||||||
|
|
||||||
|
def __exit__(self, *exc) -> None:
|
||||||
|
_current_backend.reset(self._token)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Dispatch to decode or extend based on q_len.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [batch, q_len, n_heads, head_dim]
|
||||||
|
k: [batch, q_len, n_kv_heads, head_dim]
|
||||||
|
v: [batch, q_len, n_kv_heads, head_dim]
|
||||||
|
kv_cache: cache dataclass, or None for training (no cache).
|
||||||
|
layer_id: transformer layer index for buffer access.
|
||||||
|
attn_mask: pre-built attention mask compatible with SDPA.
|
||||||
|
is_causal: whether to apply causal masking.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, q_len, n_heads * head_dim]
|
||||||
|
"""
|
||||||
|
if kv_cache is not None and q.size(1) == 1:
|
||||||
|
return self.fwd_decode(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
|
return self.fwd_prefill(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def fwd_decode(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Single-token decode with KV cache."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def fwd_prefill(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Multi-token prefill or training forward."""
|
||||||
|
|
||||||
|
|
||||||
|
class TorchNativeBackend(AttentionBackend):
|
||||||
|
"""Reference backend using torch SDPA with indirect KV cache indexing.
|
||||||
|
|
||||||
|
Writes new K/V into the cache buffers, gathers the full sequence K/V
|
||||||
|
via ``req_to_token`` indirect indexing, then calls
|
||||||
|
``F.scaled_dot_product_attention``.
|
||||||
|
|
||||||
|
For training (``kv_cache is None``), skips cache I/O entirely and
|
||||||
|
runs SDPA directly on the projected q/k/v.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def fwd_decode(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
|
|
||||||
|
def fwd_prefill(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
|
|
||||||
|
def _forward(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
if kv_cache is not None:
|
||||||
|
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||||
|
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||||
|
|
||||||
|
max_len = kv_cache.max_len
|
||||||
|
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
||||||
|
# Zero out padding positions so gather never touches invalid slots.
|
||||||
|
# Decode: attn_mask[:,0,0] is exactly the per-position validity
|
||||||
|
# mask ([B, max_len], True=keep). Prefill: fall back to seq_lens.
|
||||||
|
if q.size(1) == 1 and attn_mask is not None and attn_mask.dim() == 4:
|
||||||
|
pos_mask = attn_mask[:, 0, 0]
|
||||||
|
else:
|
||||||
|
pos_mask = (
|
||||||
|
torch.arange(max_len, device=q.device)[None, :]
|
||||||
|
< kv_cache.seq_lens[:, None]
|
||||||
|
)
|
||||||
|
indices = torch.where(pos_mask, indices, torch.zeros_like(indices))
|
||||||
|
k = kv_cache.k_buffer[layer_id, indices]
|
||||||
|
v = kv_cache.v_buffer[layer_id, indices]
|
||||||
|
|
||||||
|
n_rep = q.size(2) // k.size(2)
|
||||||
|
if n_rep > 1:
|
||||||
|
k = repeat_kv(k, n_rep)
|
||||||
|
v = repeat_kv(v, n_rep)
|
||||||
|
|
||||||
|
q = q.permute(0, 2, 1, 3)
|
||||||
|
k = k.permute(0, 2, 1, 3)
|
||||||
|
v = v.permute(0, 2, 1, 3)
|
||||||
|
|
||||||
|
out = F.scaled_dot_product_attention(q, k, v, attn_mask, is_causal=is_causal)
|
||||||
|
out = out.permute(0, 2, 1, 3).contiguous().flatten(2)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
_default_backend = TorchNativeBackend()
|
||||||
|
|
||||||
|
|
||||||
|
class CudaBackend(AttentionBackend):
|
||||||
|
"""CUDA kernel backend with direct KV cache access.
|
||||||
|
|
||||||
|
Decode path: writes K/V to 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) is not handled — use
|
||||||
|
``TorchNativeBackend`` for training.
|
||||||
|
|
||||||
|
Raises ``RuntimeError`` if the required kernel is not available.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def fwd_decode(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
if kv_cache is None:
|
||||||
|
raise RuntimeError("CudaBackend does not support training (kv_cache=None)")
|
||||||
|
|
||||||
|
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||||
|
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||||
|
|
||||||
|
b = q.size(0)
|
||||||
|
q_3d = q.squeeze(1)
|
||||||
|
|
||||||
|
kv_indptr = kv_cache.kv_indptr
|
||||||
|
|
||||||
|
out = attn_paged_decode(
|
||||||
|
q_3d,
|
||||||
|
kv_cache.k_buffer[layer_id],
|
||||||
|
kv_cache.v_buffer[layer_id],
|
||||||
|
kv_cache.req_to_token,
|
||||||
|
kv_cache.req_pool_indices,
|
||||||
|
kv_indptr,
|
||||||
|
kv_cache.max_len,
|
||||||
|
mask=attn_mask,
|
||||||
|
is_causal=is_causal,
|
||||||
|
)
|
||||||
|
return out.unsqueeze(1).flatten(2)
|
||||||
|
|
||||||
|
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)")
|
||||||
|
|
||||||
|
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||||
|
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||||
|
|
||||||
|
b = q.size(0)
|
||||||
|
q_len = q.size(1)
|
||||||
|
|
||||||
|
kv_indptr = kv_cache.kv_indptr
|
||||||
|
qo_indptr = torch.arange(b + 1, dtype=torch.int32, device=q.device) * q_len
|
||||||
|
|
||||||
|
q_flat = q.reshape(b * q_len, q.size(2), q.size(3))
|
||||||
|
|
||||||
|
out = attn_paged_prefill(
|
||||||
|
q_flat,
|
||||||
|
kv_cache.k_buffer[layer_id],
|
||||||
|
kv_cache.v_buffer[layer_id],
|
||||||
|
kv_cache.req_to_token,
|
||||||
|
kv_cache.req_pool_indices,
|
||||||
|
kv_indptr,
|
||||||
|
qo_indptr,
|
||||||
|
attn_mask,
|
||||||
|
q_len,
|
||||||
|
is_causal=is_causal,
|
||||||
|
)
|
||||||
|
return out.reshape(b, q_len, q.size(2), q.size(3)).flatten(2)
|
||||||
|
|
||||||
|
|
||||||
|
_BACKEND_REGISTRY: dict[ATTN_BACKEND, type[AttentionBackend]] = {
|
||||||
|
ATTN_BACKEND.TORCH_NATIVE: TorchNativeBackend,
|
||||||
|
ATTN_BACKEND.CUDA: CudaBackend,
|
||||||
|
}
|
||||||
@@ -0,0 +1,185 @@
|
|||||||
|
"""Attention kernel wrapper functions — one entry point per compiled kernel.
|
||||||
|
|
||||||
|
Each wrapper calls its CUDA kernel directly. If the kernel is not
|
||||||
|
available, raises ``RuntimeError``. Fallback to torch SDPA is the
|
||||||
|
responsibility of the attention backend, not this module.
|
||||||
|
|
||||||
|
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
||||||
|
(blhd). Scale is always ``1/sqrt(head_dim)``.
|
||||||
|
|
||||||
|
Interface (all functions):
|
||||||
|
is_causal: True = causal mask; False = non-causal
|
||||||
|
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
|
||||||
|
"""
|
||||||
|
|
||||||
|
import enum
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.extension.loader import _available, _modules
|
||||||
|
|
||||||
|
|
||||||
|
class TensorLayout(enum.IntEnum):
|
||||||
|
"""Q/K/V tensor layout, mirrors the C++ ``TensorLayout`` enum in ``attn_common.h``.
|
||||||
|
|
||||||
|
Kernels internally operate on BHLD; BLHD inputs are transposed at entry.
|
||||||
|
"""
|
||||||
|
|
||||||
|
BHLD = 0 # [batch, n_heads, seq_len, head_dim]
|
||||||
|
BLHD = 1 # [batch, seq_len, n_heads, head_dim]
|
||||||
|
|
||||||
|
|
||||||
|
def _check_available(name: str):
|
||||||
|
if not _available.get(name):
|
||||||
|
raise RuntimeError(
|
||||||
|
f"CUDA kernel '{name}' is not available. "
|
||||||
|
f"Build with CSRC_KERNELS=true or use a torch-native backend."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def attn_decode(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
mask: Optional[torch.Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""GQA decode attention (q_len == 1).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [batch, 1, n_heads, head_dim] (blhd, bf16)
|
||||||
|
k: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||||
|
v: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||||
|
mask: 2D [batch, kv_len] or 3D [batch, 1, kv_len] (bool, True=keep)
|
||||||
|
is_causal: apply causal mask
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, 1, n_heads, head_dim] (blhd, bf16)
|
||||||
|
"""
|
||||||
|
_check_available("attn_decode")
|
||||||
|
causal_offset = (k.size(1) - 1) if is_causal else -1
|
||||||
|
return _modules["attn_decode"].attn_decode(
|
||||||
|
q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def attn_prefill(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
mask: Optional[torch.Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""GQA prefill attention (q_len > 1).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
q: [batch, q_len, n_heads, head_dim] (blhd, bf16)
|
||||||
|
k: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||||
|
v: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
|
||||||
|
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
|
||||||
|
is_causal: apply causal mask
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, q_len, n_heads, head_dim] (blhd, bf16)
|
||||||
|
"""
|
||||||
|
_check_available("attn_prefill")
|
||||||
|
causal_offset = (k.size(1) - q.size(1)) if is_causal else -1
|
||||||
|
return _modules["attn_prefill"].attn_prefill(
|
||||||
|
q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def attn_paged_decode(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k_cache: torch.Tensor,
|
||||||
|
v_cache: torch.Tensor,
|
||||||
|
req_to_token: torch.Tensor,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
kv_indptr: torch.Tensor,
|
||||||
|
max_seq_len: int,
|
||||||
|
mask: Optional[torch.Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> 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] (int64) — token -> slot
|
||||||
|
req_pool_indices: [batch] (int64) — rows into req_to_token
|
||||||
|
kv_indptr: [batch+1] (int32) — prefix sum of per-request seq_lens
|
||||||
|
max_seq_len: max per-request seq_len (Python int, for split computation)
|
||||||
|
mask: 2D [batch, max_seq_len] (bool, True=keep) or None
|
||||||
|
is_causal: apply causal mask
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, n_heads, head_dim] (bf16, 3D)
|
||||||
|
"""
|
||||||
|
_check_available("attn_paged_decode")
|
||||||
|
causal_offset = 0 if is_causal else -1
|
||||||
|
return _modules["attn_paged_decode"].attn_paged_decode(
|
||||||
|
q,
|
||||||
|
k_cache,
|
||||||
|
v_cache,
|
||||||
|
req_to_token,
|
||||||
|
req_pool_indices,
|
||||||
|
kv_indptr,
|
||||||
|
max_seq_len,
|
||||||
|
mask=mask,
|
||||||
|
causal_offset=causal_offset,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
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,
|
||||||
|
mask: Optional[torch.Tensor] = None,
|
||||||
|
max_q_len: int = 0,
|
||||||
|
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] (int64)
|
||||||
|
req_pool_indices: [batch] (int64)
|
||||||
|
kv_indptr: [batch+1] (int32) — prefix sum of per-request kv_lens
|
||||||
|
qo_indptr: [batch+1] (int32) — prefix sum of per-request q_lens
|
||||||
|
mask: 4D [batch, 1, q_len, kv_len] (bool, True=keep) or None
|
||||||
|
max_q_len: max per-request q_len (Python int, for grid computation)
|
||||||
|
is_causal: apply causal mask
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[total_q, n_heads, head_dim] (bf16, 3D)
|
||||||
|
"""
|
||||||
|
_check_available("attn_paged_prefill")
|
||||||
|
causal_offset = 0 if is_causal else -1
|
||||||
|
return _modules["attn_paged_prefill"].attn_paged_prefill(
|
||||||
|
q,
|
||||||
|
k_cache,
|
||||||
|
v_cache,
|
||||||
|
req_to_token,
|
||||||
|
req_pool_indices,
|
||||||
|
kv_indptr,
|
||||||
|
qo_indptr,
|
||||||
|
mask,
|
||||||
|
max_q_len,
|
||||||
|
causal_offset=causal_offset,
|
||||||
|
)
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
"""Compiled CUDA kernel modules (``*.so``) live here, kept separate from Python source."""
|
||||||
@@ -0,0 +1,42 @@
|
|||||||
|
"""Dynamic discovery and loading of compiled CUDA kernel modules.
|
||||||
|
|
||||||
|
Each kernel is registered in ``csrc/build.py`` and built into a ``.so`` placed
|
||||||
|
in this package directory. On import we try to load each one; kernels that
|
||||||
|
failed to build (or are running on a CPU-only machine) are marked unavailable
|
||||||
|
so the wrapper functions can fall back to ``torch`` SDPA.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import importlib
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
KERNEL_NAMES = [
|
||||||
|
"attn_decode",
|
||||||
|
"attn_prefill",
|
||||||
|
"attn_paged_decode",
|
||||||
|
"attn_paged_prefill",
|
||||||
|
"rotary_emb",
|
||||||
|
]
|
||||||
|
|
||||||
|
_available: dict[str, bool] = {}
|
||||||
|
_modules: dict[str, object] = {}
|
||||||
|
|
||||||
|
for _name in KERNEL_NAMES:
|
||||||
|
try:
|
||||||
|
_mod = importlib.import_module(f".lib.{_name}", package=__package__)
|
||||||
|
_available[_name] = True
|
||||||
|
_modules[_name] = _mod
|
||||||
|
except ImportError:
|
||||||
|
_available[_name] = False
|
||||||
|
_modules[_name] = None
|
||||||
|
|
||||||
|
|
||||||
|
def is_available(name: str) -> bool:
|
||||||
|
"""Return ``True`` if the compiled kernel ``name`` was loaded."""
|
||||||
|
return _available.get(name, False)
|
||||||
|
|
||||||
|
|
||||||
|
def get_module(name: str) -> object:
|
||||||
|
"""Return the loaded kernel module for ``name``, or ``None`` if unavailable."""
|
||||||
|
return _modules.get(name)
|
||||||
@@ -0,0 +1,54 @@
|
|||||||
|
"""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
|
||||||
|
|
||||||
|
_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
|
||||||
|
):
|
||||||
|
from astrai.extension.rotary_ops import rotary_emb as _cuda_rotary
|
||||||
|
|
||||||
|
return _cuda_rotary(x, freqs_cis)
|
||||||
|
return _torch_apply(x, freqs_cis)
|
||||||
@@ -0,0 +1,39 @@
|
|||||||
|
"""Rotary embedding CUDA kernel wrapper.
|
||||||
|
|
||||||
|
Calls the compiled CUDA kernel directly. If the kernel is not available,
|
||||||
|
raises ``RuntimeError``. Fallback to torch complex multiply is the
|
||||||
|
responsibility of ``astrai.extension.rotary_backend.apply_rotary_emb``.
|
||||||
|
|
||||||
|
Layout: x is [batch, seq_len, n_heads, head_dim] (bf16, contiguous).
|
||||||
|
freqs_cis is [batch, seq_len, head_dim/2, 2] (f32, contiguous) — [cos, sin] pairs.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.extension.loader import _available, _modules
|
||||||
|
|
||||||
|
|
||||||
|
def _check_available():
|
||||||
|
if not _available.get("rotary_emb"):
|
||||||
|
raise RuntimeError(
|
||||||
|
"CUDA kernel 'rotary_emb' is not available. "
|
||||||
|
"Build with CSRC_KERNELS=true or use the torch fallback."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Fused rotary embedding kernel.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
x: [batch, seq_len, n_heads, head_dim] (bf16, contiguous)
|
||||||
|
freqs_cis: [batch, seq_len, head_dim/2, 2] (f32, contiguous) — [cos, sin] pairs
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[batch, seq_len, n_heads, head_dim] (bf16)
|
||||||
|
"""
|
||||||
|
_check_available()
|
||||||
|
if not x.is_contiguous():
|
||||||
|
x = x.contiguous()
|
||||||
|
if not freqs_cis.is_contiguous():
|
||||||
|
freqs_cis = freqs_cis.contiguous()
|
||||||
|
return _modules["rotary_emb"].rotary_emb(x, freqs_cis)
|
||||||
+99
-156
@@ -1,210 +1,153 @@
|
|||||||
"""Base factory class for extensible component registration."""
|
"""Base factory with decorator-based registration and kwarg-filtered instantiation."""
|
||||||
|
|
||||||
|
import inspect
|
||||||
|
import sys
|
||||||
from abc import ABC
|
from abc import ABC
|
||||||
from typing import Callable, Dict, Generic, List, Optional, Tuple, Type, TypeVar
|
from typing import (
|
||||||
|
Callable,
|
||||||
|
Dict,
|
||||||
|
ForwardRef,
|
||||||
|
Generic,
|
||||||
|
List,
|
||||||
|
Optional,
|
||||||
|
Type,
|
||||||
|
TypeVar,
|
||||||
|
Union,
|
||||||
|
get_args,
|
||||||
|
get_origin,
|
||||||
|
)
|
||||||
|
|
||||||
T = TypeVar("T")
|
T = TypeVar("T")
|
||||||
|
|
||||||
|
|
||||||
class Registry:
|
def _resolve_base_type(
|
||||||
"""Flexible registry for component classes with category and priority support.
|
arg: Union[Type, str, ForwardRef], factory_cls: type
|
||||||
|
) -> Optional[Type]:
|
||||||
|
"""Resolve the generic type-arg T to a concrete class.
|
||||||
|
|
||||||
This registry stores component classes with optional metadata (category, priority).
|
- Concrete class (``BaseFactory[MyBase]``): returned directly.
|
||||||
It provides methods for registration, retrieval, and listing with filtering.
|
- Forward reference (``BaseFactory["MyBase"]``): ``Base["X"]``
|
||||||
|
produces a ``ForwardRef("X")`` at class-creation time. We
|
||||||
|
extract the name and evaluate it in the factory module's
|
||||||
|
global namespace — the same mechanism ``typing.get_type_hints``
|
||||||
|
uses internally.
|
||||||
"""
|
"""
|
||||||
|
if isinstance(arg, type):
|
||||||
|
return arg
|
||||||
|
|
||||||
def __init__(self):
|
if isinstance(arg, str):
|
||||||
self._entries = {} # name -> (component_cls, category, priority)
|
name = arg
|
||||||
|
elif isinstance(arg, ForwardRef):
|
||||||
|
name = arg.__forward_arg__
|
||||||
|
else:
|
||||||
|
return None
|
||||||
|
|
||||||
def register(
|
mod = sys.modules.get(factory_cls.__module__)
|
||||||
self,
|
if mod is None:
|
||||||
name: str,
|
return None
|
||||||
component_cls: Type,
|
try:
|
||||||
category: Optional[str] = None,
|
return eval(name, vars(mod)) # noqa: S307
|
||||||
priority: int = 0,
|
except NameError:
|
||||||
) -> None:
|
return None
|
||||||
"""Register a component class with optional category and priority."""
|
|
||||||
if name in self._entries:
|
|
||||||
raise ValueError(f"Component '{name}' is already registered")
|
|
||||||
self._entries[name] = (component_cls, category, priority)
|
|
||||||
|
|
||||||
def get(self, name: str) -> Type:
|
|
||||||
"""Get component class by name."""
|
|
||||||
if name not in self._entries:
|
|
||||||
raise KeyError(f"Component '{name}' not found in registry")
|
|
||||||
return self._entries[name][0]
|
|
||||||
|
|
||||||
def get_with_metadata(self, name: str) -> Tuple[Type, Optional[str], int]:
|
def _validate_component(component_cls: Type, base: Optional[Type]) -> None:
|
||||||
"""Get component class with its metadata."""
|
"""Validate that *component_cls* inherits from *base*.
|
||||||
entry = self._entries.get(name)
|
|
||||||
if entry is None:
|
|
||||||
raise KeyError(f"Component '{name}' not found in registry")
|
|
||||||
return entry
|
|
||||||
|
|
||||||
def contains(self, name: str) -> bool:
|
No-op when *base* is ``None`` (e.g. forward-ref resolution failed).
|
||||||
"""Check if a name is registered."""
|
"""
|
||||||
return name in self._entries
|
if base is not None and not issubclass(component_cls, base):
|
||||||
|
raise TypeError(f"{component_cls.__name__} must inherit from {base.__name__}")
|
||||||
def list_names(self) -> List[str]:
|
|
||||||
"""Return list of registered component names."""
|
|
||||||
return sorted(self._entries.keys())
|
|
||||||
|
|
||||||
def list_by_category(self, category: str) -> List[str]:
|
|
||||||
"""Return names of components belonging to a specific category."""
|
|
||||||
return sorted(
|
|
||||||
name for name, (_, cat, _) in self._entries.items() if cat == category
|
|
||||||
)
|
|
||||||
|
|
||||||
def list_by_priority(self, reverse: bool = False) -> List[str]:
|
|
||||||
"""Return names sorted by priority (default ascending)."""
|
|
||||||
return sorted(
|
|
||||||
self._entries.keys(),
|
|
||||||
key=lambda name: self._entries[name][2],
|
|
||||||
reverse=reverse,
|
|
||||||
)
|
|
||||||
|
|
||||||
def entries(self) -> Dict[str, Tuple[Type, Optional[str], int]]:
|
|
||||||
"""Return raw entries dictionary."""
|
|
||||||
return self._entries.copy()
|
|
||||||
|
|
||||||
|
|
||||||
class BaseFactory(ABC, Generic[T]):
|
class BaseFactory(ABC, Generic[T]):
|
||||||
"""Generic factory class for component registration and creation.
|
"""Generic factory with decorator-based registration.
|
||||||
|
|
||||||
This base class provides a decorator-based registration pattern
|
Create a factory by subclassing with the desired base type::
|
||||||
for creating extensible component factories.
|
|
||||||
|
|
||||||
Example usage:
|
class MyFactory(BaseFactory[MyBase]):
|
||||||
class MyFactory(BaseFactory[MyBaseClass]):
|
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
Register components with the ``register`` decorator::
|
||||||
|
|
||||||
@MyFactory.register("custom")
|
@MyFactory.register("custom")
|
||||||
class CustomComponent(MyBaseClass):
|
class CustomComponent(MyBase):
|
||||||
...
|
...
|
||||||
|
|
||||||
component = MyFactory.create("custom", *args, **kwargs)
|
obj = MyFactory.create("custom", *args, **kwargs)
|
||||||
|
|
||||||
|
``create()`` filters kwargs to match the component's ``__init__``
|
||||||
|
signature so components don't need ``**kwargs`` just to absorb
|
||||||
|
unrelated parameters.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
_registry: Registry
|
_entries: Dict[str, Type[T]]
|
||||||
|
|
||||||
def __init_subclass__(cls, **kwargs):
|
def __init_subclass__(cls, **kwargs):
|
||||||
super().__init_subclass__(**kwargs)
|
super().__init_subclass__(**kwargs)
|
||||||
cls._registry = Registry()
|
for orig_base in getattr(cls, "__orig_bases__", ()):
|
||||||
|
if get_origin(orig_base) is BaseFactory:
|
||||||
|
(arg,) = get_args(orig_base)
|
||||||
|
cls._entries = {}
|
||||||
|
cls._component_base = _resolve_base_type(arg, cls)
|
||||||
|
return
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def register(
|
def register(cls, name: str) -> Callable[[Type[T]], Type[T]]:
|
||||||
cls, name: str, category: Optional[str] = None, priority: int = 0
|
"""Decorator to register a component class.
|
||||||
) -> Callable[[Type[T]], Type[T]]:
|
|
||||||
"""Decorator to register a component class with optional category and priority.
|
|
||||||
|
|
||||||
Args:
|
Validates that the decorated class inherits from the generic
|
||||||
name: Registration name for the component
|
type parameter ``T`` declared on the factory.
|
||||||
category: Optional category for grouping components
|
|
||||||
priority: Priority for ordering (default 0)
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Decorator function that registers the component class
|
|
||||||
|
|
||||||
Raises:
|
|
||||||
TypeError: If the decorated class doesn't inherit from the base type
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def decorator(component_cls: Type[T]) -> Type[T]:
|
def decorator(component_cls: Type[T]) -> Type[T]:
|
||||||
cls._validate_component(component_cls)
|
_validate_component(component_cls, cls._component_base)
|
||||||
cls._registry.register(
|
if name in cls._entries:
|
||||||
name, component_cls, category=category, priority=priority
|
raise ValueError(f"Component '{name}' is already registered")
|
||||||
)
|
cls._entries[name] = component_cls
|
||||||
return component_cls
|
return component_cls
|
||||||
|
|
||||||
return decorator
|
return decorator
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def create(cls, name: str, *args, **kwargs) -> T:
|
def create(cls, name: str, *args, **kwargs) -> T:
|
||||||
"""Create a component instance by name.
|
"""Create a component instance by name, filtering kwargs to match
|
||||||
|
the component's ``__init__`` signature.
|
||||||
Args:
|
|
||||||
name: Registered name of the component
|
|
||||||
*args: Positional arguments passed to component constructor
|
|
||||||
**kwargs: Keyword arguments passed to component constructor
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Component instance
|
|
||||||
|
|
||||||
Raises:
|
|
||||||
ValueError: If the component name is not registered
|
|
||||||
"""
|
"""
|
||||||
if not cls._registry.contains(name):
|
component_cls = cls._entries.get(name)
|
||||||
|
if component_cls is None:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Unknown component: '{name}'. "
|
f"Unknown component: '{name}'. Supported types: {sorted(cls._entries)}"
|
||||||
f"Supported types: {sorted(cls._registry.list_names())}"
|
|
||||||
)
|
)
|
||||||
component_cls = cls._registry.get(name)
|
sig = inspect.signature(component_cls.__init__)
|
||||||
|
has_var_kwargs = any(
|
||||||
|
p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()
|
||||||
|
)
|
||||||
|
if not has_var_kwargs:
|
||||||
|
valid = {
|
||||||
|
p.name
|
||||||
|
for p in sig.parameters.values()
|
||||||
|
if p.name != "self" and p.kind != inspect.Parameter.VAR_KEYWORD
|
||||||
|
}
|
||||||
|
kwargs = {k: v for k, v in kwargs.items() if k in valid}
|
||||||
return component_cls(*args, **kwargs)
|
return component_cls(*args, **kwargs)
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def _validate_component(cls, component_cls: Type[T]) -> None:
|
|
||||||
"""Validate that the component class is valid for this factory.
|
|
||||||
|
|
||||||
Override this method in subclasses to add custom validation.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
component_cls: Component class to validate
|
|
||||||
|
|
||||||
Raises:
|
|
||||||
TypeError: If the component class is invalid
|
|
||||||
"""
|
|
||||||
pass
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_component_class(cls, name: str) -> Type[T]:
|
def get_component_class(cls, name: str) -> Type[T]:
|
||||||
"""Get the registered component class by name without instantiating it.
|
"""Get the registered component class without instantiating it."""
|
||||||
|
entry = cls._entries.get(name)
|
||||||
Args:
|
if entry is None:
|
||||||
name: Registered name of the component
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
The component class itself
|
|
||||||
|
|
||||||
Raises:
|
|
||||||
ValueError: If the component name is not registered
|
|
||||||
"""
|
|
||||||
if not cls._registry.contains(name):
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Unknown component: '{name}'. "
|
f"Unknown component: '{name}'. Supported types: {sorted(cls._entries)}"
|
||||||
f"Supported types: {sorted(cls._registry.list_names())}"
|
|
||||||
)
|
)
|
||||||
return cls._registry.get(name)
|
return entry
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def list_registered(cls) -> list:
|
def list_registered(cls) -> List[str]:
|
||||||
"""List all registered component names.
|
"""List all registered component names."""
|
||||||
|
return sorted(cls._entries)
|
||||||
Returns:
|
|
||||||
List of registered component names
|
|
||||||
"""
|
|
||||||
return cls._registry.list_names()
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def is_registered(cls, name: str) -> bool:
|
def is_registered(cls, name: str) -> bool:
|
||||||
"""Check if a component name is registered.
|
"""Check if a component name is registered."""
|
||||||
|
return name in cls._entries
|
||||||
Args:
|
|
||||||
name: Component name to check
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True if registered, False otherwise
|
|
||||||
"""
|
|
||||||
return cls._registry.contains(name)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def list_by_category(cls, category: str) -> List[str]:
|
|
||||||
"""List registered component names in a category."""
|
|
||||||
return cls._registry.list_by_category(category)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def list_by_priority(cls, reverse: bool = False) -> List[str]:
|
|
||||||
"""List registered component names sorted by priority."""
|
|
||||||
return cls._registry.list_by_priority(reverse)
|
|
||||||
|
|
||||||
|
|
||||||
__all__ = ["Registry", "BaseFactory"]
|
|
||||||
|
|||||||
@@ -1,47 +1,51 @@
|
|||||||
"""Inference module for continuous batching.
|
"""Inference module for continuous batching.
|
||||||
|
|
||||||
Layers:
|
Layers:
|
||||||
- core/: Core inference loop (cache, executor, scheduler, task)
|
- core/: Core inference loop (cache, executor, scheduler, task)
|
||||||
- api/: HTTP protocol handlers (OpenAI, Anthropic)
|
- api/: HTTP orchestration (ProtocolHandler, server)
|
||||||
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
|
- protocols/: Response builders (OpenAI, Anthropic)
|
||||||
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy)
|
- transport/: SSE transport utilities
|
||||||
|
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
|
||||||
|
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy, FrequencyPenaltyStrategy)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from astrai.inference.api import (
|
from astrai.inference.api import (
|
||||||
AnthropicHandler,
|
|
||||||
AnthropicMessage,
|
AnthropicMessage,
|
||||||
|
BaseToolParser,
|
||||||
ChatCompletionRequest,
|
ChatCompletionRequest,
|
||||||
ChatMessage,
|
ChatMessage,
|
||||||
|
FunctionDef,
|
||||||
|
GenContext,
|
||||||
MessagesRequest,
|
MessagesRequest,
|
||||||
OpenAIHandler,
|
|
||||||
ProtocolHandler,
|
ProtocolHandler,
|
||||||
|
SimpleJsonToolParser,
|
||||||
StopChecker,
|
StopChecker,
|
||||||
StreamContext,
|
ToolDef,
|
||||||
app,
|
ToolParserFactory,
|
||||||
|
get_app,
|
||||||
run_server,
|
run_server,
|
||||||
)
|
)
|
||||||
|
from astrai.inference.api.anthropic import AnthropicResponseBuilder
|
||||||
|
from astrai.inference.api.openai import OpenAIResponseBuilder
|
||||||
from astrai.inference.core import (
|
from astrai.inference.core import (
|
||||||
STOP,
|
STOP,
|
||||||
Allocator,
|
Allocator,
|
||||||
Executor,
|
Executor,
|
||||||
InferenceScheduler,
|
InferenceScheduler,
|
||||||
KVCache,
|
KVCache,
|
||||||
KvcacheView,
|
KVStorage,
|
||||||
PagePool,
|
PagePool,
|
||||||
PrefixCache,
|
PrefixCache,
|
||||||
Storage,
|
ReqToTokenPool,
|
||||||
Task,
|
Task,
|
||||||
TaskManager,
|
TaskManager,
|
||||||
TaskStatus,
|
TaskStatus,
|
||||||
TaskTable,
|
|
||||||
page_hash,
|
page_hash,
|
||||||
)
|
)
|
||||||
from astrai.inference.engine import (
|
from astrai.inference.engine import GenerationRequest, InferenceEngine
|
||||||
GenerationRequest,
|
|
||||||
InferenceEngine,
|
|
||||||
)
|
|
||||||
from astrai.inference.sample import (
|
from astrai.inference.sample import (
|
||||||
BaseSamplingStrategy,
|
BaseSamplingStrategy,
|
||||||
|
FrequencyPenaltyStrategy,
|
||||||
SamplingPipeline,
|
SamplingPipeline,
|
||||||
TemperatureStrategy,
|
TemperatureStrategy,
|
||||||
TopKStrategy,
|
TopKStrategy,
|
||||||
@@ -50,43 +54,42 @@ from astrai.inference.sample import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
# Engine / Requests
|
|
||||||
"InferenceEngine",
|
"InferenceEngine",
|
||||||
"GenerationRequest",
|
"GenerationRequest",
|
||||||
# Core scheduler
|
|
||||||
"InferenceScheduler",
|
"InferenceScheduler",
|
||||||
"Executor",
|
"Executor",
|
||||||
"STOP",
|
"STOP",
|
||||||
"Task",
|
"Task",
|
||||||
"TaskManager",
|
"TaskManager",
|
||||||
"TaskStatus",
|
"TaskStatus",
|
||||||
# Core cache
|
|
||||||
"Allocator",
|
"Allocator",
|
||||||
"KVCache",
|
"KVCache",
|
||||||
"KvcacheView",
|
"KVStorage",
|
||||||
"PagePool",
|
"PagePool",
|
||||||
"PrefixCache",
|
"PrefixCache",
|
||||||
"Storage",
|
"ReqToTokenPool",
|
||||||
"TaskTable",
|
|
||||||
"page_hash",
|
"page_hash",
|
||||||
# Sampling (Strategy pattern)
|
|
||||||
"sample",
|
"sample",
|
||||||
"BaseSamplingStrategy",
|
"BaseSamplingStrategy",
|
||||||
"TemperatureStrategy",
|
"TemperatureStrategy",
|
||||||
"TopKStrategy",
|
"TopKStrategy",
|
||||||
"TopPStrategy",
|
"TopPStrategy",
|
||||||
|
"FrequencyPenaltyStrategy",
|
||||||
"SamplingPipeline",
|
"SamplingPipeline",
|
||||||
# Protocol
|
|
||||||
"ProtocolHandler",
|
"ProtocolHandler",
|
||||||
"StopChecker",
|
"StopChecker",
|
||||||
"StreamContext",
|
"GenContext",
|
||||||
"AnthropicHandler",
|
"BaseToolParser",
|
||||||
"OpenAIHandler",
|
"SimpleJsonToolParser",
|
||||||
# Server
|
"ToolParserFactory",
|
||||||
|
"OpenAIResponseBuilder",
|
||||||
|
"AnthropicResponseBuilder",
|
||||||
"ChatMessage",
|
"ChatMessage",
|
||||||
"ChatCompletionRequest",
|
"ChatCompletionRequest",
|
||||||
|
"FunctionDef",
|
||||||
|
"ToolDef",
|
||||||
"AnthropicMessage",
|
"AnthropicMessage",
|
||||||
"MessagesRequest",
|
"MessagesRequest",
|
||||||
"app",
|
"get_app",
|
||||||
"run_server",
|
"run_server",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -1,31 +1,39 @@
|
|||||||
"""Inference API: protocol handlers and FastAPI server."""
|
"""Inference API: protocol handler, stop checker, tool parsers, and FastAPI server.
|
||||||
|
|
||||||
from astrai.inference.api.protocol import (
|
``app`` is no longer a module-level global. Use :func:`get_app` to access the
|
||||||
AnthropicHandler,
|
lazy singleton FastAPI instance.
|
||||||
OpenAIHandler,
|
"""
|
||||||
ProtocolHandler,
|
|
||||||
StopChecker,
|
from astrai.inference.api.protocol import GenContext, ProtocolHandler, StopChecker
|
||||||
StreamContext,
|
|
||||||
)
|
|
||||||
from astrai.inference.api.server import (
|
from astrai.inference.api.server import (
|
||||||
AnthropicMessage,
|
AnthropicMessage,
|
||||||
ChatCompletionRequest,
|
ChatCompletionRequest,
|
||||||
ChatMessage,
|
ChatMessage,
|
||||||
|
FunctionDef,
|
||||||
MessagesRequest,
|
MessagesRequest,
|
||||||
app,
|
ToolDef,
|
||||||
|
get_app,
|
||||||
run_server,
|
run_server,
|
||||||
)
|
)
|
||||||
|
from astrai.inference.api.tool_parser import (
|
||||||
|
BaseToolParser,
|
||||||
|
SimpleJsonToolParser,
|
||||||
|
ToolParserFactory,
|
||||||
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"AnthropicHandler",
|
|
||||||
"OpenAIHandler",
|
|
||||||
"ProtocolHandler",
|
"ProtocolHandler",
|
||||||
"StopChecker",
|
"StopChecker",
|
||||||
"StreamContext",
|
"GenContext",
|
||||||
|
"BaseToolParser",
|
||||||
|
"SimpleJsonToolParser",
|
||||||
|
"ToolParserFactory",
|
||||||
"AnthropicMessage",
|
"AnthropicMessage",
|
||||||
"ChatCompletionRequest",
|
"ChatCompletionRequest",
|
||||||
"ChatMessage",
|
"ChatMessage",
|
||||||
|
"FunctionDef",
|
||||||
|
"ToolDef",
|
||||||
"MessagesRequest",
|
"MessagesRequest",
|
||||||
"app",
|
"get_app",
|
||||||
"run_server",
|
"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.api.protocol import (
|
||||||
|
GenContext,
|
||||||
|
ResponseBuilder,
|
||||||
|
StopInfo,
|
||||||
|
sse_event,
|
||||||
|
)
|
||||||
|
from astrai.inference.engine import InferenceEngine
|
||||||
|
|
||||||
|
|
||||||
|
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,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.api.protocol import (
|
||||||
|
GenContext,
|
||||||
|
ResponseBuilder,
|
||||||
|
StopInfo,
|
||||||
|
sse_event,
|
||||||
|
)
|
||||||
|
from astrai.inference.api.tool_parser import BaseToolParser, ToolParserFactory
|
||||||
|
from astrai.inference.engine import InferenceEngine
|
||||||
|
|
||||||
|
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,
|
||||||
|
},
|
||||||
|
}
|
||||||
+109
-343
@@ -1,15 +1,13 @@
|
|||||||
"""Protocol handlers for OpenAI and Anthropic chat completion APIs.
|
"""Orchestration layer: ProtocolHandler, StopChecker, GenContext, StopInfo, ResponseBuilder, SSE utils.
|
||||||
|
|
||||||
Template Method + Builder patterns eliminate the 45% code duplication between
|
ProtocolHandler orchestrates the async generation loop and delegates
|
||||||
stream/non-stream branches and across protocol adapters.
|
protocol-specific formatting to a ResponseBuilder.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
import time
|
|
||||||
import uuid
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, Dict, List, Optional, Union
|
from typing import Any, AsyncGenerator, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
from fastapi.responses import StreamingResponse
|
from fastapi.responses import StreamingResponse
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
@@ -17,7 +15,7 @@ from pydantic import BaseModel
|
|||||||
from astrai.inference.engine import InferenceEngine
|
from astrai.inference.engine import InferenceEngine
|
||||||
|
|
||||||
|
|
||||||
def _sse_event(data: Dict[str, Any], event: Optional[str] = None) -> str:
|
def sse_event(data: Dict[str, Any], event: Optional[str] = None) -> str:
|
||||||
lines: List[str] = []
|
lines: List[str] = []
|
||||||
if event:
|
if event:
|
||||||
lines.append(f"event: {event}")
|
lines.append(f"event: {event}")
|
||||||
@@ -26,22 +24,28 @@ def _sse_event(data: Dict[str, Any], event: Optional[str] = None) -> str:
|
|||||||
return "\n".join(lines)
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
|
||||||
def _sse_done() -> str:
|
def sse_done() -> str:
|
||||||
return "data: [DONE]\n\n"
|
return "data: [DONE]\n\n"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class StreamContext:
|
class GenContext:
|
||||||
"""Shared state across the streaming generation lifecycle."""
|
"""Per-generation metadata passed to builder format methods."""
|
||||||
|
|
||||||
resp_id: str
|
resp_id: str
|
||||||
created: int
|
created: int
|
||||||
model: str
|
model: str
|
||||||
prompt_tokens: int
|
prompt_tokens: int = 0
|
||||||
completion_tokens: int = 0
|
completion_tokens: int = 0
|
||||||
accumulated: str = ""
|
|
||||||
stop_matched: Optional[str] = None
|
|
||||||
last_yield_trimmed: str = ""
|
@dataclass
|
||||||
|
class StopInfo:
|
||||||
|
"""Stop-check result passed to format_stream_end / format_response."""
|
||||||
|
|
||||||
|
matched: Optional[str] = None
|
||||||
|
body: str = ""
|
||||||
|
yielded: str = ""
|
||||||
|
|
||||||
|
|
||||||
class StopChecker:
|
class StopChecker:
|
||||||
@@ -56,129 +60,116 @@ class StopChecker:
|
|||||||
return seq
|
return seq
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def trim(self, text: str, matched: str) -> str:
|
|
||||||
idx = text.rfind(matched)
|
|
||||||
return text[:idx] if idx != -1 else text
|
|
||||||
|
|
||||||
@property
|
class ResponseBuilder(ABC):
|
||||||
def has_sequences(self) -> bool:
|
"""Interface for protocol-specific response formatting.
|
||||||
return len(self._sequences) > 0
|
|
||||||
|
|
||||||
|
A new protocol requires one concrete builder implementing 5 methods.
|
||||||
class ProtocolHandler(ABC):
|
|
||||||
"""Template-method base for API protocol handlers.
|
|
||||||
|
|
||||||
Subclasses implement format hooks; the base class orchestrates the
|
|
||||||
generate-async loop and SSE/JSON response construction.
|
|
||||||
|
|
||||||
Lifecycle::
|
|
||||||
|
|
||||||
handle()
|
|
||||||
├─ build_prompt() # protocol-specific prompt assembly
|
|
||||||
├─ create_response_id() # unique response identifier
|
|
||||||
├─ [stream]
|
|
||||||
│ ├─ format_stream_start()
|
|
||||||
│ ├─ format_stream_token() × N
|
|
||||||
│ │ └─ on_token() hook for stop-sequence interception
|
|
||||||
│ └─ format_stream_end()
|
|
||||||
└─ [non-stream]
|
|
||||||
├─ (accumulate tokens)
|
|
||||||
└─ format_non_stream_response()
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
request_model: type[BaseModel]
|
@abstractmethod
|
||||||
|
def prepare(
|
||||||
|
self, request: BaseModel, engine: InferenceEngine
|
||||||
|
) -> Tuple[str, GenContext, List[str]]:
|
||||||
|
"""Return (prompt, ctx, stop_sequences) for a generation request."""
|
||||||
|
|
||||||
def __init__(self, request: BaseModel, engine: InferenceEngine):
|
@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.request = request
|
||||||
self.engine = engine
|
self.engine = engine
|
||||||
|
self.builder = builder
|
||||||
@abstractmethod
|
|
||||||
def build_prompt(self) -> str:
|
|
||||||
"""Build the full prompt string from the request messages."""
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def create_response_id(self) -> str:
|
|
||||||
"""Generate a unique response ID following the protocol convention."""
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def format_stream_start(self, ctx: StreamContext) -> List[str]:
|
|
||||||
"""Yield SSE events that open the stream (role marker, metadata)."""
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def format_stream_token(self, ctx: StreamContext, token: str) -> str:
|
|
||||||
"""Yield an SSE event for a single generated token."""
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def format_stream_end(self, ctx: StreamContext) -> List[str]:
|
|
||||||
"""Yield SSE events that close the stream (finish reason, usage stats)."""
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def format_non_stream_response(
|
|
||||||
self, ctx: StreamContext, content: str
|
|
||||||
) -> Dict[str, Any]:
|
|
||||||
"""Build the JSON response body for non-streaming mode."""
|
|
||||||
|
|
||||||
def get_stop_sequences(self) -> List[str]:
|
|
||||||
return []
|
|
||||||
|
|
||||||
def create_stop_checker(self) -> StopChecker:
|
|
||||||
return StopChecker(self.get_stop_sequences())
|
|
||||||
|
|
||||||
def on_token(
|
|
||||||
self, ctx: StreamContext, token: str, stop_checker: StopChecker
|
|
||||||
) -> Optional[str]:
|
|
||||||
"""Hook after each token is appended to accumulated.
|
|
||||||
|
|
||||||
Return a matched stop-sequence string to break the loop,
|
|
||||||
or None to continue.
|
|
||||||
|
|
||||||
"""
|
|
||||||
return None
|
|
||||||
|
|
||||||
async def handle(self) -> Union[StreamingResponse, Dict[str, Any]]:
|
async def handle(self) -> Union[StreamingResponse, Dict[str, Any]]:
|
||||||
ctx = StreamContext(
|
prompt, ctx, stop_sequences = self.builder.prepare(self.request, self.engine)
|
||||||
resp_id=self.create_response_id(),
|
ctx.prompt_tokens = len(self.engine.tokenizer.encode(prompt))
|
||||||
created=int(time.time()),
|
|
||||||
model=self.request.model,
|
|
||||||
prompt_tokens=self._count_prompt_tokens(),
|
|
||||||
)
|
|
||||||
|
|
||||||
agen = self.engine.generate_async(
|
agen = self.engine.generate_async(
|
||||||
prompt=self.build_prompt(),
|
prompt=prompt,
|
||||||
max_tokens=self.request.max_tokens,
|
max_tokens=self.request.max_tokens,
|
||||||
temperature=self.request.temperature,
|
temperature=self.request.temperature,
|
||||||
top_p=self.request.top_p,
|
top_p=self.request.top_p,
|
||||||
top_k=self.request.top_k,
|
top_k=self.request.top_k,
|
||||||
|
frequency_penalty=getattr(self.request, "frequency_penalty", 0.0),
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.request.stream:
|
if self.request.stream:
|
||||||
return self._handle_stream(agen, ctx)
|
return self._handle_stream(agen, ctx, stop_sequences)
|
||||||
else:
|
else:
|
||||||
return await self._handle_non_stream(agen, ctx)
|
return await self._handle_non_stream(agen, ctx, stop_sequences)
|
||||||
|
|
||||||
def _count_prompt_tokens(self) -> int:
|
def _handle_stream(
|
||||||
return len(self.engine.tokenizer.encode(self.build_prompt()))
|
self, agen: AsyncGenerator, ctx: GenContext, stop_sequences: List[str]
|
||||||
|
) -> StreamingResponse:
|
||||||
def _handle_stream(self, agen, ctx: StreamContext) -> StreamingResponse:
|
checker = StopChecker(stop_sequences)
|
||||||
stop_checker = self.create_stop_checker()
|
|
||||||
|
|
||||||
async def event_stream():
|
async def event_stream():
|
||||||
for event in self.format_stream_start(ctx):
|
for event in self.builder.format_stream_start(ctx):
|
||||||
yield event
|
yield event
|
||||||
|
|
||||||
|
body = ""
|
||||||
|
yielded = ""
|
||||||
|
matched = None
|
||||||
|
token_ids: List[int] = []
|
||||||
async for token in agen:
|
async for token in agen:
|
||||||
ctx.completion_tokens += 1
|
body += token
|
||||||
ctx.accumulated += token
|
|
||||||
|
|
||||||
matched = self.on_token(ctx, token, stop_checker)
|
new_ids = self.engine.tokenizer.encode(token)
|
||||||
|
token_ids.extend(new_ids)
|
||||||
|
|
||||||
|
matched = checker.check(body)
|
||||||
if matched:
|
if matched:
|
||||||
break
|
break
|
||||||
|
|
||||||
yield self.format_stream_token(ctx, token)
|
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
|
||||||
|
|
||||||
for event in self.format_stream_end(ctx):
|
stop = StopInfo(matched=matched, body=body, yielded=yielded)
|
||||||
|
for event in self.builder.format_stream_end(ctx, stop):
|
||||||
yield event
|
yield event
|
||||||
yield _sse_done()
|
yield sse_done()
|
||||||
|
|
||||||
return StreamingResponse(
|
return StreamingResponse(
|
||||||
event_stream(),
|
event_stream(),
|
||||||
@@ -186,249 +177,24 @@ class ProtocolHandler(ABC):
|
|||||||
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _handle_non_stream(self, agen, ctx: StreamContext) -> Dict[str, Any]:
|
async def _handle_non_stream(
|
||||||
stop_checker = self.create_stop_checker()
|
self, agen: AsyncGenerator, ctx: GenContext, stop_sequences: List[str]
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
checker = StopChecker(stop_sequences)
|
||||||
chunks: List[str] = []
|
chunks: List[str] = []
|
||||||
|
body = ""
|
||||||
|
matched = None
|
||||||
|
|
||||||
async for token in agen:
|
async for token in agen:
|
||||||
ctx.completion_tokens += 1
|
|
||||||
ctx.accumulated += token
|
|
||||||
chunks.append(token)
|
chunks.append(token)
|
||||||
|
body += token
|
||||||
|
|
||||||
matched = self.on_token(ctx, token, stop_checker)
|
matched = checker.check(body)
|
||||||
if matched:
|
if matched:
|
||||||
break
|
break
|
||||||
|
|
||||||
|
ctx.completion_tokens += 1
|
||||||
|
|
||||||
content = "".join(chunks)
|
content = "".join(chunks)
|
||||||
return self.format_non_stream_response(ctx, content)
|
stop = StopInfo(matched=matched, body=body)
|
||||||
|
return self.builder.format_response(ctx, content, stop)
|
||||||
|
|
||||||
def _extract_text_content(content: Union[str, List[Dict[str, Any]]]) -> str:
|
|
||||||
"""Extract plain text from an Anthropic content block (string or list)."""
|
|
||||||
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 OpenAIHandler(ProtocolHandler):
|
|
||||||
"""OpenAI-compatible /v1/chat/completions handler."""
|
|
||||||
|
|
||||||
def build_prompt(self) -> str:
|
|
||||||
messages = [
|
|
||||||
{"role": m.role, "content": m.content} for m in self.request.messages
|
|
||||||
]
|
|
||||||
return self.engine.tokenizer.apply_chat_template(messages, tokenize=False)
|
|
||||||
|
|
||||||
def create_response_id(self) -> str:
|
|
||||||
return f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
|
||||||
|
|
||||||
def format_stream_start(self, ctx: StreamContext) -> List[str]:
|
|
||||||
return [
|
|
||||||
_sse_event(
|
|
||||||
{
|
|
||||||
"id": ctx.resp_id,
|
|
||||||
"object": "chat.completion.chunk",
|
|
||||||
"created": ctx.created,
|
|
||||||
"model": ctx.model,
|
|
||||||
"choices": [
|
|
||||||
{
|
|
||||||
"index": 0,
|
|
||||||
"delta": {"role": "assistant"},
|
|
||||||
"finish_reason": None,
|
|
||||||
}
|
|
||||||
],
|
|
||||||
}
|
|
||||||
)
|
|
||||||
]
|
|
||||||
|
|
||||||
def format_stream_token(self, ctx: StreamContext, token: str) -> str:
|
|
||||||
return _sse_event(
|
|
||||||
{
|
|
||||||
"id": ctx.resp_id,
|
|
||||||
"object": "chat.completion.chunk",
|
|
||||||
"created": ctx.created,
|
|
||||||
"model": ctx.model,
|
|
||||||
"choices": [
|
|
||||||
{"index": 0, "delta": {"content": token}, "finish_reason": None}
|
|
||||||
],
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
def format_stream_end(self, ctx: StreamContext) -> List[str]:
|
|
||||||
return [
|
|
||||||
_sse_event(
|
|
||||||
{
|
|
||||||
"id": ctx.resp_id,
|
|
||||||
"object": "chat.completion.chunk",
|
|
||||||
"created": ctx.created,
|
|
||||||
"model": ctx.model,
|
|
||||||
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
|
|
||||||
}
|
|
||||||
),
|
|
||||||
_sse_event(
|
|
||||||
{
|
|
||||||
"prompt_tokens": ctx.prompt_tokens,
|
|
||||||
"completion_tokens": ctx.completion_tokens,
|
|
||||||
"total_tokens": ctx.prompt_tokens + ctx.completion_tokens,
|
|
||||||
}
|
|
||||||
),
|
|
||||||
]
|
|
||||||
|
|
||||||
def format_non_stream_response(
|
|
||||||
self, ctx: StreamContext, content: str
|
|
||||||
) -> Dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"id": ctx.resp_id,
|
|
||||||
"object": "chat.completion",
|
|
||||||
"created": ctx.created,
|
|
||||||
"model": ctx.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,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
class AnthropicHandler(ProtocolHandler):
|
|
||||||
"""Anthropic-compatible /v1/messages handler."""
|
|
||||||
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
self._yielded = ""
|
|
||||||
|
|
||||||
def build_prompt(self) -> str:
|
|
||||||
messages: List[Dict[str, str]] = []
|
|
||||||
system = getattr(self.request, "system", None)
|
|
||||||
if system:
|
|
||||||
messages.append({"role": "system", "content": system})
|
|
||||||
for m in self.request.messages:
|
|
||||||
content = _extract_text_content(m.content)
|
|
||||||
if content:
|
|
||||||
messages.append({"role": m.role, "content": content})
|
|
||||||
return self.engine.tokenizer.apply_chat_template(messages, tokenize=False)
|
|
||||||
|
|
||||||
def create_response_id(self) -> str:
|
|
||||||
return f"msg_{uuid.uuid4().hex[:24]}"
|
|
||||||
|
|
||||||
def get_stop_sequences(self) -> List[str]:
|
|
||||||
return getattr(self.request, "stop_sequences", None) or []
|
|
||||||
|
|
||||||
def on_token(
|
|
||||||
self, ctx: StreamContext, token: str, stop_checker: StopChecker
|
|
||||||
) -> Optional[str]:
|
|
||||||
matched = stop_checker.check(ctx.accumulated)
|
|
||||||
if not matched:
|
|
||||||
return None
|
|
||||||
|
|
||||||
ctx.stop_matched = matched
|
|
||||||
trimmed = ctx.accumulated[: ctx.accumulated.rfind(matched)]
|
|
||||||
unyielded = trimmed[len(self._yielded) :]
|
|
||||||
if unyielded:
|
|
||||||
ctx.last_yield_trimmed = unyielded
|
|
||||||
return matched
|
|
||||||
|
|
||||||
def format_stream_start(self, ctx: StreamContext) -> 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_stream_token(self, ctx: StreamContext, token: str) -> str:
|
|
||||||
self._yielded += token
|
|
||||||
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: StreamContext) -> List[str]:
|
|
||||||
matched = ctx.stop_matched
|
|
||||||
events: List[str] = []
|
|
||||||
last_yielded = ctx.last_yield_trimmed
|
|
||||||
if last_yielded:
|
|
||||||
events.append(
|
|
||||||
_sse_event(
|
|
||||||
{
|
|
||||||
"type": "content_block_delta",
|
|
||||||
"index": 0,
|
|
||||||
"delta": {"type": "text_delta", "text": last_yielded},
|
|
||||||
},
|
|
||||||
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 matched else "end_turn",
|
|
||||||
"stop_sequence": matched,
|
|
||||||
},
|
|
||||||
"usage": {"output_tokens": ctx.completion_tokens},
|
|
||||||
},
|
|
||||||
event="message_delta",
|
|
||||||
)
|
|
||||||
)
|
|
||||||
events.append(_sse_event({"type": "message_stop"}, event="message_stop"))
|
|
||||||
return events
|
|
||||||
|
|
||||||
def format_non_stream_response(
|
|
||||||
self, ctx: StreamContext, content: str
|
|
||||||
) -> Dict[str, Any]:
|
|
||||||
matched = ctx.stop_matched
|
|
||||||
if matched:
|
|
||||||
content = content[: content.rfind(matched)]
|
|
||||||
return {
|
|
||||||
"id": ctx.resp_id,
|
|
||||||
"type": "message",
|
|
||||||
"role": "assistant",
|
|
||||||
"model": ctx.model,
|
|
||||||
"content": [{"type": "text", "text": content}],
|
|
||||||
"stop_reason": "stop_sequence" if matched else "end_turn",
|
|
||||||
"stop_sequence": matched,
|
|
||||||
"usage": {
|
|
||||||
"input_tokens": ctx.prompt_tokens,
|
|
||||||
"output_tokens": ctx.completion_tokens,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -3,6 +3,9 @@ OpenAI / Anthropic-compatible chat completion server backed by continuous-batchi
|
|||||||
|
|
||||||
Protocol-specific formatting is delegated to ``astrai.inference.protocol``.
|
Protocol-specific formatting is delegated to ``astrai.inference.protocol``.
|
||||||
This module owns the FastAPI app, request/response schemas, and dependency wiring.
|
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
|
import logging
|
||||||
@@ -12,22 +15,37 @@ from typing import Any, Dict, List, Optional, Union
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
import uvicorn
|
import uvicorn
|
||||||
from fastapi import FastAPI, HTTPException, Request
|
from fastapi import APIRouter, FastAPI, HTTPException
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
from astrai.inference.api.protocol import AnthropicHandler, OpenAIHandler
|
from astrai.inference.api.anthropic import AnthropicResponseBuilder
|
||||||
|
from astrai.inference.api.openai import OpenAIResponseBuilder
|
||||||
|
from astrai.inference.api.protocol import ProtocolHandler
|
||||||
from astrai.inference.engine import InferenceEngine
|
from astrai.inference.engine import InferenceEngine
|
||||||
from astrai.model import AutoModel
|
from astrai.model import AutoModel
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
_project_root = Path(__file__).parent.parent.parent
|
_app_instance: Optional[FastAPI] = None
|
||||||
|
|
||||||
|
|
||||||
class ChatMessage(BaseModel):
|
class ChatMessage(BaseModel):
|
||||||
role: str
|
role: str
|
||||||
content: 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):
|
class ChatCompletionRequest(BaseModel):
|
||||||
@@ -46,6 +64,8 @@ class ChatCompletionRequest(BaseModel):
|
|||||||
frequency_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
|
logit_bias: Optional[Dict[int, float]] = None
|
||||||
user: Optional[str] = None
|
user: Optional[str] = None
|
||||||
|
tools: Optional[List[ToolDef]] = None
|
||||||
|
tool_choice: Optional[Union[str, Dict[str, Any]]] = "auto"
|
||||||
|
|
||||||
|
|
||||||
class AnthropicMessage(BaseModel):
|
class AnthropicMessage(BaseModel):
|
||||||
@@ -67,31 +87,6 @@ class MessagesRequest(BaseModel):
|
|||||||
stop_sequences: Optional[List[str]] = None
|
stop_sequences: Optional[List[str]] = None
|
||||||
|
|
||||||
|
|
||||||
def _create_engine(
|
|
||||||
param_path: Optional[Path] = None,
|
|
||||||
device: str = "cuda",
|
|
||||||
dtype: torch.dtype = torch.bfloat16,
|
|
||||||
max_batch_size: int = 16,
|
|
||||||
) -> InferenceEngine:
|
|
||||||
if param_path is None:
|
|
||||||
param_path = _project_root / "params"
|
|
||||||
if not param_path.exists():
|
|
||||||
raise FileNotFoundError(f"Parameter directory not found: {param_path}")
|
|
||||||
|
|
||||||
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
|
||||||
model = AutoModel.from_pretrained(param_path)
|
|
||||||
model.to(device=device, dtype=dtype)
|
|
||||||
logger.info(f"Model loaded on {device} with dtype {dtype}")
|
|
||||||
|
|
||||||
engine = InferenceEngine(
|
|
||||||
model=model,
|
|
||||||
tokenizer=tokenizer,
|
|
||||||
max_batch_size=max_batch_size,
|
|
||||||
)
|
|
||||||
logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}")
|
|
||||||
return engine
|
|
||||||
|
|
||||||
|
|
||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
async def lifespan(app: FastAPI):
|
async def lifespan(app: FastAPI):
|
||||||
config = app.state.server_config
|
config = app.state.server_config
|
||||||
@@ -107,60 +102,105 @@ async def lifespan(app: FastAPI):
|
|||||||
logger.info("Inference engine shutdown complete")
|
logger.info("Inference engine shutdown complete")
|
||||||
|
|
||||||
|
|
||||||
app = FastAPI(title="AstrAI Inference Server", version="0.2.0", lifespan=lifespan)
|
router = APIRouter()
|
||||||
|
|
||||||
|
|
||||||
def _get_engine(request: Request) -> InferenceEngine:
|
def _create_engine(
|
||||||
engine = request.app.state.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:
|
if engine is None:
|
||||||
raise HTTPException(status_code=503, detail="Engine not initialized")
|
raise HTTPException(status_code=503, detail="Engine not initialized")
|
||||||
return engine
|
return engine
|
||||||
|
|
||||||
|
|
||||||
@app.get("/health")
|
@router.get("/health")
|
||||||
async def health(request: Request):
|
async def health():
|
||||||
|
app = get_app()
|
||||||
return {
|
return {
|
||||||
"status": "ok",
|
"status": "ok",
|
||||||
"model_loaded": request.app.state.engine is not None,
|
"model_loaded": app.state.engine is not None,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@app.get("/stats")
|
@router.get("/stats")
|
||||||
async def get_stats(request: Request):
|
async def get_stats():
|
||||||
return _get_engine(request).get_stats()
|
return _get_engine().get_stats()
|
||||||
|
|
||||||
|
|
||||||
@app.post("/v1/chat/completions")
|
@router.post("/v1/chat/completions")
|
||||||
async def chat_completion(request: ChatCompletionRequest, req: Request):
|
async def chat_completion(request: ChatCompletionRequest):
|
||||||
engine = _get_engine(req)
|
engine = _get_engine()
|
||||||
handler = OpenAIHandler(request, engine)
|
handler = ProtocolHandler(request, engine, OpenAIResponseBuilder())
|
||||||
return await handler.handle()
|
return await handler.handle()
|
||||||
|
|
||||||
|
|
||||||
@app.post("/v1/messages")
|
@router.post("/v1/messages")
|
||||||
async def create_message(request: MessagesRequest, req: Request):
|
async def create_message(request: MessagesRequest):
|
||||||
engine = _get_engine(req)
|
engine = _get_engine()
|
||||||
handler = AnthropicHandler(request, engine)
|
handler = ProtocolHandler(request, engine, AnthropicResponseBuilder())
|
||||||
return await handler.handle()
|
return await handler.handle()
|
||||||
|
|
||||||
|
|
||||||
def run_server(
|
def run_server(
|
||||||
|
param_path: Path,
|
||||||
host: str = "0.0.0.0",
|
host: str = "0.0.0.0",
|
||||||
port: int = 8000,
|
port: int = 8000,
|
||||||
reload: bool = False,
|
reload: bool = False,
|
||||||
device: str = "cuda",
|
device: str = "cuda",
|
||||||
dtype: torch.dtype = torch.bfloat16,
|
dtype: torch.dtype = torch.bfloat16,
|
||||||
param_path: Optional[Path] = None,
|
|
||||||
max_batch_size: int = 16,
|
max_batch_size: int = 16,
|
||||||
|
max_seq_len: Optional[int] = None,
|
||||||
):
|
):
|
||||||
|
app = get_app()
|
||||||
app.state.server_config = {
|
app.state.server_config = {
|
||||||
"device": device,
|
"device": device,
|
||||||
"dtype": dtype,
|
"dtype": dtype,
|
||||||
"param_path": param_path,
|
"param_path": param_path,
|
||||||
"max_batch_size": max_batch_size,
|
"max_batch_size": max_batch_size,
|
||||||
|
"max_seq_len": max_seq_len,
|
||||||
}
|
}
|
||||||
uvicorn.run(
|
uvicorn.run(
|
||||||
app,
|
app,
|
||||||
host=host,
|
host=host,
|
||||||
port=port,
|
port=port,
|
||||||
|
reload=reload,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -3,11 +3,10 @@
|
|||||||
from astrai.inference.core.cache import (
|
from astrai.inference.core.cache import (
|
||||||
Allocator,
|
Allocator,
|
||||||
KVCache,
|
KVCache,
|
||||||
KvcacheView,
|
KVStorage,
|
||||||
PagePool,
|
PagePool,
|
||||||
PrefixCache,
|
PrefixCache,
|
||||||
Storage,
|
ReqToTokenPool,
|
||||||
TaskTable,
|
|
||||||
page_hash,
|
page_hash,
|
||||||
)
|
)
|
||||||
from astrai.inference.core.executor import Executor
|
from astrai.inference.core.executor import Executor
|
||||||
@@ -17,11 +16,10 @@ from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
|||||||
__all__ = [
|
__all__ = [
|
||||||
"Allocator",
|
"Allocator",
|
||||||
"KVCache",
|
"KVCache",
|
||||||
"KvcacheView",
|
"KVStorage",
|
||||||
"PagePool",
|
"PagePool",
|
||||||
"PrefixCache",
|
"PrefixCache",
|
||||||
"Storage",
|
"ReqToTokenPool",
|
||||||
"TaskTable",
|
|
||||||
"page_hash",
|
"page_hash",
|
||||||
"Executor",
|
"Executor",
|
||||||
"InferenceScheduler",
|
"InferenceScheduler",
|
||||||
|
|||||||
+332
-214
@@ -1,6 +1,21 @@
|
|||||||
|
"""KV cache architecture: three-layer separation (SGLang-inspired).
|
||||||
|
|
||||||
|
Layer 1 — KVStorage: flat token-level K/V buffers [n_layers, size, H, D]
|
||||||
|
Layer 2 — ReqToTokenPool: index table [req_idx, pos] → physical token slot
|
||||||
|
Layer 3 — Allocator: slot/page allocation with ref-counting and LRU
|
||||||
|
|
||||||
|
PagePool orchestrates all three plus PrefixCache (content addressing).
|
||||||
|
KVCache is a pure dataclass passed to the model for direct buffer access.
|
||||||
|
|
||||||
|
Two modes:
|
||||||
|
- contiguous (default): pre-allocated per-request blocks, no dynamic alloc
|
||||||
|
- paged: shared pool with on-demand allocation, prefix caching support
|
||||||
|
"""
|
||||||
|
|
||||||
import threading
|
import threading
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
from typing import Callable, Dict, List, Optional, Tuple
|
from dataclasses import dataclass
|
||||||
|
from typing import Callable, Dict, List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
@@ -42,7 +57,7 @@ class Allocator:
|
|||||||
return idx
|
return idx
|
||||||
return -1
|
return -1
|
||||||
|
|
||||||
def free(self, idx: int, keep_cached: bool = False) -> None:
|
def free(self, idx: int, keep_cached: bool = False):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._refs[idx] -= 1
|
self._refs[idx] -= 1
|
||||||
if self._refs[idx] == 0:
|
if self._refs[idx] == 0:
|
||||||
@@ -51,7 +66,7 @@ class Allocator:
|
|||||||
else:
|
else:
|
||||||
self._free_mask |= 1 << idx
|
self._free_mask |= 1 << idx
|
||||||
|
|
||||||
def inc_ref(self, idx: int) -> None:
|
def inc_ref(self, idx: int):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._refs[idx] += 1
|
self._refs[idx] += 1
|
||||||
self._lru.pop(idx, None)
|
self._lru.pop(idx, None)
|
||||||
@@ -60,9 +75,10 @@ class Allocator:
|
|||||||
with self._lock:
|
with self._lock:
|
||||||
return self._refs[idx]
|
return self._refs[idx]
|
||||||
|
|
||||||
def touch(self, idx: int) -> None:
|
def touch(self, idx: int):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._lru.move_to_end(idx)
|
if idx in self._lru:
|
||||||
|
self._lru.move_to_end(idx)
|
||||||
|
|
||||||
|
|
||||||
class PrefixCache:
|
class PrefixCache:
|
||||||
@@ -74,7 +90,7 @@ class PrefixCache:
|
|||||||
self._hash_to_page: Dict[int, int] = {}
|
self._hash_to_page: Dict[int, int] = {}
|
||||||
self._lock = threading.Lock()
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
def evict(self, idx: int) -> None:
|
def evict(self, idx: int):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
h = self._page_to_hash.pop(idx, None)
|
h = self._page_to_hash.pop(idx, None)
|
||||||
if h is not None:
|
if h is not None:
|
||||||
@@ -96,9 +112,7 @@ class PrefixCache:
|
|||||||
hits.append(p)
|
hits.append(p)
|
||||||
return hits
|
return hits
|
||||||
|
|
||||||
def record(
|
def record(self, page_idx: int, token_ids: List[int], logical_page_idx: int):
|
||||||
self, page_idx: int, token_ids: List[int], logical_page_idx: int
|
|
||||||
) -> None:
|
|
||||||
with self._lock:
|
with self._lock:
|
||||||
h = page_hash(token_ids, logical_page_idx, self._page_size)
|
h = page_hash(token_ids, logical_page_idx, self._page_size)
|
||||||
old_h = self._page_to_hash.pop(page_idx, None)
|
old_h = self._page_to_hash.pop(page_idx, None)
|
||||||
@@ -108,265 +122,369 @@ class PrefixCache:
|
|||||||
self._hash_to_page[h] = page_idx
|
self._hash_to_page[h] = page_idx
|
||||||
|
|
||||||
|
|
||||||
class PagePool:
|
class ReqToTokenPool:
|
||||||
"""Orchestrates allocator (page management) and PrefixCache (content addressing)."""
|
"""Maps [req_idx, pos] -> physical token slot in KV storage.
|
||||||
|
|
||||||
def __init__(self, allocator: Allocator, prefix: PrefixCache):
|
Each row is one request; each column is a sequence position. The value
|
||||||
self._alloc = allocator
|
at [req_idx, pos] is the flat index into the KV storage buffers.
|
||||||
self._prefix = prefix
|
"""
|
||||||
self._alloc.on_evict = prefix.evict
|
|
||||||
|
|
||||||
@property
|
def __init__(self, size: int, max_context_len: int, device: torch.device):
|
||||||
def allocator(self) -> Allocator:
|
self.size = size
|
||||||
return self._alloc
|
self.max_context_len = max_context_len
|
||||||
|
self.req_to_token = torch.zeros(
|
||||||
@property
|
(size, max_context_len), dtype=torch.long, device=device
|
||||||
def prefix(self) -> PrefixCache:
|
)
|
||||||
return self._prefix
|
self.free_slots = list(range(size))
|
||||||
|
|
||||||
def alloc(self) -> int:
|
|
||||||
return self._alloc.alloc()
|
|
||||||
|
|
||||||
def free(self, idx: int) -> None:
|
|
||||||
keep = self._prefix.has_page(idx)
|
|
||||||
self._alloc.free(idx, keep_cached=keep)
|
|
||||||
if not keep:
|
|
||||||
self._prefix.evict(idx)
|
|
||||||
|
|
||||||
def inc_ref(self, idx: int) -> None:
|
|
||||||
self._alloc.inc_ref(idx)
|
|
||||||
|
|
||||||
def lookup(self, token_ids: List[int]) -> List[int]:
|
|
||||||
hits = self._prefix.lookup(token_ids)
|
|
||||||
for p in hits:
|
|
||||||
self._alloc.touch(p)
|
|
||||||
return hits
|
|
||||||
|
|
||||||
def record(
|
|
||||||
self, page_idx: int, token_ids: List[int], logical_page_idx: int
|
|
||||||
) -> None:
|
|
||||||
self._prefix.record(page_idx, token_ids, logical_page_idx)
|
|
||||||
|
|
||||||
|
|
||||||
class TaskTable:
|
|
||||||
"""Maps task_ids to page tables and cached token counts."""
|
|
||||||
|
|
||||||
def __init__(self, page_size: int):
|
|
||||||
self._page_size = page_size
|
|
||||||
self._pages: Dict[str, List[int]] = {}
|
|
||||||
self._cached: Dict[str, int] = {}
|
|
||||||
self._lock = threading.Lock()
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
def set(self, task_id: str, page_table: List[int], cached: int) -> None:
|
def alloc(self, num_reqs: int) -> Optional[List[int]]:
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self._pages[task_id] = page_table
|
if num_reqs > len(self.free_slots):
|
||||||
self._cached[task_id] = cached
|
return None
|
||||||
|
slots = self.free_slots[:num_reqs]
|
||||||
|
self.free_slots = self.free_slots[num_reqs:]
|
||||||
|
return slots
|
||||||
|
|
||||||
def get(self, task_id: str) -> List[int]:
|
def free(self, req_indices: List[int]):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
return self._pages.get(task_id, [])
|
self.free_slots.extend(req_indices)
|
||||||
|
|
||||||
def get_cached(self, task_id: str) -> int:
|
def write(self, indices, values):
|
||||||
with self._lock:
|
self.req_to_token[indices] = values
|
||||||
return self._cached.get(task_id, 0)
|
|
||||||
|
|
||||||
def pop(self, task_id: str) -> Tuple[List[int], int]:
|
|
||||||
with self._lock:
|
|
||||||
pages = self._pages.pop(task_id, [])
|
|
||||||
cached = self._cached.pop(task_id, 0)
|
|
||||||
return pages, cached
|
|
||||||
|
|
||||||
def get_ref(self, task_id: str) -> List[int]:
|
|
||||||
with self._lock:
|
|
||||||
return self._pages.setdefault(task_id, [])
|
|
||||||
|
|
||||||
def table_tensor(self, task_ids: List[str], device: torch.device) -> Tensor:
|
|
||||||
with self._lock:
|
|
||||||
states = [self._pages.get(tid, []) for tid in task_ids]
|
|
||||||
max_pages = max((len(s) for s in states), default=0)
|
|
||||||
rows = [s + [-1] * (max_pages - len(s)) for s in states]
|
|
||||||
return torch.tensor(rows, dtype=torch.long, device=device)
|
|
||||||
|
|
||||||
|
|
||||||
class Storage:
|
class KVStorage:
|
||||||
"""KV-cache tensor storage with paged write/gather."""
|
"""Token-level KV cache storage.
|
||||||
|
|
||||||
|
Buffers: [n_layers, size, n_kv_heads, head_dim]. Each token occupies
|
||||||
|
one slot indexed by ReqToTokenPool.
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
|
size: int,
|
||||||
n_layers: int,
|
n_layers: int,
|
||||||
n_pages: int,
|
|
||||||
page_size: int,
|
|
||||||
n_kv_heads: int,
|
n_kv_heads: int,
|
||||||
head_dim: int,
|
head_dim: int,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
dtype: torch.dtype,
|
dtype: torch.dtype,
|
||||||
):
|
):
|
||||||
self.page_size = page_size
|
self.size = size
|
||||||
self.k_cache = torch.empty(
|
self.k_buffer = torch.empty(
|
||||||
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
|
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
|
||||||
device=device,
|
|
||||||
dtype=dtype,
|
|
||||||
)
|
)
|
||||||
self.v_cache = torch.empty(
|
self.v_buffer = torch.empty(
|
||||||
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
|
(n_layers, size, n_kv_heads, head_dim), device=device, dtype=dtype
|
||||||
device=device,
|
|
||||||
dtype=dtype,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def write(
|
def get_key_buffer(self, layer_id: int) -> Tensor:
|
||||||
self,
|
return self.k_buffer[layer_id]
|
||||||
layer_id: int,
|
|
||||||
page_table: Tensor,
|
|
||||||
start_pos: int,
|
|
||||||
k: Tensor,
|
|
||||||
v: Tensor,
|
|
||||||
) -> None:
|
|
||||||
seq_len = k.size(1)
|
|
||||||
if seq_len == 0:
|
|
||||||
return
|
|
||||||
page_size = self.page_size
|
|
||||||
written = 0
|
|
||||||
first_page = start_pos // page_size
|
|
||||||
last_page = (start_pos + seq_len - 1) // page_size
|
|
||||||
for pi in range(first_page, last_page + 1):
|
|
||||||
phys_pages = page_table[:, pi]
|
|
||||||
page_start = pi * page_size
|
|
||||||
write_start = max(page_start, start_pos)
|
|
||||||
write_end = min(page_start + page_size, start_pos + seq_len)
|
|
||||||
offset = write_start - page_start
|
|
||||||
chunk = write_end - write_start
|
|
||||||
valid = phys_pages >= 0
|
|
||||||
if not valid.all():
|
|
||||||
if valid.any():
|
|
||||||
valid_pages = phys_pages[valid]
|
|
||||||
self.k_cache[layer_id, valid_pages, offset : offset + chunk] = k[
|
|
||||||
valid, written : written + chunk
|
|
||||||
]
|
|
||||||
self.v_cache[layer_id, valid_pages, offset : offset + chunk] = v[
|
|
||||||
valid, written : written + chunk
|
|
||||||
]
|
|
||||||
written += chunk
|
|
||||||
continue
|
|
||||||
self.k_cache[layer_id, phys_pages, offset : offset + chunk] = k[
|
|
||||||
:, written : written + chunk
|
|
||||||
]
|
|
||||||
self.v_cache[layer_id, phys_pages, offset : offset + chunk] = v[
|
|
||||||
:, written : written + chunk
|
|
||||||
]
|
|
||||||
written += chunk
|
|
||||||
|
|
||||||
def gather(
|
def get_value_buffer(self, layer_id: int) -> Tensor:
|
||||||
self, layer_id: int, page_table: Tensor, total_len: int
|
return self.v_buffer[layer_id]
|
||||||
) -> Tuple[Tensor, Tensor]:
|
|
||||||
safe = page_table.clamp(min=0)
|
|
||||||
k = self.k_cache[layer_id, safe]
|
|
||||||
v = self.v_cache[layer_id, safe]
|
|
||||||
k = k.flatten(1, 2)
|
|
||||||
v = v.flatten(1, 2)
|
|
||||||
if (page_table < 0).any():
|
|
||||||
invalid = (
|
|
||||||
(page_table < 0)
|
|
||||||
.unsqueeze(-1)
|
|
||||||
.expand(-1, -1, self.page_size)
|
|
||||||
.flatten(1, 2)
|
|
||||||
)
|
|
||||||
invalid = invalid[:, :, None, None].expand_as(k)
|
|
||||||
k = k.masked_fill(invalid, 0.0)
|
|
||||||
v = v.masked_fill(invalid, 0.0)
|
|
||||||
k = k[:, :total_len]
|
|
||||||
v = v[:, :total_len]
|
|
||||||
return k, v
|
|
||||||
|
|
||||||
|
def set_kv_buffer(self, layer_id: int, loc: Tensor, k: Tensor, v: Tensor) -> None:
|
||||||
class KvcacheView:
|
self.k_buffer[layer_id, loc] = k
|
||||||
"""Bundles Storage + page_table + total_len for attention layers."""
|
self.v_buffer[layer_id, loc] = v
|
||||||
|
|
||||||
def __init__(self, storage: Storage, page_table: Tensor, total_len: int = 0):
|
|
||||||
self._storage = storage
|
|
||||||
self._page_table = page_table
|
|
||||||
self._total_len = total_len
|
|
||||||
|
|
||||||
def write(self, layer_id: int, k: Tensor, v: Tensor) -> None:
|
|
||||||
start_pos = self._total_len - k.size(1)
|
|
||||||
self._storage.write(layer_id, self._page_table, start_pos, k, v)
|
|
||||||
|
|
||||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
|
||||||
return self._storage.gather(layer_id, self._page_table, self._total_len)
|
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
class KVCache:
|
class KVCache:
|
||||||
"""Facade: page management + KV-cache I/O for continuous batching."""
|
"""Pure data struct passed to model for KV cache I/O.
|
||||||
|
|
||||||
|
The attention layer does raw buffer indexing — no methods, no abstraction.
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
k_buffer: [n_layers, size, n_kv_heads, head_dim]
|
||||||
|
v_buffer: [n_layers, size, n_kv_heads, head_dim]
|
||||||
|
req_to_token: [num_reqs, max_ctx_len] — index table
|
||||||
|
req_pool_indices: [batch_size] — row indices into req_to_token
|
||||||
|
seq_lens: [batch_size] — per-request total sequence lengths
|
||||||
|
out_cache_loc: [batch, new_seq_len] or [batch, 1] — write indices
|
||||||
|
max_len: max(seq_lens) as Python int — avoids GPU sync in decode
|
||||||
|
kv_indptr: [batch+1] int32 — prefix sum of seq_lens, precomputed once
|
||||||
|
per step so the attention backend avoids rebuilding it per layer.
|
||||||
|
"""
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
class PagePool:
|
||||||
|
"""Top-level KV cache manager.
|
||||||
|
|
||||||
|
Combines KVStorage + ReqToTokenPool + Allocator + PrefixCache.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
n_layers: Number of transformer layers.
|
||||||
|
n_kv_heads: Number of KV attention heads.
|
||||||
|
head_dim: Dimension per head.
|
||||||
|
max_batch_size: Maximum concurrent requests.
|
||||||
|
max_seq_len: Maximum sequence length per request.
|
||||||
|
device, dtype: Tensor device and dtype.
|
||||||
|
page_size: Page size for paged mode (1 = token-level).
|
||||||
|
n_tokens: Total token slots for paged mode. None = contiguous mode
|
||||||
|
(pre-allocates max_batch_size * max_seq_len).
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
n_layers: int,
|
n_layers: int,
|
||||||
n_pages: int,
|
|
||||||
page_size: int,
|
|
||||||
n_kv_heads: int,
|
n_kv_heads: int,
|
||||||
head_dim: int,
|
head_dim: int,
|
||||||
|
max_batch_size: int,
|
||||||
|
max_seq_len: int,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
dtype: torch.dtype,
|
dtype: torch.dtype,
|
||||||
|
page_size: int = 1,
|
||||||
|
n_tokens: Optional[int] = None,
|
||||||
):
|
):
|
||||||
self.page_size = page_size
|
self.page_size = page_size
|
||||||
self._pool = PagePool(Allocator(n_pages), PrefixCache(page_size))
|
self.max_batch_size = max_batch_size
|
||||||
self._table = TaskTable(page_size)
|
self.max_seq_len = max_seq_len
|
||||||
self._storage = Storage(
|
self.device = device
|
||||||
n_layers, n_pages, page_size, n_kv_heads, head_dim, device, dtype
|
self.dtype = dtype
|
||||||
|
self.n_layers = n_layers
|
||||||
|
self.n_kv_heads = n_kv_heads
|
||||||
|
self.head_dim = head_dim
|
||||||
|
|
||||||
|
self.contiguous = n_tokens is None
|
||||||
|
if self.contiguous:
|
||||||
|
self.n_tokens = max_batch_size * max_seq_len
|
||||||
|
else:
|
||||||
|
self.n_tokens = n_tokens
|
||||||
|
|
||||||
|
self._storage = KVStorage(
|
||||||
|
self.n_tokens, n_layers, n_kv_heads, head_dim, device, dtype
|
||||||
)
|
)
|
||||||
|
self._req_pool = ReqToTokenPool(max_batch_size, max_seq_len, device)
|
||||||
|
|
||||||
|
if self.contiguous:
|
||||||
|
for i in range(max_batch_size):
|
||||||
|
self._req_pool.req_to_token[i] = torch.arange(
|
||||||
|
i * max_seq_len, (i + 1) * max_seq_len, device=device
|
||||||
|
)
|
||||||
|
self._alloc: Optional[Allocator] = None
|
||||||
|
self._prefix: Optional[PrefixCache] = None
|
||||||
|
else:
|
||||||
|
n_pages = self.n_tokens // page_size
|
||||||
|
self._alloc = Allocator(n_pages)
|
||||||
|
self._prefix = PrefixCache(page_size) if page_size > 1 else None
|
||||||
|
if self._prefix is not None:
|
||||||
|
self._alloc.on_evict = self._prefix.evict
|
||||||
|
|
||||||
|
self._task_req: Dict[str, int] = {}
|
||||||
|
self._task_len: Dict[int, int] = {}
|
||||||
|
self._task_cached: Dict[str, int] = {}
|
||||||
|
self._task_slots: Dict[str, List[int]] = {}
|
||||||
|
self._task_pages: Dict[str, List[int]] = {}
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
# ---- task lifecycle ----
|
||||||
|
|
||||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
||||||
hits = self._pool.lookup(prompt_ids)
|
req_slots = self._req_pool.alloc(1)
|
||||||
cached = len(hits) * self.page_size
|
if req_slots is None:
|
||||||
for p in hits:
|
return False
|
||||||
self._pool.inc_ref(p)
|
req_idx = req_slots[0]
|
||||||
|
self._task_req[task_id] = req_idx
|
||||||
|
|
||||||
remaining = len(prompt_ids) - cached
|
if self.contiguous:
|
||||||
n_new = (
|
self._task_len[req_idx] = len(prompt_ids)
|
||||||
(remaining + self.page_size - 1) // self.page_size if remaining > 0 else 0
|
self._task_cached[task_id] = 0
|
||||||
)
|
return True
|
||||||
new_pages: List[int] = []
|
|
||||||
if n_new > 0:
|
n_tokens_needed = len(prompt_ids)
|
||||||
for _ in range(n_new):
|
cached = 0
|
||||||
p = self._pool.alloc()
|
|
||||||
if p < 0:
|
if self._prefix is not None:
|
||||||
for hp in hits:
|
hits = self._prefix.lookup(prompt_ids)
|
||||||
self._pool.free(hp)
|
cached = len(hits) * self.page_size
|
||||||
for np in new_pages:
|
for p in hits:
|
||||||
self._pool.free(np)
|
self._alloc.inc_ref(p)
|
||||||
|
self._task_pages[task_id] = list(hits)
|
||||||
|
self._task_slots[task_id] = []
|
||||||
|
else:
|
||||||
|
self._task_pages[task_id] = []
|
||||||
|
self._task_slots[task_id] = []
|
||||||
|
|
||||||
|
remaining = n_tokens_needed - cached
|
||||||
|
if remaining > 0:
|
||||||
|
if self.page_size == 1:
|
||||||
|
slots = self._alloc_tokens(remaining)
|
||||||
|
if slots is None:
|
||||||
|
for p in self._task_pages[task_id]:
|
||||||
|
self._alloc.free(p)
|
||||||
|
self._req_pool.free([req_idx])
|
||||||
|
del self._task_req[task_id]
|
||||||
return False
|
return False
|
||||||
new_pages.append(p)
|
self._task_slots[task_id] = slots
|
||||||
|
else:
|
||||||
|
n_new_pages = (remaining + self.page_size - 1) // self.page_size
|
||||||
|
new_pages = []
|
||||||
|
for _ in range(n_new_pages):
|
||||||
|
p = self._alloc.alloc()
|
||||||
|
if p < 0:
|
||||||
|
for hp in self._task_pages[task_id]:
|
||||||
|
self._alloc.free(hp)
|
||||||
|
for np_ in new_pages:
|
||||||
|
self._alloc.free(np_)
|
||||||
|
self._req_pool.free([req_idx])
|
||||||
|
del self._task_req[task_id]
|
||||||
|
return False
|
||||||
|
new_pages.append(p)
|
||||||
|
self._task_pages[task_id].extend(new_pages)
|
||||||
|
|
||||||
self._table.set(task_id, hits + new_pages, cached)
|
self._write_req_to_token(task_id, prompt_ids, cached)
|
||||||
|
self._task_len[req_idx] = len(prompt_ids)
|
||||||
|
self._task_cached[task_id] = cached
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def task_free(self, task_id: str) -> None:
|
def task_free(self, task_id: str):
|
||||||
page_table, _ = self._table.pop(task_id)
|
req_idx = self._task_req.pop(task_id, None)
|
||||||
for idx in page_table:
|
if req_idx is None:
|
||||||
self._pool.free(idx)
|
return
|
||||||
|
self._task_len.pop(req_idx, None)
|
||||||
|
self._task_cached.pop(task_id, None)
|
||||||
|
|
||||||
|
if not self.contiguous:
|
||||||
|
if self._prefix is not None:
|
||||||
|
for p in self._task_pages.get(task_id, []):
|
||||||
|
keep = self._prefix.has_page(p)
|
||||||
|
self._alloc.free(p, keep_cached=keep)
|
||||||
|
if not keep:
|
||||||
|
self._prefix.evict(p)
|
||||||
|
else:
|
||||||
|
for p in self._task_pages.get(task_id, []):
|
||||||
|
self._alloc.free(p)
|
||||||
|
self._task_pages.pop(task_id, None)
|
||||||
|
self._task_slots.pop(task_id, None)
|
||||||
|
|
||||||
|
self._req_pool.free([req_idx])
|
||||||
|
|
||||||
def task_extend(self, task_id: str, pos: int) -> bool:
|
def task_extend(self, task_id: str, pos: int) -> bool:
|
||||||
page_table = self._table.get(task_id)
|
req_idx = self._task_req.get(task_id)
|
||||||
needed = (pos + 1 + self.page_size - 1) // self.page_size
|
if req_idx is None:
|
||||||
while len(page_table) < needed:
|
return False
|
||||||
p = self._pool.alloc()
|
|
||||||
if p < 0:
|
if self.contiguous:
|
||||||
|
return pos < self.max_seq_len
|
||||||
|
|
||||||
|
if self.page_size == 1:
|
||||||
|
slots = self._alloc_tokens(1)
|
||||||
|
if slots is None:
|
||||||
return False
|
return False
|
||||||
page_table.append(p)
|
self._task_slots.setdefault(task_id, []).extend(slots)
|
||||||
|
self._req_pool.req_to_token[req_idx, pos] = slots[0]
|
||||||
|
else:
|
||||||
|
page_idx = pos // self.page_size
|
||||||
|
existing = self._task_pages.get(task_id, [])
|
||||||
|
if page_idx >= len(existing):
|
||||||
|
p = self._alloc.alloc()
|
||||||
|
if p < 0:
|
||||||
|
return False
|
||||||
|
existing.append(p)
|
||||||
|
self._task_pages[task_id] = existing
|
||||||
|
page_offset = pos % self.page_size
|
||||||
|
page = existing[page_idx]
|
||||||
|
token_slot = page * self.page_size + page_offset
|
||||||
|
self._req_pool.req_to_token[req_idx, pos] = token_slot
|
||||||
|
|
||||||
|
self._task_len[req_idx] = pos + 1
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def task_cached(self, task_id: str) -> int:
|
def task_cached(self, task_id: str) -> int:
|
||||||
return self._table.get_cached(task_id)
|
return self._task_cached.get(task_id, 0)
|
||||||
|
|
||||||
def task_record_hashes(
|
def task_record_hashes(
|
||||||
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
||||||
) -> None:
|
):
|
||||||
page_table = self._table.get(task_id)
|
if self._prefix is None or self.contiguous:
|
||||||
|
return
|
||||||
|
pages = self._task_pages.get(task_id, [])
|
||||||
full_pages = len(prompt_ids) // self.page_size
|
full_pages = len(prompt_ids) // self.page_size
|
||||||
for i in range(start_logical_page, full_pages):
|
for i in range(start_logical_page, min(full_pages, len(pages))):
|
||||||
self._pool.record(page_table[i], prompt_ids, i)
|
self._prefix.record(pages[i], prompt_ids, i)
|
||||||
|
|
||||||
def make_table_tensor(self, task_ids: List[str], device: torch.device) -> Tensor:
|
# ---- bind for forward ----
|
||||||
return self._table.table_tensor(task_ids, device)
|
|
||||||
|
|
||||||
def bind(self, page_table: Tensor, total_len: int = 0) -> KvcacheView:
|
def bind_tasks(
|
||||||
return KvcacheView(self._storage, page_table, total_len)
|
self,
|
||||||
|
task_ids: List[str],
|
||||||
|
seq_lens: List[int],
|
||||||
|
device: torch.device,
|
||||||
|
start_pos: Optional[int] = None,
|
||||||
|
) -> KVCache:
|
||||||
|
req_indices = [self._task_req[tid] for tid in task_ids]
|
||||||
|
req_pool_indices = torch.tensor(req_indices, dtype=torch.long, device=device)
|
||||||
|
seq_lens_t = torch.tensor(seq_lens, dtype=torch.long, device=device)
|
||||||
|
|
||||||
|
if start_pos is not None:
|
||||||
|
seq_len = seq_lens[0]
|
||||||
|
out_cache_loc = self._req_pool.req_to_token[
|
||||||
|
req_pool_indices, start_pos:seq_len
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
write_pos = seq_lens_t - 1
|
||||||
|
out_cache_loc = self._req_pool.req_to_token[
|
||||||
|
req_pool_indices, write_pos
|
||||||
|
].unsqueeze(-1)
|
||||||
|
|
||||||
|
kv_indptr = torch.zeros(len(seq_lens) + 1, dtype=torch.int32, device=device)
|
||||||
|
kv_indptr[1:] = seq_lens_t.cumsum(0).to(torch.int32)
|
||||||
|
|
||||||
|
return KVCache(
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
# ---- internals ----
|
||||||
|
|
||||||
|
def _alloc_tokens(self, n: int) -> Optional[List[int]]:
|
||||||
|
if self.page_size != 1:
|
||||||
|
raise RuntimeError("_alloc_tokens is for page_size=1 only")
|
||||||
|
slots = []
|
||||||
|
for _ in range(n):
|
||||||
|
p = self._alloc.alloc()
|
||||||
|
if p < 0:
|
||||||
|
for s in slots:
|
||||||
|
self._alloc.free(s)
|
||||||
|
return None
|
||||||
|
slots.append(p)
|
||||||
|
return slots
|
||||||
|
|
||||||
|
def _write_req_to_token(self, task_id: str, prompt_ids: List[int], cached: int):
|
||||||
|
req_idx = self._task_req[task_id]
|
||||||
|
total = len(prompt_ids)
|
||||||
|
|
||||||
|
if self.contiguous:
|
||||||
|
return
|
||||||
|
|
||||||
|
if self.page_size == 1:
|
||||||
|
slots = self._task_slots.get(task_id, [])
|
||||||
|
all_slots = slots[: total - cached]
|
||||||
|
if all_slots:
|
||||||
|
self._req_pool.req_to_token[req_idx, cached:total] = torch.tensor(
|
||||||
|
all_slots, dtype=torch.long, device=self.device
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
pages = self._task_pages.get(task_id, [])
|
||||||
|
for pos in range(cached, total):
|
||||||
|
page_idx = pos // self.page_size
|
||||||
|
page_offset = pos % self.page_size
|
||||||
|
if page_idx < len(pages):
|
||||||
|
token_slot = pages[page_idx] * self.page_size + page_offset
|
||||||
|
self._req_pool.req_to_token[req_idx, pos] = token_slot
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ from typing import List, Optional
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.inference.core.cache import KVCache
|
from astrai.inference.core.cache import PagePool
|
||||||
from astrai.inference.core.task import Task
|
from astrai.inference.core.task import Task
|
||||||
from astrai.inference.sample import sample
|
from astrai.inference.sample import sample
|
||||||
from astrai.model.automodel import AutoModel
|
from astrai.model.automodel import AutoModel
|
||||||
@@ -19,19 +19,17 @@ class Executor:
|
|||||||
self,
|
self,
|
||||||
model: AutoModel,
|
model: AutoModel,
|
||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
page_cache: KVCache,
|
kv_cache: PagePool,
|
||||||
device: Optional[str] = None,
|
device: Optional[str] = None,
|
||||||
dtype: Optional[torch.dtype] = None,
|
dtype: Optional[torch.dtype] = None,
|
||||||
):
|
):
|
||||||
self.model = model
|
self.model = model
|
||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
self.page_cache = page_cache
|
self.kv_cache = kv_cache
|
||||||
self.device = device or next(model.parameters()).device
|
self.device = device or next(model.parameters()).device
|
||||||
self.dtype = dtype or next(model.parameters()).dtype
|
self.dtype = dtype or next(model.parameters()).dtype
|
||||||
|
|
||||||
def execute_prefill(
|
def execute_prefill(self, tasks: List[Task], prompt_len: int, start_pos: int = 0):
|
||||||
self, tasks: List[Task], prompt_len: int, start_pos: int = 0
|
|
||||||
) -> None:
|
|
||||||
if start_pos >= prompt_len:
|
if start_pos >= prompt_len:
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -45,20 +43,42 @@ class Executor:
|
|||||||
)
|
)
|
||||||
|
|
||||||
task_ids = [t.task_id for t in tasks]
|
task_ids = [t.task_id for t in tasks]
|
||||||
page_tables = self.page_cache.make_table_tensor(task_ids, self.device)
|
position_ids = (
|
||||||
|
torch.arange(start_pos, prompt_len, dtype=torch.long, device=self.device)
|
||||||
|
.unsqueeze(0)
|
||||||
|
.expand(batch_sz, -1)
|
||||||
|
)
|
||||||
|
input_mask = position_ids.unsqueeze(-1) >= torch.arange(
|
||||||
|
prompt_len, device=self.device
|
||||||
|
)
|
||||||
|
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
self.model(
|
self.model(
|
||||||
input_ids,
|
input_ids,
|
||||||
position_ids=torch.arange(
|
input_mask=input_mask,
|
||||||
start_pos, prompt_len, dtype=torch.long, device=self.device
|
position_ids=position_ids,
|
||||||
)
|
kv_cache=self.kv_cache.bind_tasks(
|
||||||
.unsqueeze(0)
|
task_ids, [prompt_len] * batch_sz, self.device, start_pos=start_pos
|
||||||
.expand(batch_sz, -1),
|
),
|
||||||
paged_cache=self.page_cache.bind(page_tables, total_len=prompt_len),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def execute_decode(self, tasks: List[Task]) -> List[int]:
|
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:
|
if not tasks:
|
||||||
return []
|
return []
|
||||||
|
|
||||||
@@ -71,26 +91,84 @@ class Executor:
|
|||||||
position_ids = torch.tensor(
|
position_ids = torch.tensor(
|
||||||
[t.next_pos for t in tasks], dtype=torch.long, device=self.device
|
[t.next_pos for t in tasks], dtype=torch.long, device=self.device
|
||||||
)
|
)
|
||||||
total_len = position_ids.max().item() + 1
|
total_len = max(t.next_pos for t in tasks) + 1
|
||||||
|
input_mask = position_ids[:, None, None] >= torch.arange(
|
||||||
|
total_len, device=self.device
|
||||||
|
)
|
||||||
|
|
||||||
task_ids = [t.task_id for t in tasks]
|
task_ids = [t.task_id for t in tasks]
|
||||||
page_tables = self.page_cache.make_table_tensor(task_ids, self.device)
|
|
||||||
|
|
||||||
temperatures = torch.tensor([t.temperature for t in tasks], device=self.device)
|
temperatures = torch.tensor([t.temperature for t in tasks], device=self.device)
|
||||||
top_ks = torch.tensor([t.top_k for t in tasks], device=self.device)
|
top_ks = torch.tensor([t.top_k for t in tasks], device=self.device)
|
||||||
top_ps = torch.tensor([t.top_p for t in tasks], device=self.device)
|
top_ps = torch.tensor([t.top_p for t in tasks], device=self.device)
|
||||||
|
freq_penalties = torch.tensor(
|
||||||
|
[t.frequency_penalty for t in tasks], device=self.device
|
||||||
|
)
|
||||||
|
|
||||||
|
has_freq = bool((freq_penalties != 0).any())
|
||||||
|
if has_freq:
|
||||||
|
history_lists = []
|
||||||
|
history_lens = []
|
||||||
|
for t in tasks:
|
||||||
|
window = t.rep_window
|
||||||
|
prompt_part = t.prompt_ids[-window:]
|
||||||
|
ids = prompt_part + t.output_ids
|
||||||
|
history_lists.append(ids)
|
||||||
|
history_lens.append(len(ids))
|
||||||
|
|
||||||
|
max_len = max(history_lens) if history_lens else 0
|
||||||
|
padded_ids = torch.zeros(
|
||||||
|
len(tasks), max_len, dtype=torch.long, device=self.device
|
||||||
|
)
|
||||||
|
padded_mask = torch.zeros(
|
||||||
|
len(tasks), max_len, dtype=torch.bool, device=self.device
|
||||||
|
)
|
||||||
|
for i, h in enumerate(history_lists):
|
||||||
|
L = history_lens[i]
|
||||||
|
padded_ids[i, :L] = torch.as_tensor(
|
||||||
|
h, dtype=torch.long, device=self.device
|
||||||
|
)
|
||||||
|
padded_mask[i, :L] = True
|
||||||
|
else:
|
||||||
|
padded_ids = None
|
||||||
|
padded_mask = None
|
||||||
|
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
outputs = self.model(
|
outputs = self.model(
|
||||||
input_ids.unsqueeze(1),
|
input_ids.unsqueeze(1),
|
||||||
paged_cache=self.page_cache.bind(page_tables, total_len=total_len),
|
input_mask=input_mask,
|
||||||
|
kv_cache=self.kv_cache.bind_tasks(
|
||||||
|
task_ids,
|
||||||
|
[t.next_pos + 1 for t in tasks],
|
||||||
|
self.device,
|
||||||
|
),
|
||||||
position_ids=position_ids.unsqueeze(1),
|
position_ids=position_ids.unsqueeze(1),
|
||||||
)
|
)
|
||||||
logits = outputs["logits"][:, -1, :]
|
logits = outputs["logits"][:, -1, :]
|
||||||
|
|
||||||
|
if return_logprobs:
|
||||||
|
tokens, logprobs = sample(
|
||||||
|
logits,
|
||||||
|
temperature=temperatures,
|
||||||
|
top_k=top_ks,
|
||||||
|
top_p=top_ps,
|
||||||
|
frequency_penalty=freq_penalties,
|
||||||
|
input_ids=padded_ids,
|
||||||
|
input_mask=padded_mask,
|
||||||
|
return_logprobs=True,
|
||||||
|
)
|
||||||
|
tokens_list = tokens.tolist()
|
||||||
|
logprobs_list = logprobs.tolist()
|
||||||
|
for t, lp in zip(tasks, logprobs_list):
|
||||||
|
t.output_logprobs.append(float(lp))
|
||||||
|
return list(zip(tokens_list, logprobs_list))
|
||||||
|
|
||||||
return sample(
|
return sample(
|
||||||
logits,
|
logits,
|
||||||
temperature=temperatures,
|
temperature=temperatures,
|
||||||
top_k=top_ks,
|
top_k=top_ks,
|
||||||
top_p=top_ps,
|
top_p=top_ps,
|
||||||
|
frequency_penalty=freq_penalties,
|
||||||
|
input_ids=padded_ids,
|
||||||
|
input_mask=padded_mask,
|
||||||
).tolist()
|
).tolist()
|
||||||
|
|||||||
@@ -1,10 +1,11 @@
|
|||||||
import logging
|
import logging
|
||||||
import threading
|
import threading
|
||||||
|
import uuid
|
||||||
from typing import Any, Dict, List, Optional, Tuple
|
from typing import Any, Dict, List, Optional, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.inference.core.cache import KVCache
|
from astrai.inference.core.cache import PagePool
|
||||||
from astrai.inference.core.executor import Executor
|
from astrai.inference.core.executor import Executor
|
||||||
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
||||||
from astrai.model.automodel import AutoModel
|
from astrai.model.automodel import AutoModel
|
||||||
@@ -14,7 +15,7 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
|
|
||||||
class InferenceScheduler:
|
class InferenceScheduler:
|
||||||
"""Four-phase continuous batching loop: cleanup -> refill -> prefill -> decode."""
|
"""Continuous batching loop: cleanup -> refill -> prefill -> decode (all groups)."""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -22,65 +23,74 @@ class InferenceScheduler:
|
|||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
max_batch_size: int = 16,
|
max_batch_size: int = 16,
|
||||||
max_seq_len: Optional[int] = None,
|
max_seq_len: Optional[int] = None,
|
||||||
max_prompt_len: int = 512,
|
|
||||||
page_size: int = 64,
|
|
||||||
device: Optional[str] = None,
|
device: Optional[str] = None,
|
||||||
dtype: Optional[torch.dtype] = None,
|
dtype: Optional[torch.dtype] = None,
|
||||||
|
cache: Optional[PagePool] = None,
|
||||||
):
|
):
|
||||||
config = model.config
|
config = model.config
|
||||||
|
|
||||||
self.max_seq_len = max_seq_len or config.max_len
|
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.device = device or next(model.parameters()).device
|
||||||
self.dtype = dtype or next(model.parameters()).dtype
|
self.dtype = dtype or next(model.parameters()).dtype
|
||||||
|
|
||||||
n_pages = (
|
head_dim = config.hidden_size // config.num_attention_heads
|
||||||
max_batch_size * (self.max_seq_len + page_size) + page_size - 1
|
|
||||||
) // page_size
|
|
||||||
|
|
||||||
self._page_cache = KVCache(
|
if cache is not None:
|
||||||
config.n_layers,
|
self._cache = cache
|
||||||
n_pages,
|
else:
|
||||||
page_size,
|
self._cache = PagePool(
|
||||||
config.n_kv_heads,
|
n_layers=config.num_hidden_layers,
|
||||||
config.dim // config.n_heads,
|
n_kv_heads=config.num_key_value_heads,
|
||||||
self.device,
|
head_dim=head_dim,
|
||||||
self.dtype,
|
max_batch_size=max_batch_size,
|
||||||
)
|
max_seq_len=self.max_seq_len,
|
||||||
|
device=self.device,
|
||||||
|
dtype=self.dtype,
|
||||||
|
)
|
||||||
|
|
||||||
self._task_mgr = TaskManager(
|
self._task_mgr = TaskManager(
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
max_batch_size=max_batch_size,
|
max_batch_size=max_batch_size,
|
||||||
max_seq_len=self.max_seq_len,
|
max_seq_len=self.max_seq_len,
|
||||||
max_prompt_len=max_prompt_len,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
self._executor = Executor(
|
self._executor = Executor(
|
||||||
model=model,
|
model=model,
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
page_cache=self._page_cache,
|
kv_cache=self._cache,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
dtype=self.dtype,
|
dtype=self.dtype,
|
||||||
)
|
)
|
||||||
|
|
||||||
self._running = False
|
self._stop_event = threading.Event()
|
||||||
|
self._loop_thread: Optional[threading.Thread] = None
|
||||||
|
|
||||||
def add_task(self, prompt: str, **kwargs) -> str:
|
def add_task(self, prompt: str, **kwargs) -> str:
|
||||||
return self._task_mgr.add_task(prompt, **kwargs)
|
return self._task_mgr.add_task(prompt, **kwargs)
|
||||||
|
|
||||||
def remove_task(self, task_id: str) -> None:
|
def remove_task(self, task_id: str):
|
||||||
for task in self._task_mgr.remove_task(task_id):
|
for task in self._task_mgr.remove_task(task_id):
|
||||||
self._page_cache.task_free(task.task_id)
|
self._cache.task_free(task.task_id)
|
||||||
|
|
||||||
def get_stats(self) -> Dict[str, Any]:
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
return self._task_mgr.get_stats()
|
return self._task_mgr.get_stats()
|
||||||
|
|
||||||
def _run_generation_loop(self) -> None:
|
def _run_generation_loop(self):
|
||||||
stop_ids = self._task_mgr.tokenizer.stop_ids
|
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||||
|
cache = self._cache
|
||||||
try:
|
try:
|
||||||
while self._running:
|
while not self._stop_event.is_set():
|
||||||
finished = self._task_mgr.remove_finished_tasks(stop_ids)
|
finished = self._task_mgr.remove_finished_tasks(stop_ids)
|
||||||
for task in finished:
|
for task in finished:
|
||||||
self._page_cache.task_free(task.task_id)
|
cache.task_free(task.task_id)
|
||||||
|
|
||||||
active = self._task_mgr.get_active_tasks()
|
active = self._task_mgr.get_active_tasks()
|
||||||
available = self._task_mgr.max_batch_size - len(active)
|
available = self._task_mgr.max_batch_size - len(active)
|
||||||
@@ -88,7 +98,7 @@ class InferenceScheduler:
|
|||||||
candidates = self._task_mgr.pull_candidates(available)
|
candidates = self._task_mgr.pull_candidates(available)
|
||||||
failed = []
|
failed = []
|
||||||
for task in candidates:
|
for task in candidates:
|
||||||
if self._page_cache.task_alloc(task.task_id, task.prompt_ids):
|
if cache.task_alloc(task.task_id, task.prompt_ids):
|
||||||
self._task_mgr.activate(task)
|
self._task_mgr.activate(task)
|
||||||
else:
|
else:
|
||||||
failed.append(task)
|
failed.append(task)
|
||||||
@@ -99,8 +109,13 @@ class InferenceScheduler:
|
|||||||
self._task_mgr.wait_for_tasks(timeout=1.0)
|
self._task_mgr.wait_for_tasks(timeout=1.0)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
active = self._task_mgr.get_active_tasks()
|
||||||
|
|
||||||
to_prefill = [
|
to_prefill = [
|
||||||
t for t in self._task_mgr.get_active_tasks() if t.output_tokens == 0
|
t
|
||||||
|
for t in active
|
||||||
|
if t.output_tokens == 0
|
||||||
|
and cache.task_cached(t.task_id) < len(t.prompt_ids)
|
||||||
]
|
]
|
||||||
if to_prefill:
|
if to_prefill:
|
||||||
for t in to_prefill:
|
for t in to_prefill:
|
||||||
@@ -110,78 +125,187 @@ class InferenceScheduler:
|
|||||||
for t in to_prefill:
|
for t in to_prefill:
|
||||||
key = (
|
key = (
|
||||||
len(t.prompt_ids),
|
len(t.prompt_ids),
|
||||||
self._page_cache.task_cached(t.task_id),
|
cache.task_cached(t.task_id),
|
||||||
)
|
)
|
||||||
groups.setdefault(key, []).append(t)
|
groups.setdefault(key, []).append(t)
|
||||||
|
|
||||||
for (prompt_len, start_pos), group in groups.items():
|
for (prompt_len, start_pos), group in groups.items():
|
||||||
self._executor.execute_prefill(group, prompt_len, start_pos)
|
self._executor.execute_prefill(group, prompt_len, start_pos)
|
||||||
start_logical_page = start_pos // self._page_cache.page_size
|
start_logical_page = start_pos // getattr(
|
||||||
|
cache, "page_size", 64
|
||||||
|
)
|
||||||
for t in group:
|
for t in group:
|
||||||
self._page_cache.task_record_hashes(
|
cache.task_record_hashes(
|
||||||
t.task_id,
|
t.task_id, t.prompt_ids, start_logical_page
|
||||||
t.prompt_ids,
|
|
||||||
start_logical_page=start_logical_page,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
pos_groups: Dict[int, List[Task]] = {}
|
decode_tasks = active
|
||||||
for t in self._task_mgr.get_active_tasks():
|
|
||||||
pos_groups.setdefault(t.next_pos, []).append(t)
|
|
||||||
|
|
||||||
if pos_groups:
|
valid: List[Task] = []
|
||||||
best_key = max(pos_groups, key=lambda k: len(pos_groups[k]))
|
for t in decode_tasks:
|
||||||
group = sorted(pos_groups[best_key], key=lambda t: t.task_id)
|
if cache.task_extend(t.task_id, t.next_pos):
|
||||||
|
valid.append(t)
|
||||||
|
else:
|
||||||
|
t.status = TaskStatus.ABORTED
|
||||||
|
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||||
|
|
||||||
valid: List[Task] = []
|
if valid:
|
||||||
for t in group:
|
next_tokens = self._executor.execute_decode(valid)
|
||||||
if self._page_cache.task_extend(t.task_id, t.next_pos):
|
|
||||||
valid.append(t)
|
|
||||||
else:
|
|
||||||
t.status = TaskStatus.ABORTED
|
|
||||||
if t.stream_callback:
|
|
||||||
t.stream_callback(STOP)
|
|
||||||
|
|
||||||
if valid:
|
for t, ntok in zip(valid, next_tokens):
|
||||||
next_tokens = self._executor.execute_decode(valid)
|
t.output_ids.append(ntok)
|
||||||
|
t.output_tokens += 1
|
||||||
|
new_text = t.decode_new_token(self._task_mgr.tokenizer)
|
||||||
|
if new_text:
|
||||||
|
self._task_mgr.invoke_callback(t.task_id, new_text)
|
||||||
|
|
||||||
for t, ntok in zip(valid, next_tokens):
|
for t in valid:
|
||||||
t.output_ids.append(ntok)
|
if t.is_finished(stop_ids):
|
||||||
t.output_tokens += 1
|
remaining = t.flush_remaining(self._task_mgr.tokenizer)
|
||||||
pos = t.input_tokens + t.output_tokens
|
if remaining:
|
||||||
self._page_cache.task_extend(t.task_id, pos)
|
self._task_mgr.invoke_callback(t.task_id, remaining)
|
||||||
if t.stream_callback:
|
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||||
t.stream_callback(
|
|
||||||
self._task_mgr.tokenizer.decode([ntok])
|
|
||||||
)
|
|
||||||
|
|
||||||
for t in valid:
|
|
||||||
if t.is_finished(stop_ids):
|
|
||||||
if t.stream_callback:
|
|
||||||
t.stream_callback(STOP)
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
self._stop_event.set()
|
||||||
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
||||||
for task in self._task_mgr.get_active_tasks():
|
for task in self._task_mgr.get_active_tasks():
|
||||||
if task.stream_callback:
|
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||||
task.stream_callback(STOP)
|
cache.task_free(task.task_id)
|
||||||
self._page_cache.task_free(task.task_id)
|
for task in self._task_mgr.get_waiting_tasks():
|
||||||
|
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||||
self._task_mgr.clear_queues()
|
self._task_mgr.clear_queues()
|
||||||
raise
|
|
||||||
|
|
||||||
def start(self) -> None:
|
def start(self):
|
||||||
if not self._running:
|
if self._loop_thread is not None and self._loop_thread.is_alive():
|
||||||
self._running = True
|
return
|
||||||
t = threading.Thread(target=self._run_generation_loop, daemon=True)
|
self._stop_event.clear()
|
||||||
t.start()
|
t = threading.Thread(target=self._run_generation_loop, daemon=True)
|
||||||
self._loop_thread = t
|
t.start()
|
||||||
|
self._loop_thread = t
|
||||||
|
|
||||||
def stop(self) -> None:
|
def stop(self):
|
||||||
self._running = False
|
self._stop_event.set()
|
||||||
self._task_mgr.wake()
|
self._task_mgr.wake()
|
||||||
if hasattr(self, "_loop_thread"):
|
if self._loop_thread is not None:
|
||||||
self._loop_thread.join(timeout=2.0)
|
self._loop_thread.join(timeout=2.0)
|
||||||
|
self._loop_thread = None
|
||||||
for task in self._task_mgr.get_active_tasks():
|
for task in self._task_mgr.get_active_tasks():
|
||||||
self._page_cache.task_free(task.task_id)
|
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||||
|
self._cache.task_free(task.task_id)
|
||||||
|
for task in self._task_mgr.get_waiting_tasks():
|
||||||
|
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||||
|
self._cache.task_free(task.task_id)
|
||||||
self._task_mgr.clear_queues()
|
self._task_mgr.clear_queues()
|
||||||
if torch.cuda.is_available():
|
if torch.cuda.is_available():
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
def run_batch(
|
||||||
|
self,
|
||||||
|
prompt_ids_list: List[List[int]],
|
||||||
|
*,
|
||||||
|
max_tokens: Optional[int] = None,
|
||||||
|
temperature: float = 1.0,
|
||||||
|
top_p: float = 1.0,
|
||||||
|
top_k: int = 50,
|
||||||
|
frequency_penalty: float = 0.0,
|
||||||
|
rep_window: int = 64,
|
||||||
|
return_logprobs: bool = False,
|
||||||
|
) -> List[List[int]]:
|
||||||
|
"""Synchronous batch generation without the scheduler thread.
|
||||||
|
|
||||||
|
Accepts already-tokenized prompts (no string round-trip) and runs
|
||||||
|
prefill + decode to completion on the calling thread. Designed for
|
||||||
|
RL rollout, where logprobs of the behaviour policy must be collected
|
||||||
|
alongside generated tokens.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prompt_ids_list: ``B`` prompts, each a list of token IDs.
|
||||||
|
max_tokens: Maximum tokens to generate per prompt. ``None``
|
||||||
|
uses ``self.max_seq_len - len(prompt_ids)``.
|
||||||
|
temperature/top_p/top_k/frequency_penalty/rep_window: Sampling
|
||||||
|
parameters (uniform across the batch).
|
||||||
|
return_logprobs: If ``True``, return ``(token_ids, logprobs)``
|
||||||
|
tuples per prompt (logprobs aligned 1-to-1 with token_ids).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
``List[List[int]]`` of generated token IDs per prompt, or —
|
||||||
|
when ``return_logprobs`` is ``True`` —
|
||||||
|
``List[Tuple[List[int], List[float]]]``.
|
||||||
|
"""
|
||||||
|
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||||
|
cache = self._cache
|
||||||
|
seq_cap = self.max_seq_len
|
||||||
|
|
||||||
|
tasks: List[Task] = []
|
||||||
|
for ids in prompt_ids_list:
|
||||||
|
if len(ids) >= seq_cap:
|
||||||
|
tasks.append(None)
|
||||||
|
continue
|
||||||
|
t_max = max_tokens
|
||||||
|
if t_max is None:
|
||||||
|
t_max = seq_cap - len(ids)
|
||||||
|
else:
|
||||||
|
t_max = min(t_max, seq_cap - len(ids))
|
||||||
|
task = Task(
|
||||||
|
task_id=f"batch_{uuid.uuid4().hex[:8]}",
|
||||||
|
prompt_ids=list(ids),
|
||||||
|
max_tokens=t_max,
|
||||||
|
temperature=temperature,
|
||||||
|
top_p=top_p,
|
||||||
|
top_k=top_k,
|
||||||
|
frequency_penalty=frequency_penalty,
|
||||||
|
rep_window=rep_window,
|
||||||
|
)
|
||||||
|
if not cache.task_alloc(task.task_id, task.prompt_ids):
|
||||||
|
tasks.append(None)
|
||||||
|
continue
|
||||||
|
task.input_tokens = len(task.prompt_ids)
|
||||||
|
tasks.append(task)
|
||||||
|
|
||||||
|
try:
|
||||||
|
live = [t for t in tasks if t is not None]
|
||||||
|
prefill_groups: Dict[Tuple[int, int], List[Task]] = {}
|
||||||
|
for t in live:
|
||||||
|
key = (len(t.prompt_ids), cache.task_cached(t.task_id))
|
||||||
|
prefill_groups.setdefault(key, []).append(t)
|
||||||
|
for (prompt_len, start_pos), group in prefill_groups.items():
|
||||||
|
self._executor.execute_prefill(group, prompt_len, start_pos)
|
||||||
|
|
||||||
|
while live:
|
||||||
|
valid: List[Task] = []
|
||||||
|
for t in sorted(live, key=lambda x: x.task_id):
|
||||||
|
if cache.task_extend(t.task_id, t.next_pos):
|
||||||
|
valid.append(t)
|
||||||
|
else:
|
||||||
|
t.status = TaskStatus.ABORTED
|
||||||
|
if not valid:
|
||||||
|
break
|
||||||
|
|
||||||
|
step_out = self._executor.execute_decode(
|
||||||
|
valid, return_logprobs=return_logprobs
|
||||||
|
)
|
||||||
|
if return_logprobs:
|
||||||
|
for t, (ntok, _lp) in zip(valid, step_out):
|
||||||
|
t.output_ids.append(ntok)
|
||||||
|
t.output_tokens += 1
|
||||||
|
else:
|
||||||
|
for t, ntok in zip(valid, step_out):
|
||||||
|
t.output_ids.append(ntok)
|
||||||
|
t.output_tokens += 1
|
||||||
|
|
||||||
|
live = [t for t in valid if not t.is_finished(stop_ids)]
|
||||||
|
finally:
|
||||||
|
for t in tasks:
|
||||||
|
if t is not None:
|
||||||
|
cache.task_free(t.task_id)
|
||||||
|
|
||||||
|
results: List[Any] = []
|
||||||
|
for t in tasks:
|
||||||
|
if t is None:
|
||||||
|
results.append(([], []) if return_logprobs else [])
|
||||||
|
elif return_logprobs:
|
||||||
|
results.append((list(t.output_ids), list(t.output_logprobs)))
|
||||||
|
else:
|
||||||
|
results.append(list(t.output_ids))
|
||||||
|
return results
|
||||||
|
|||||||
@@ -6,6 +6,8 @@ from collections import deque
|
|||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import Any, Callable, Deque, Dict, List, Optional
|
from typing import Any, Callable, Deque, Dict, List, Optional
|
||||||
|
|
||||||
|
from tokenizers.decoders import DecodeStream
|
||||||
|
|
||||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -13,6 +15,33 @@ logger = logging.getLogger(__name__)
|
|||||||
STOP = object()
|
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):
|
class TaskStatus(Enum):
|
||||||
"""Task lifecycle states."""
|
"""Task lifecycle states."""
|
||||||
|
|
||||||
@@ -33,7 +62,8 @@ class Task:
|
|||||||
temperature: float = 1.0,
|
temperature: float = 1.0,
|
||||||
top_p: float = 1.0,
|
top_p: float = 1.0,
|
||||||
top_k: int = 50,
|
top_k: int = 50,
|
||||||
stream_callback: Optional[Callable[[str], None]] = None,
|
frequency_penalty: float = 0.0,
|
||||||
|
rep_window: int = 64,
|
||||||
):
|
):
|
||||||
self.task_id = task_id
|
self.task_id = task_id
|
||||||
self.prompt_ids = prompt_ids
|
self.prompt_ids = prompt_ids
|
||||||
@@ -41,14 +71,37 @@ class Task:
|
|||||||
self.temperature = temperature
|
self.temperature = temperature
|
||||||
self.top_p = top_p
|
self.top_p = top_p
|
||||||
self.top_k = top_k
|
self.top_k = top_k
|
||||||
|
self.frequency_penalty = frequency_penalty
|
||||||
|
self.rep_window = rep_window
|
||||||
|
|
||||||
self.status = TaskStatus.PENDING
|
self.status = TaskStatus.PENDING
|
||||||
self.output_ids: List[int] = []
|
self.output_ids: List[int] = []
|
||||||
|
self.output_logprobs: List[float] = []
|
||||||
self.input_tokens: int = 0
|
self.input_tokens: int = 0
|
||||||
self.output_tokens: int = 0
|
self.output_tokens: int = 0
|
||||||
self.arrival_time = time.time()
|
self.arrival_time = time.time()
|
||||||
self.finish_time: Optional[float] = None
|
self.finish_time: Optional[float] = None
|
||||||
self.stream_callback = stream_callback
|
self._decoder: Optional[StreamDecoder] = None
|
||||||
|
|
||||||
|
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])
|
||||||
|
|
||||||
|
def flush_remaining(self, tokenizer: AutoTokenizer) -> str:
|
||||||
|
"""Emit any text still buffered in the decoder.
|
||||||
|
|
||||||
|
With the Rust-native DecodeStream, the stream is always in a
|
||||||
|
correct state — any completed text was already emitted by the
|
||||||
|
last ``push``. A trailing incomplete multi-byte sequence has no
|
||||||
|
valid text to emit, so this is a no-op.
|
||||||
|
"""
|
||||||
|
return ""
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def next_pos(self) -> int:
|
def next_pos(self) -> int:
|
||||||
@@ -70,15 +123,14 @@ class TaskManager:
|
|||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
max_batch_size: int = 16,
|
max_batch_size: int = 16,
|
||||||
max_seq_len: int = 8192,
|
max_seq_len: int = 8192,
|
||||||
max_prompt_len: int = 512,
|
|
||||||
):
|
):
|
||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
self.max_batch_size = max_batch_size
|
self.max_batch_size = max_batch_size
|
||||||
self.max_seq_len = max_seq_len
|
self.max_seq_len = max_seq_len
|
||||||
self.max_prompt_len = max_prompt_len
|
|
||||||
|
|
||||||
self.waiting_queue: Deque[Task] = deque()
|
self.waiting_queue: Deque[Task] = deque()
|
||||||
self.active_tasks: List[Task] = []
|
self.active_tasks: List[Task] = []
|
||||||
|
self._callbacks: Dict[str, Callable[[str], None]] = {}
|
||||||
|
|
||||||
self._task_event = threading.Event()
|
self._task_event = threading.Event()
|
||||||
self._lock = threading.Lock()
|
self._lock = threading.Lock()
|
||||||
@@ -93,14 +145,16 @@ class TaskManager:
|
|||||||
temperature: float = 1.0,
|
temperature: float = 1.0,
|
||||||
top_p: float = 1.0,
|
top_p: float = 1.0,
|
||||||
top_k: int = 50,
|
top_k: int = 50,
|
||||||
|
frequency_penalty: float = 0.0,
|
||||||
|
rep_window: int = 64,
|
||||||
stream_callback: Optional[Callable[[str], None]] = None,
|
stream_callback: Optional[Callable[[str], None]] = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
|
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
|
||||||
prompt_ids = self.tokenizer.encode(prompt)
|
prompt_ids = self.tokenizer.encode(prompt)
|
||||||
if len(prompt_ids) > self.max_prompt_len:
|
if len(prompt_ids) > self.max_seq_len:
|
||||||
prompt_ids = prompt_ids[-self.max_prompt_len :]
|
prompt_ids = prompt_ids[-self.max_seq_len :]
|
||||||
|
|
||||||
if len(prompt_ids) >= self.max_seq_len:
|
if len(prompt_ids) > self.max_seq_len:
|
||||||
if stream_callback:
|
if stream_callback:
|
||||||
stream_callback(STOP)
|
stream_callback(STOP)
|
||||||
return task_id
|
return task_id
|
||||||
@@ -117,12 +171,15 @@ class TaskManager:
|
|||||||
temperature=temperature,
|
temperature=temperature,
|
||||||
top_p=top_p,
|
top_p=top_p,
|
||||||
top_k=top_k,
|
top_k=top_k,
|
||||||
stream_callback=stream_callback,
|
frequency_penalty=frequency_penalty,
|
||||||
|
rep_window=rep_window,
|
||||||
)
|
)
|
||||||
|
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self.waiting_queue.append(task)
|
self.waiting_queue.append(task)
|
||||||
self._total_tasks += 1
|
self._total_tasks += 1
|
||||||
|
if stream_callback:
|
||||||
|
self._callbacks[task_id] = stream_callback
|
||||||
|
|
||||||
self._task_event.set()
|
self._task_event.set()
|
||||||
return task_id
|
return task_id
|
||||||
@@ -134,8 +191,14 @@ class TaskManager:
|
|||||||
t for t in self.waiting_queue if t.task_id != task_id
|
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.active_tasks = [t for t in self.active_tasks if t.task_id != task_id]
|
||||||
|
self._callbacks.pop(task_id, None)
|
||||||
return removed_active
|
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]:
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
return {
|
return {
|
||||||
"total_tasks": self._total_tasks,
|
"total_tasks": self._total_tasks,
|
||||||
@@ -172,12 +235,12 @@ class TaskManager:
|
|||||||
to_add.append(self.waiting_queue.popleft())
|
to_add.append(self.waiting_queue.popleft())
|
||||||
return to_add
|
return to_add
|
||||||
|
|
||||||
def activate(self, task: Task) -> None:
|
def activate(self, task: Task):
|
||||||
task.status = TaskStatus.RUNNING
|
task.status = TaskStatus.RUNNING
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self.active_tasks.append(task)
|
self.active_tasks.append(task)
|
||||||
|
|
||||||
def return_to_waiting(self, tasks: List[Task]) -> None:
|
def return_to_waiting(self, tasks: List[Task]):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
for task in reversed(tasks):
|
for task in reversed(tasks):
|
||||||
self.waiting_queue.appendleft(task)
|
self.waiting_queue.appendleft(task)
|
||||||
@@ -185,18 +248,26 @@ class TaskManager:
|
|||||||
def has_work(self) -> bool:
|
def has_work(self) -> bool:
|
||||||
return bool(self.active_tasks or self.waiting_queue)
|
return bool(self.active_tasks or self.waiting_queue)
|
||||||
|
|
||||||
def wait_for_tasks(self, timeout: float = 1.0) -> None:
|
def wait_for_tasks(self, timeout: float = 1.0):
|
||||||
self._task_event.clear()
|
with self._lock:
|
||||||
|
if self.waiting_queue or self.active_tasks:
|
||||||
|
return
|
||||||
|
self._task_event.clear()
|
||||||
self._task_event.wait(timeout=timeout)
|
self._task_event.wait(timeout=timeout)
|
||||||
|
|
||||||
def get_active_tasks(self) -> List[Task]:
|
def get_active_tasks(self) -> List[Task]:
|
||||||
with self._lock:
|
with self._lock:
|
||||||
return list(self.active_tasks)
|
return list(self.active_tasks)
|
||||||
|
|
||||||
def clear_queues(self) -> None:
|
def get_waiting_tasks(self) -> List[Task]:
|
||||||
|
with self._lock:
|
||||||
|
return list(self.waiting_queue)
|
||||||
|
|
||||||
|
def clear_queues(self):
|
||||||
with self._lock:
|
with self._lock:
|
||||||
self.waiting_queue.clear()
|
self.waiting_queue.clear()
|
||||||
self.active_tasks.clear()
|
self.active_tasks.clear()
|
||||||
|
self._callbacks.clear()
|
||||||
|
|
||||||
def wake(self) -> None:
|
def wake(self):
|
||||||
self._task_event.set()
|
self._task_event.set()
|
||||||
|
|||||||
+74
-25
@@ -8,22 +8,12 @@ from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Tuple,
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
|
from astrai.inference.core.cache import PagePool
|
||||||
from astrai.inference.core.scheduler import InferenceScheduler
|
from astrai.inference.core.scheduler import InferenceScheduler
|
||||||
from astrai.inference.core.task import STOP
|
from astrai.inference.core.task import STOP
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
def _validate_sampling_params(
|
|
||||||
top_k: int, top_p: float, temperature: float, max_tokens: Optional[int] = None
|
|
||||||
):
|
|
||||||
if not (isinstance(top_k, int) and top_k >= 0):
|
|
||||||
raise ValueError("top_k must be a non-negative integer")
|
|
||||||
if not (0.0 <= top_p <= 1.0):
|
|
||||||
raise ValueError("top_p must be a float between 0.0 and 1.0")
|
|
||||||
if not (isinstance(temperature, (int, float)) and temperature >= 0):
|
|
||||||
raise ValueError("temperature must be a non-negative number")
|
|
||||||
|
|
||||||
|
|
||||||
class GenerateResult:
|
class GenerateResult:
|
||||||
"""Thread-safe token accumulator for streaming and non-streaming modes."""
|
"""Thread-safe token accumulator for streaming and non-streaming modes."""
|
||||||
|
|
||||||
@@ -59,7 +49,7 @@ class GenerateResult:
|
|||||||
def wait(self, timeout: Optional[float] = None) -> bool:
|
def wait(self, timeout: Optional[float] = None) -> bool:
|
||||||
return self._event.wait(timeout=timeout)
|
return self._event.wait(timeout=timeout)
|
||||||
|
|
||||||
def wait_completion(self, timeout: float = 300.0) -> None:
|
def wait_completion(self, timeout: float = 300.0):
|
||||||
with self._cond:
|
with self._cond:
|
||||||
if not self._cond.wait_for(
|
if not self._cond.wait_for(
|
||||||
lambda: self._completed >= self._total, timeout=timeout
|
lambda: self._completed >= self._total, timeout=timeout
|
||||||
@@ -84,15 +74,31 @@ class GenerationRequest:
|
|||||||
top_p: float = 1.0,
|
top_p: float = 1.0,
|
||||||
temperature: float = 1.0,
|
temperature: float = 1.0,
|
||||||
max_tokens: Optional[int] = None,
|
max_tokens: Optional[int] = None,
|
||||||
|
frequency_penalty: float = 0.0,
|
||||||
|
rep_window: int = 64,
|
||||||
stream: bool = False,
|
stream: bool = False,
|
||||||
):
|
):
|
||||||
_validate_sampling_params(top_k, top_p, temperature, max_tokens)
|
if not (isinstance(top_k, int) and top_k >= 0):
|
||||||
|
raise ValueError("top_k must be a non-negative integer")
|
||||||
|
if not (0.0 <= top_p <= 1.0):
|
||||||
|
raise ValueError("top_p must be a float between 0.0 and 1.0")
|
||||||
|
if not (isinstance(temperature, (int, float)) and temperature >= 0):
|
||||||
|
raise ValueError("temperature must be a non-negative number")
|
||||||
|
if not (
|
||||||
|
isinstance(frequency_penalty, (int, float))
|
||||||
|
and -2.0 <= frequency_penalty <= 2.0
|
||||||
|
):
|
||||||
|
raise ValueError("frequency_penalty must be between -2.0 and 2.0")
|
||||||
|
if not (isinstance(rep_window, int) and rep_window > 0):
|
||||||
|
raise ValueError("rep_window must be a positive integer")
|
||||||
|
|
||||||
self.messages = messages
|
self.messages = messages
|
||||||
self.top_k = top_k
|
self.top_k = top_k
|
||||||
self.top_p = top_p
|
self.top_p = top_p
|
||||||
self.temperature = temperature
|
self.temperature = temperature
|
||||||
self.max_tokens = max_tokens
|
self.max_tokens = max_tokens
|
||||||
|
self.frequency_penalty = frequency_penalty
|
||||||
|
self.rep_window = rep_window
|
||||||
self.stream = stream
|
self.stream = stream
|
||||||
|
|
||||||
|
|
||||||
@@ -105,8 +111,7 @@ class InferenceEngine:
|
|||||||
tokenizer: AutoTokenizer,
|
tokenizer: AutoTokenizer,
|
||||||
max_batch_size: int = 1,
|
max_batch_size: int = 1,
|
||||||
max_seq_len: Optional[int] = None,
|
max_seq_len: Optional[int] = None,
|
||||||
max_prompt_len: int = 2048,
|
cache: Optional[PagePool] = None,
|
||||||
page_size: int = 128,
|
|
||||||
):
|
):
|
||||||
self.model = model
|
self.model = model
|
||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
@@ -115,8 +120,7 @@ class InferenceEngine:
|
|||||||
tokenizer=self.tokenizer,
|
tokenizer=self.tokenizer,
|
||||||
max_batch_size=max_batch_size,
|
max_batch_size=max_batch_size,
|
||||||
max_seq_len=max_seq_len,
|
max_seq_len=max_seq_len,
|
||||||
max_prompt_len=max_prompt_len,
|
cache=cache,
|
||||||
page_size=page_size,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
self.scheduler.start()
|
self.scheduler.start()
|
||||||
@@ -136,18 +140,33 @@ class InferenceEngine:
|
|||||||
temperature: float = 1.0,
|
temperature: float = 1.0,
|
||||||
top_p: float = 1.0,
|
top_p: float = 1.0,
|
||||||
top_k: int = 50,
|
top_k: int = 50,
|
||||||
|
frequency_penalty: float = 0.0,
|
||||||
|
rep_window: int = 64,
|
||||||
) -> Union[Generator, str, List[str]]:
|
) -> Union[Generator, str, List[str]]:
|
||||||
_validate_sampling_params(top_k, top_p, temperature, max_tokens)
|
|
||||||
is_batch = isinstance(prompt, list)
|
is_batch = isinstance(prompt, list)
|
||||||
prompts = prompt if is_batch else [prompt]
|
prompts = prompt if is_batch else [prompt]
|
||||||
|
|
||||||
if stream:
|
if stream:
|
||||||
return self._generate_streaming(
|
return self._generate_streaming(
|
||||||
prompts, is_batch, max_tokens, temperature, top_p, top_k
|
prompts,
|
||||||
|
is_batch,
|
||||||
|
max_tokens,
|
||||||
|
temperature,
|
||||||
|
top_p,
|
||||||
|
top_k,
|
||||||
|
frequency_penalty,
|
||||||
|
rep_window,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
return self._generate_non_streaming(
|
return self._generate_non_streaming(
|
||||||
prompts, is_batch, max_tokens, temperature, top_p, top_k
|
prompts,
|
||||||
|
is_batch,
|
||||||
|
max_tokens,
|
||||||
|
temperature,
|
||||||
|
top_p,
|
||||||
|
top_k,
|
||||||
|
frequency_penalty,
|
||||||
|
rep_window,
|
||||||
)
|
)
|
||||||
|
|
||||||
def generate_async(
|
def generate_async(
|
||||||
@@ -157,10 +176,18 @@ class InferenceEngine:
|
|||||||
temperature: float = 1.0,
|
temperature: float = 1.0,
|
||||||
top_p: float = 1.0,
|
top_p: float = 1.0,
|
||||||
top_k: int = 50,
|
top_k: int = 50,
|
||||||
|
frequency_penalty: float = 0.0,
|
||||||
|
rep_window: int = 64,
|
||||||
) -> AsyncGenerator[str, None]:
|
) -> AsyncGenerator[str, None]:
|
||||||
_validate_sampling_params(top_k, top_p, temperature, max_tokens)
|
|
||||||
sync_gen = self._generate_streaming(
|
sync_gen = self._generate_streaming(
|
||||||
[prompt], False, max_tokens, temperature, top_p, top_k
|
[prompt],
|
||||||
|
False,
|
||||||
|
max_tokens,
|
||||||
|
temperature,
|
||||||
|
top_p,
|
||||||
|
top_k,
|
||||||
|
frequency_penalty,
|
||||||
|
rep_window,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _agen():
|
async def _agen():
|
||||||
@@ -191,6 +218,8 @@ class InferenceEngine:
|
|||||||
temperature=request.temperature,
|
temperature=request.temperature,
|
||||||
top_p=request.top_p,
|
top_p=request.top_p,
|
||||||
top_k=request.top_k,
|
top_k=request.top_k,
|
||||||
|
frequency_penalty=request.frequency_penalty,
|
||||||
|
rep_window=request.rep_window,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _submit_tasks(
|
def _submit_tasks(
|
||||||
@@ -200,6 +229,8 @@ class InferenceEngine:
|
|||||||
temperature: float,
|
temperature: float,
|
||||||
top_p: float,
|
top_p: float,
|
||||||
top_k: int,
|
top_k: int,
|
||||||
|
frequency_penalty: float,
|
||||||
|
rep_window: int,
|
||||||
) -> Tuple[GenerateResult, List[str]]:
|
) -> Tuple[GenerateResult, List[str]]:
|
||||||
n = len(prompts)
|
n = len(prompts)
|
||||||
result = GenerateResult(count=n)
|
result = GenerateResult(count=n)
|
||||||
@@ -212,6 +243,8 @@ class InferenceEngine:
|
|||||||
temperature=temperature,
|
temperature=temperature,
|
||||||
top_p=top_p,
|
top_p=top_p,
|
||||||
top_k=top_k,
|
top_k=top_k,
|
||||||
|
frequency_penalty=frequency_penalty,
|
||||||
|
rep_window=rep_window,
|
||||||
stream_callback=cb,
|
stream_callback=cb,
|
||||||
)
|
)
|
||||||
task_ids.append(task_id)
|
task_ids.append(task_id)
|
||||||
@@ -232,9 +265,17 @@ class InferenceEngine:
|
|||||||
temperature: float,
|
temperature: float,
|
||||||
top_p: float,
|
top_p: float,
|
||||||
top_k: int,
|
top_k: int,
|
||||||
|
frequency_penalty: float,
|
||||||
|
rep_window: int,
|
||||||
) -> Generator:
|
) -> Generator:
|
||||||
result, task_ids = self._submit_tasks(
|
result, task_ids = self._submit_tasks(
|
||||||
prompts, max_tokens, temperature, top_p, top_k
|
prompts,
|
||||||
|
max_tokens,
|
||||||
|
temperature,
|
||||||
|
top_p,
|
||||||
|
top_k,
|
||||||
|
frequency_penalty,
|
||||||
|
rep_window,
|
||||||
)
|
)
|
||||||
n = len(prompts)
|
n = len(prompts)
|
||||||
remaining = n
|
remaining = n
|
||||||
@@ -268,9 +309,17 @@ class InferenceEngine:
|
|||||||
temperature: float,
|
temperature: float,
|
||||||
top_p: float,
|
top_p: float,
|
||||||
top_k: int,
|
top_k: int,
|
||||||
|
frequency_penalty: float,
|
||||||
|
rep_window: int,
|
||||||
) -> Union[str, List[str]]:
|
) -> Union[str, List[str]]:
|
||||||
result, task_ids = self._submit_tasks(
|
result, task_ids = self._submit_tasks(
|
||||||
prompts, max_tokens, temperature, top_p, top_k
|
prompts,
|
||||||
|
max_tokens,
|
||||||
|
temperature,
|
||||||
|
top_p,
|
||||||
|
top_k,
|
||||||
|
frequency_penalty,
|
||||||
|
rep_window,
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -289,7 +338,7 @@ class InferenceEngine:
|
|||||||
def get_stats(self) -> Dict[str, Any]:
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
return self.scheduler.get_stats()
|
return self.scheduler.get_stats()
|
||||||
|
|
||||||
def shutdown(self) -> None:
|
def shutdown(self):
|
||||||
self.scheduler.stop()
|
self.scheduler.stop()
|
||||||
if torch.cuda.is_available():
|
if torch.cuda.is_available():
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
|
|||||||
+244
-28
@@ -1,15 +1,15 @@
|
|||||||
"""Composable sampling strategies for logit transformation.
|
"""Composable sampling strategies for logit transformation.
|
||||||
|
|
||||||
Implements the Strategy pattern: each sampling technique
|
Implements the Strategy pattern: each sampling technique
|
||||||
(temperature, top-k, top-p) is a pluggable strategy that
|
(temperature, top-k, top-p, frequency penalty) is a pluggable
|
||||||
can be composed into a pipeline.
|
strategy that can be composed into a pipeline.
|
||||||
|
|
||||||
All strategies accept both scalar and per-sample tensor
|
All strategies accept both scalar and per-sample tensor
|
||||||
parameters, so a single pipeline works for any batch size.
|
parameters, so a single pipeline works for any batch size.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import List, Union
|
from typing import List, Optional, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
@@ -19,16 +19,28 @@ class BaseSamplingStrategy(ABC):
|
|||||||
"""Abstract base for a logit transformation strategy."""
|
"""Abstract base for a logit transformation strategy."""
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
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.
|
"""Applies the strategy to logits.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
logits: Raw logits tensor (batch, vocab_size).
|
logits: Raw logits tensor (batch, vocab_size).
|
||||||
filter_value: Value assigned to filtered-out positions.
|
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:
|
Returns:
|
||||||
Transformed logits tensor.
|
Transformed logits tensor.
|
||||||
"""
|
"""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
|
||||||
class TemperatureStrategy(BaseSamplingStrategy):
|
class TemperatureStrategy(BaseSamplingStrategy):
|
||||||
@@ -41,13 +53,21 @@ class TemperatureStrategy(BaseSamplingStrategy):
|
|||||||
def __init__(self, temperature: Union[float, Tensor] = 1.0):
|
def __init__(self, temperature: Union[float, Tensor] = 1.0):
|
||||||
self.temperature = temperature
|
self.temperature = temperature
|
||||||
|
|
||||||
def apply(self, logits, filter_value=-float("inf")):
|
def apply(
|
||||||
|
self,
|
||||||
|
logits: Tensor,
|
||||||
|
filter_value: float = -float("inf"),
|
||||||
|
input_ids: Optional[Tensor] = None,
|
||||||
|
input_mask: Optional[Tensor] = None,
|
||||||
|
) -> Tensor:
|
||||||
t = self.temperature
|
t = self.temperature
|
||||||
if isinstance(t, Tensor):
|
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():
|
if (t != 1.0).any():
|
||||||
logits = logits / t.to(logits.device, non_blocking=True).view(-1, 1)
|
logits = logits / t
|
||||||
elif t != 1.0:
|
elif t != 1.0:
|
||||||
logits = logits / t
|
logits = logits / max(t, 1e-8)
|
||||||
return logits
|
return logits
|
||||||
|
|
||||||
|
|
||||||
@@ -61,7 +81,13 @@ class TopKStrategy(BaseSamplingStrategy):
|
|||||||
def __init__(self, top_k: Union[int, Tensor] = 0):
|
def __init__(self, top_k: Union[int, Tensor] = 0):
|
||||||
self.top_k = top_k
|
self.top_k = top_k
|
||||||
|
|
||||||
def apply(self, logits, filter_value=-float("inf")):
|
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
|
tk = self.top_k
|
||||||
if isinstance(tk, Tensor):
|
if isinstance(tk, Tensor):
|
||||||
tk = tk.to(logits.device, non_blocking=True).long().clamp(min=0)
|
tk = tk.to(logits.device, non_blocking=True).long().clamp(min=0)
|
||||||
@@ -98,7 +124,9 @@ class TopPStrategy(BaseSamplingStrategy):
|
|||||||
def __init__(self, top_p: Union[float, Tensor] = 1.0):
|
def __init__(self, top_p: Union[float, Tensor] = 1.0):
|
||||||
self.top_p = top_p
|
self.top_p = top_p
|
||||||
|
|
||||||
def _apply(self, logits, top_p, filter_value):
|
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)
|
sorted_logits, sorted_indices = torch.sort(logits, descending=True, dim=-1)
|
||||||
cum_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
|
cum_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
|
||||||
remove = cum_probs > top_p
|
remove = cum_probs > top_p
|
||||||
@@ -109,7 +137,13 @@ class TopPStrategy(BaseSamplingStrategy):
|
|||||||
logits[mask] = filter_value
|
logits[mask] = filter_value
|
||||||
return logits
|
return logits
|
||||||
|
|
||||||
def apply(self, logits, filter_value=-float("inf")):
|
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
|
tp = self.top_p
|
||||||
if isinstance(tp, Tensor):
|
if isinstance(tp, Tensor):
|
||||||
tp = tp.to(logits.device, non_blocking=True)
|
tp = tp.to(logits.device, non_blocking=True)
|
||||||
@@ -120,6 +154,84 @@ class TopPStrategy(BaseSamplingStrategy):
|
|||||||
return logits
|
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):
|
class SamplingPipeline(BaseSamplingStrategy):
|
||||||
"""Composes multiple sampling strategies into a single transformation.
|
"""Composes multiple sampling strategies into a single transformation.
|
||||||
|
|
||||||
@@ -140,25 +252,76 @@ class SamplingPipeline(BaseSamplingStrategy):
|
|||||||
def __init__(self, strategies: List[BaseSamplingStrategy]):
|
def __init__(self, strategies: List[BaseSamplingStrategy]):
|
||||||
self.strategies = strategies
|
self.strategies = strategies
|
||||||
|
|
||||||
def apply(self, logits, filter_value=-float("inf")):
|
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:
|
for strategy in self.strategies:
|
||||||
logits = strategy.apply(logits, filter_value)
|
logits = strategy.apply(logits, filter_value, input_ids, input_mask)
|
||||||
return logits
|
return logits
|
||||||
|
|
||||||
@torch.no_grad()
|
@staticmethod
|
||||||
def sample(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
def _is_greedy(temperature: Union[float, Tensor]) -> bool:
|
||||||
|
if isinstance(temperature, Tensor):
|
||||||
|
return temperature.numel() == 1 and temperature.item() == 0
|
||||||
|
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).
|
"""Apply strategies then sample (softmax + multinomial).
|
||||||
|
|
||||||
|
Short-circuits to ``argmax`` when temperature is exactly 0
|
||||||
|
(deterministic / greedy decode).
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
logits: Raw logits ``[batch, vocab_size]``.
|
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:
|
Returns:
|
||||||
Sampled token IDs ``[batch]``.
|
Sampled token IDs ``[batch]``, or — when ``return_logprobs``
|
||||||
|
is ``True`` — a ``(token_ids, chosen_logprobs)`` tuple.
|
||||||
"""
|
"""
|
||||||
return torch.multinomial(
|
if self._is_greedy_pipeline():
|
||||||
torch.softmax(self.apply(logits, filter_value), dim=-1),
|
tokens = logits.argmax(dim=-1)
|
||||||
num_samples=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)
|
||||||
|
log_probs = torch.log_softmax(transformed.float(), dim=-1)
|
||||||
|
tokens = torch.multinomial(
|
||||||
|
torch.softmax(transformed, dim=-1), num_samples=1
|
||||||
).squeeze(-1)
|
).squeeze(-1)
|
||||||
|
if not return_logprobs:
|
||||||
|
return tokens
|
||||||
|
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()
|
@torch.inference_mode()
|
||||||
@@ -167,22 +330,75 @@ def sample(
|
|||||||
temperature: Union[float, Tensor] = 1.0,
|
temperature: Union[float, Tensor] = 1.0,
|
||||||
top_k: Union[int, Tensor] = 0,
|
top_k: Union[int, Tensor] = 0,
|
||||||
top_p: Union[float, Tensor] = 1.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"),
|
filter_value: float = -float("inf"),
|
||||||
) -> Tensor:
|
return_logprobs: bool = False,
|
||||||
|
):
|
||||||
"""Apply sampling strategies then sample (softmax + multinomial).
|
"""Apply sampling strategies then sample (softmax + multinomial).
|
||||||
|
|
||||||
Shortcut for ``SamplingPipeline(...).sample(logits)``.
|
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:
|
Args:
|
||||||
logits: Raw logits ``[batch, vocab_size]``.
|
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:
|
Returns:
|
||||||
Sampled token IDs ``[batch]``.
|
Sampled token IDs ``[batch]``, or — when ``return_logprobs`` is
|
||||||
|
``True`` — a ``(token_ids, chosen_logprobs)`` tuple where
|
||||||
|
``chosen_logprobs`` has shape ``[batch]``.
|
||||||
"""
|
"""
|
||||||
return SamplingPipeline(
|
greedy = (
|
||||||
[
|
(
|
||||||
TemperatureStrategy(temperature),
|
isinstance(temperature, Tensor)
|
||||||
TopKStrategy(top_k),
|
and temperature.numel() == 1
|
||||||
TopPStrategy(top_p),
|
and temperature.item() == 0
|
||||||
]
|
)
|
||||||
).sample(logits, filter_value)
|
if isinstance(temperature, Tensor)
|
||||||
|
else temperature == 0
|
||||||
|
)
|
||||||
|
|
||||||
|
if greedy:
|
||||||
|
tokens = logits.argmax(dim=-1)
|
||||||
|
if not return_logprobs:
|
||||||
|
return tokens
|
||||||
|
log_probs = torch.log_softmax(logits.float(), dim=-1)
|
||||||
|
chosen = torch.gather(log_probs, -1, tokens.unsqueeze(-1)).squeeze(-1)
|
||||||
|
return tokens, chosen
|
||||||
|
|
||||||
|
has_freq = (
|
||||||
|
(isinstance(frequency_penalty, Tensor) and (frequency_penalty != 0).any())
|
||||||
|
if isinstance(frequency_penalty, Tensor)
|
||||||
|
else frequency_penalty != 0
|
||||||
|
)
|
||||||
|
|
||||||
|
strategies: List[BaseSamplingStrategy] = [
|
||||||
|
TemperatureStrategy(temperature),
|
||||||
|
TopKStrategy(top_k),
|
||||||
|
TopPStrategy(top_p),
|
||||||
|
]
|
||||||
|
if has_freq:
|
||||||
|
strategies.append(FrequencyPenaltyStrategy(frequency_penalty))
|
||||||
|
|
||||||
|
return SamplingPipeline(strategies).sample(
|
||||||
|
logits,
|
||||||
|
filter_value=filter_value,
|
||||||
|
input_ids=input_ids,
|
||||||
|
input_mask=input_mask,
|
||||||
|
return_logprobs=return_logprobs,
|
||||||
|
)
|
||||||
|
|||||||
@@ -1,12 +1,18 @@
|
|||||||
from astrai.model.automodel import AutoModel
|
from astrai.model.automodel import AutoModel
|
||||||
from astrai.model.module import (
|
from astrai.model.components.attention import GQA
|
||||||
GQA,
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
MLP,
|
from astrai.model.components.linear import Linear
|
||||||
DecoderBlock,
|
from astrai.model.components.lora import (
|
||||||
Linear,
|
LoRAConfig,
|
||||||
RMSNorm,
|
inject_lora,
|
||||||
|
load_lora,
|
||||||
|
merge_lora,
|
||||||
|
save_lora,
|
||||||
)
|
)
|
||||||
from astrai.model.transformer import Transformer
|
from astrai.model.components.mlp import MLP
|
||||||
|
from astrai.model.components.norm import RMSNorm
|
||||||
|
from astrai.model.encoder import EmbeddingEncoder
|
||||||
|
from astrai.model.transformer import AutoRegressiveLM
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
# Modules
|
# Modules
|
||||||
@@ -16,6 +22,13 @@ __all__ = [
|
|||||||
"GQA",
|
"GQA",
|
||||||
"DecoderBlock",
|
"DecoderBlock",
|
||||||
# Models
|
# Models
|
||||||
"Transformer",
|
"AutoRegressiveLM",
|
||||||
|
"EmbeddingEncoder",
|
||||||
"AutoModel",
|
"AutoModel",
|
||||||
|
# LoRA
|
||||||
|
"LoRAConfig",
|
||||||
|
"inject_lora",
|
||||||
|
"merge_lora",
|
||||||
|
"save_lora",
|
||||||
|
"load_lora",
|
||||||
]
|
]
|
||||||
|
|||||||
+33
-36
@@ -6,16 +6,20 @@ from contextlib import contextmanager
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Self, Union
|
from typing import Self, Union
|
||||||
|
|
||||||
import safetensors.torch as st
|
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
from astrai.config import ModelConfig
|
from astrai.config.model_config import BaseModelConfig, ConfigFactory
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.serialization import load_model_config, load_model_weights, save_model
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def _disable_random_init(enable: bool = True):
|
def _disable_random_init(enable: bool = True):
|
||||||
init_functions = [
|
if not enable:
|
||||||
|
yield
|
||||||
|
return
|
||||||
|
|
||||||
|
names = (
|
||||||
"xavier_normal_",
|
"xavier_normal_",
|
||||||
"xavier_uniform_",
|
"xavier_uniform_",
|
||||||
"kaiming_normal_",
|
"kaiming_normal_",
|
||||||
@@ -25,27 +29,25 @@ def _disable_random_init(enable: bool = True):
|
|||||||
"constant_",
|
"constant_",
|
||||||
"normal_",
|
"normal_",
|
||||||
"uniform_",
|
"uniform_",
|
||||||
]
|
)
|
||||||
original_funcs = {}
|
orig = {n: getattr(nn.init, n) for n in names if hasattr(nn.init, n)}
|
||||||
for name in init_functions:
|
for n in orig:
|
||||||
if enable and hasattr(nn.init, name):
|
setattr(nn.init, n, lambda *a, **kw: None)
|
||||||
original_funcs[name] = getattr(nn.init, name)
|
|
||||||
setattr(nn.init, name, lambda *args, **kwargs: None)
|
|
||||||
try:
|
try:
|
||||||
yield
|
yield
|
||||||
finally:
|
finally:
|
||||||
if enable:
|
for n, fn in orig.items():
|
||||||
for name, orig_func in original_funcs.items():
|
setattr(nn.init, n, fn)
|
||||||
setattr(nn.init, name, orig_func)
|
|
||||||
|
|
||||||
|
|
||||||
class AutoModel(BaseFactory["AutoModel"], nn.Module):
|
class ModelFactory(BaseFactory[nn.Module]):
|
||||||
"""
|
"""Pure factory for model dispatch, separated from nn.Module state."""
|
||||||
Autoregressive language model base class.
|
|
||||||
Provides model loading/saving, registration, and generation.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, config: ModelConfig):
|
|
||||||
|
class AutoModel(nn.Module):
|
||||||
|
"""Model base class with loading/saving and generation."""
|
||||||
|
|
||||||
|
def __init__(self, config: BaseModelConfig):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.config = config
|
self.config = config
|
||||||
|
|
||||||
@@ -59,24 +61,22 @@ class AutoModel(BaseFactory["AutoModel"], nn.Module):
|
|||||||
|
|
||||||
model_path = Path(path)
|
model_path = Path(path)
|
||||||
|
|
||||||
# Load config
|
|
||||||
config = ModelConfig()
|
|
||||||
config_path = model_path / "config.json"
|
config_path = model_path / "config.json"
|
||||||
if config_path.exists():
|
if not config_path.exists():
|
||||||
config.load(str(config_path))
|
|
||||||
else:
|
|
||||||
raise FileNotFoundError(f"Config file not found: {config_path}")
|
raise FileNotFoundError(f"Config file not found: {config_path}")
|
||||||
|
|
||||||
model_type = config.model_type or "transformer"
|
raw = load_model_config(str(model_path))
|
||||||
actual_cls = AutoModel.get_component_class(model_type)
|
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):
|
with _disable_random_init(enable=disable_random_init):
|
||||||
model = actual_cls(config)
|
model = actual_cls(config)
|
||||||
|
|
||||||
# Load weights
|
|
||||||
weights_path = model_path / "model.safetensors"
|
weights_path = model_path / "model.safetensors"
|
||||||
if weights_path.exists():
|
if weights_path.exists():
|
||||||
state_dict = st.load_file(str(weights_path))
|
state_dict = load_model_weights(str(model_path))
|
||||||
model.load_state_dict(state_dict, strict=strict)
|
model.load_state_dict(state_dict, strict=strict)
|
||||||
|
|
||||||
return model
|
return model
|
||||||
@@ -84,15 +84,12 @@ class AutoModel(BaseFactory["AutoModel"], nn.Module):
|
|||||||
def save_pretrained(
|
def save_pretrained(
|
||||||
self,
|
self,
|
||||||
save_directory: Union[str, Path],
|
save_directory: Union[str, Path],
|
||||||
) -> None:
|
):
|
||||||
save_path = Path(save_directory)
|
save_model(
|
||||||
save_path.mkdir(parents=True, exist_ok=True)
|
config=self.config.to_dict(),
|
||||||
|
state_dict=self.state_dict(),
|
||||||
# Save config
|
save_directory=str(save_directory),
|
||||||
self.config.save(str(save_path / "config.json"))
|
)
|
||||||
|
|
||||||
# Save weights
|
|
||||||
st.save_file(self.state_dict(), str(save_path / "model.safetensors"))
|
|
||||||
|
|
||||||
def to(self, *args, **kwargs) -> Self:
|
def to(self, *args, **kwargs) -> Self:
|
||||||
"""Move model to device/dtype."""
|
"""Move model to device/dtype."""
|
||||||
|
|||||||
@@ -0,0 +1,24 @@
|
|||||||
|
from astrai.extension.rotary_backend 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
|
||||||
|
from astrai.model.components.norm import RMSNorm
|
||||||
|
from astrai.model.components.rope import (
|
||||||
|
RotaryEmbedding,
|
||||||
|
get_rotary_emb,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"Linear",
|
||||||
|
"RMSNorm",
|
||||||
|
"MLP",
|
||||||
|
"Embedding",
|
||||||
|
"GQA",
|
||||||
|
"MLA",
|
||||||
|
"DecoderBlock",
|
||||||
|
"RotaryEmbedding",
|
||||||
|
"apply_rotary_emb",
|
||||||
|
"get_rotary_emb",
|
||||||
|
]
|
||||||
@@ -0,0 +1,180 @@
|
|||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.extension import attention
|
||||||
|
from astrai.extension.rotary_backend import apply_rotary_emb
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.inference.core.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:
|
||||||
|
batch_size, seq_len, _ = x.shape
|
||||||
|
x = x.reshape(batch_size, seq_len, n_heads, self.head_dim)
|
||||||
|
return x
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
x: Tensor,
|
||||||
|
rotary_emb: Tensor,
|
||||||
|
attn_mask: Tensor = None,
|
||||||
|
kv_cache: Optional[KVCache] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> 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)
|
||||||
|
|
||||||
|
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,
|
||||||
|
) -> Tensor:
|
||||||
|
bsz, seq_len, _ = x.size()
|
||||||
|
|
||||||
|
q = self.q_proj(x)
|
||||||
|
q = q.view(bsz, seq_len, self.n_heads, self.head_dim)
|
||||||
|
|
||||||
|
kv_compressed = self.kv_a_proj(x)
|
||||||
|
kv_compressed = self.kv_norm(kv_compressed)
|
||||||
|
|
||||||
|
kv = self.kv_b_proj(kv_compressed)
|
||||||
|
kv = kv.view(bsz, seq_len, self.n_kv_heads, -1)
|
||||||
|
|
||||||
|
k_nope, k_rope, v = torch.split(
|
||||||
|
kv, [self.qk_nope_head_dim, self.qk_rope_head_dim, self.head_dim], dim=-1
|
||||||
|
)
|
||||||
|
|
||||||
|
q_nope, q_rope = (
|
||||||
|
q[..., : self.qk_nope_head_dim],
|
||||||
|
q[..., self.qk_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)
|
||||||
|
|
||||||
|
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,49 @@
|
|||||||
|
from dataclasses import asdict
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch.nn as nn
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.inference.core.cache import KVCache
|
||||||
|
from astrai.model.components.attention import AttnFactory
|
||||||
|
from astrai.model.components.mlp import FFNFactory
|
||||||
|
from astrai.model.components.norm import RMSNorm
|
||||||
|
|
||||||
|
|
||||||
|
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)
|
||||||
|
self.mlp = FFNFactory.create(config.ffn_type, **cfg)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
x: Tensor,
|
||||||
|
rotary_emb: Tensor,
|
||||||
|
attention_mask: Optional[Tensor] = None,
|
||||||
|
kv_cache: Optional[KVCache] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
attn_output = self.attention(
|
||||||
|
self.input_norm(x),
|
||||||
|
rotary_emb,
|
||||||
|
attention_mask,
|
||||||
|
kv_cache,
|
||||||
|
is_causal,
|
||||||
|
)
|
||||||
|
x = attn_output + x
|
||||||
|
x = self.mlp(self.post_attention_norm(x)) + x
|
||||||
|
|
||||||
|
return x
|
||||||
@@ -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,100 @@
|
|||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
@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) -> Tensor:
|
||||||
|
gated = self.up(x) * F.silu(self.gate(x))
|
||||||
|
out = self.down(gated)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
@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,
|
||||||
|
):
|
||||||
|
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.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, dim_ffn, down_init_std=down_init_std)
|
||||||
|
for _ in range(n_shared_experts)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
self.routed_experts = nn.ModuleList(
|
||||||
|
[
|
||||||
|
MLP(dim, dim_ffn, down_init_std=down_init_std)
|
||||||
|
for _ in range(n_routed_experts)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
bsz, seq_len, dim = x.shape
|
||||||
|
x_flat = x.view(-1, dim)
|
||||||
|
|
||||||
|
shared_out = self._shared_forward(x_flat)
|
||||||
|
routed_out = self._routed_forward(x_flat)
|
||||||
|
|
||||||
|
out = (shared_out + routed_out).view(bsz, seq_len, dim)
|
||||||
|
return out
|
||||||
|
|
||||||
|
def _shared_forward(self, x: Tensor) -> Tensor:
|
||||||
|
if self.n_shared_experts == 0:
|
||||||
|
return torch.zeros_like(x)
|
||||||
|
return sum(e(x) for e in self.shared_experts) / self.n_shared_experts
|
||||||
|
|
||||||
|
def _routed_forward(self, x: Tensor) -> Tensor:
|
||||||
|
N, D = x.shape
|
||||||
|
K = self.n_activated_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)
|
||||||
|
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
||||||
|
|
||||||
|
output = torch.zeros(N, D, device=x.device, dtype=x.dtype)
|
||||||
|
for expert_idx in range(self.n_routed_experts):
|
||||||
|
expert_mask = topk_indices == expert_idx
|
||||||
|
token_idx, k_idx = expert_mask.nonzero(as_tuple=True)
|
||||||
|
if token_idx.numel() == 0:
|
||||||
|
continue
|
||||||
|
expert_input = x[token_idx]
|
||||||
|
expert_output = self.routed_experts[expert_idx](expert_input)
|
||||||
|
weights = topk_weights[token_idx, k_idx].unsqueeze(-1)
|
||||||
|
output.index_add_(0, token_idx, expert_output * weights)
|
||||||
|
|
||||||
|
return output
|
||||||
@@ -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,73 @@
|
|||||||
|
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:
|
||||||
|
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 = self.norm(x)
|
||||||
|
|
||||||
|
if self.pooling_type == "cls":
|
||||||
|
pooled = hidden_states[:, 0]
|
||||||
|
elif self.pooling_type == "last":
|
||||||
|
if input_mask is not None:
|
||||||
|
lengths = input_mask.sum(dim=1) - 1
|
||||||
|
pooled = hidden_states[torch.arange(B, device=x.device), lengths]
|
||||||
|
else:
|
||||||
|
pooled = hidden_states[:, -1]
|
||||||
|
else:
|
||||||
|
if input_mask is not None:
|
||||||
|
mask = input_mask.unsqueeze(-1).to(dtype=hidden_states.dtype)
|
||||||
|
pooled = (hidden_states * mask).sum(dim=1) / mask.sum(dim=1).clamp(
|
||||||
|
min=1.0
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
pooled = hidden_states.mean(dim=1)
|
||||||
|
|
||||||
|
if self.normalize_embeddings:
|
||||||
|
pooled = torch.nn.functional.normalize(pooled, p=2, dim=-1)
|
||||||
|
|
||||||
|
return pooled
|
||||||
@@ -1,330 +0,0 @@
|
|||||||
from typing import Optional
|
|
||||||
|
|
||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
import torch.nn.functional as F
|
|
||||||
from torch import Tensor
|
|
||||||
|
|
||||||
from astrai.inference.core.cache import KvcacheView
|
|
||||||
|
|
||||||
|
|
||||||
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
|
||||||
"""Repeat KV heads n_rep times for GQA."""
|
|
||||||
bs, slen, n_heads, head_dim = x.shape
|
|
||||||
if n_rep == 1:
|
|
||||||
return x
|
|
||||||
return (
|
|
||||||
x[:, :, :, None, :]
|
|
||||||
.expand(bs, slen, n_heads, n_rep, head_dim)
|
|
||||||
.reshape(bs, slen, n_heads * n_rep, head_dim)
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def get_rotary_emb(
|
|
||||||
dim: int,
|
|
||||||
max_len: int,
|
|
||||||
base: float = 10000,
|
|
||||||
device: Optional[torch.device] = None,
|
|
||||||
) -> Tensor:
|
|
||||||
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.complex(cos, sin)
|
|
||||||
|
|
||||||
|
|
||||||
def apply_rotary_emb(x: torch.Tensor, freqs_cis: Tensor) -> Tensor:
|
|
||||||
dtype = x.dtype
|
|
||||||
x_ = x.float().reshape(*x.shape[:-1], -1, 2)
|
|
||||||
x_complex = torch.view_as_complex(x_)
|
|
||||||
freqs_cis = freqs_cis.unsqueeze(2)
|
|
||||||
x_rotated = x_complex * freqs_cis
|
|
||||||
x_out = torch.view_as_real(x_rotated).flatten(-2)
|
|
||||||
return x_out.to(dtype)
|
|
||||||
|
|
||||||
|
|
||||||
class RotaryEmbedding(nn.Module):
|
|
||||||
def __init__(self, dim: int, max_len: int, base: int = 10000):
|
|
||||||
super().__init__()
|
|
||||||
self.dim = dim
|
|
||||||
self.max_len = max_len
|
|
||||||
self.base = base
|
|
||||||
self._set_rotary_buffer(self.max_len)
|
|
||||||
|
|
||||||
def _set_rotary_buffer(self, max_len: int):
|
|
||||||
rotary_emb = get_rotary_emb(self.dim, max_len, self.base)
|
|
||||||
freqs_cis = torch.view_as_real(rotary_emb)
|
|
||||||
self.register_buffer("freqs_cis", freqs_cis, persistent=False)
|
|
||||||
|
|
||||||
def forward(self, x: Tensor, position_ids: Optional[Tensor] = None) -> Tensor:
|
|
||||||
if position_ids is None:
|
|
||||||
position_ids = (
|
|
||||||
torch.arange(x.size(1), device=x.device)
|
|
||||||
.unsqueeze(0)
|
|
||||||
.expand(x.size(0), -1)
|
|
||||||
)
|
|
||||||
position_freq_cis = self.freqs_cis[position_ids].float()
|
|
||||||
return torch.view_as_complex(position_freq_cis)
|
|
||||||
|
|
||||||
|
|
||||||
class Linear(nn.Module):
|
|
||||||
def __init__(self, in_dim: int, out_dim: int, bias: bool = False):
|
|
||||||
super().__init__()
|
|
||||||
self.weight = nn.Parameter(torch.empty((out_dim, in_dim)))
|
|
||||||
self.bias = nn.Parameter(torch.zeros(out_dim)) if bias else None
|
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
|
||||||
return F.linear(x, self.weight, self.bias)
|
|
||||||
|
|
||||||
|
|
||||||
class RMSNorm(nn.Module):
|
|
||||||
def __init__(self, dim, norm_eps):
|
|
||||||
super().__init__()
|
|
||||||
self.weight = nn.Parameter(torch.ones(dim))
|
|
||||||
self.normalized_shape = (dim,)
|
|
||||||
self.norm_eps = norm_eps
|
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
|
||||||
return F.rms_norm(x, self.normalized_shape, self.weight, self.norm_eps)
|
|
||||||
|
|
||||||
|
|
||||||
class MLP(nn.Module):
|
|
||||||
def __init__(self, dim: int, dim_feed_forward: int):
|
|
||||||
super().__init__()
|
|
||||||
self.up = Linear(dim, dim_feed_forward)
|
|
||||||
self.gate = Linear(dim, dim_feed_forward)
|
|
||||||
self.down = Linear(dim_feed_forward, dim)
|
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
|
||||||
gated = self.up(x) * F.silu(self.gate(x))
|
|
||||||
out = self.down(gated)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
class GQA(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
dim: int,
|
|
||||||
n_heads: int,
|
|
||||||
n_kv_heads: int,
|
|
||||||
use_qk_norm: bool,
|
|
||||||
norm_eps: float,
|
|
||||||
use_gated_attention: bool,
|
|
||||||
layer_id: int,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
assert dim % n_heads == 0
|
|
||||||
assert n_heads % n_kv_heads == 0
|
|
||||||
|
|
||||||
self.head_dim = dim // n_heads
|
|
||||||
self.layer_id = layer_id
|
|
||||||
self.dim = dim
|
|
||||||
self.n_heads = n_heads
|
|
||||||
self.n_kv_heads = n_kv_heads
|
|
||||||
self.n_rep = n_heads // n_kv_heads
|
|
||||||
self.use_qk_norm = use_qk_norm
|
|
||||||
self.use_gated_attention = use_gated_attention
|
|
||||||
|
|
||||||
self.q_proj = Linear(dim, n_heads * self.head_dim)
|
|
||||||
self.k_proj = Linear(dim, n_kv_heads * self.head_dim)
|
|
||||||
self.v_proj = Linear(dim, n_kv_heads * self.head_dim)
|
|
||||||
self.o_proj = Linear(dim, dim)
|
|
||||||
|
|
||||||
if self.use_qk_norm:
|
|
||||||
self.q_norm = RMSNorm(self.head_dim, norm_eps)
|
|
||||||
self.k_norm = RMSNorm(self.head_dim, norm_eps)
|
|
||||||
|
|
||||||
if self.use_gated_attention:
|
|
||||||
self.gate = Linear(dim, dim)
|
|
||||||
|
|
||||||
def _split_heads(self, x: Tensor, n_heads) -> Tensor:
|
|
||||||
batch_size, seq_len, _ = x.shape
|
|
||||||
x = x.reshape(batch_size, seq_len, n_heads, self.head_dim)
|
|
||||||
return x
|
|
||||||
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
x: Tensor,
|
|
||||||
rotary_emb: Tensor,
|
|
||||||
attn_mask: Tensor = None,
|
|
||||||
paged_cache: Optional[KvcacheView] = None,
|
|
||||||
) -> Tensor:
|
|
||||||
is_causal = attn_mask is None
|
|
||||||
|
|
||||||
# (bsz, seq_len, dim) -> (bsz, seq_len, n_heads, head_dim)
|
|
||||||
q = self._split_heads(self.q_proj(x), self.n_heads)
|
|
||||||
k = self._split_heads(self.k_proj(x), self.n_kv_heads)
|
|
||||||
v = self._split_heads(self.v_proj(x), self.n_kv_heads)
|
|
||||||
q, k = apply_rotary_emb(q, rotary_emb), apply_rotary_emb(k, rotary_emb)
|
|
||||||
|
|
||||||
if self.use_qk_norm:
|
|
||||||
q, k = self.q_norm(q), self.k_norm(k)
|
|
||||||
|
|
||||||
if paged_cache is not None:
|
|
||||||
paged_cache.write(self.layer_id, k, v)
|
|
||||||
k, v = paged_cache.gather(self.layer_id)
|
|
||||||
|
|
||||||
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)
|
|
||||||
|
|
||||||
# (bsz, seq_len, n_heads, head_dim) -> (bsz, n_heads, seq_len, head_dim)
|
|
||||||
q, k, v = q.permute(0, 2, 1, 3), k.permute(0, 2, 1, 3), v.permute(0, 2, 1, 3)
|
|
||||||
sdqa_out = (
|
|
||||||
F.scaled_dot_product_attention(q, k, v, attn_mask, is_causal=is_causal)
|
|
||||||
.permute(0, 2, 1, 3)
|
|
||||||
.contiguous()
|
|
||||||
.flatten(2)
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.use_gated_attention:
|
|
||||||
sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
|
|
||||||
|
|
||||||
out = self.o_proj(sdqa_out)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
class MLA(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
dim: int,
|
|
||||||
n_heads: int,
|
|
||||||
n_kv_heads: int,
|
|
||||||
kv_lora_rank: int,
|
|
||||||
qk_nope_head_dim: int,
|
|
||||||
qk_rope_head_dim: int,
|
|
||||||
norm_eps: float,
|
|
||||||
use_gated_attention: bool,
|
|
||||||
layer_id: int,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.dim = dim
|
|
||||||
self.n_heads = n_heads
|
|
||||||
self.n_kv_heads = n_kv_heads
|
|
||||||
self.kv_lora_rank = kv_lora_rank
|
|
||||||
self.qk_nope_head_dim = qk_nope_head_dim
|
|
||||||
self.qk_rope_head_dim = qk_rope_head_dim
|
|
||||||
self.head_dim = qk_nope_head_dim + qk_rope_head_dim
|
|
||||||
self.layer_id = layer_id
|
|
||||||
self.n_rep = n_heads // n_kv_heads
|
|
||||||
self.use_gated_attention = use_gated_attention
|
|
||||||
|
|
||||||
self.q_proj = Linear(dim, n_heads * self.head_dim, bias=False)
|
|
||||||
self.kv_a_proj = Linear(dim, kv_lora_rank, bias=False)
|
|
||||||
self.kv_norm = RMSNorm(kv_lora_rank, norm_eps)
|
|
||||||
|
|
||||||
# fused KV: (k_nope, k_rope, v)
|
|
||||||
self.kv_b_proj = Linear(
|
|
||||||
kv_lora_rank,
|
|
||||||
n_kv_heads * (self.head_dim + qk_rope_head_dim + self.head_dim),
|
|
||||||
)
|
|
||||||
|
|
||||||
self.o_proj = Linear(dim, dim, bias=False)
|
|
||||||
|
|
||||||
if use_gated_attention:
|
|
||||||
self.gate = Linear(dim, dim, bias=False)
|
|
||||||
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
x: Tensor,
|
|
||||||
rotary_emb: Tensor,
|
|
||||||
attn_mask: Tensor = None,
|
|
||||||
paged_cache: Optional[KvcacheView] = None,
|
|
||||||
) -> Tensor:
|
|
||||||
bsz, seq_len, _ = x.size()
|
|
||||||
is_causal = attn_mask is None
|
|
||||||
|
|
||||||
q = self.q_proj(x)
|
|
||||||
q = q.view(bsz, seq_len, self.n_heads, self.head_dim)
|
|
||||||
|
|
||||||
kv_compressed = self.kv_a_proj(x)
|
|
||||||
kv_compressed = self.kv_norm(kv_compressed)
|
|
||||||
|
|
||||||
kv = self.kv_b_proj(kv_compressed)
|
|
||||||
kv = kv.view(bsz, seq_len, self.n_kv_heads, -1)
|
|
||||||
|
|
||||||
k_nope, k_rope, v = torch.split(
|
|
||||||
kv, [self.qk_nope_head_dim, self.qk_rope_head_dim, self.head_dim], dim=-1
|
|
||||||
)
|
|
||||||
|
|
||||||
q_nope, q_rope = (
|
|
||||||
q[..., : self.qk_nope_head_dim],
|
|
||||||
q[..., self.qk_rope_head_dim :],
|
|
||||||
)
|
|
||||||
q_rope = apply_rotary_emb(q_rope, rotary_emb)
|
|
||||||
k_rope = apply_rotary_emb(k_rope, rotary_emb)
|
|
||||||
|
|
||||||
q = torch.cat([q_nope, q_rope], dim=-1)
|
|
||||||
k = torch.cat([k_nope, k_rope], dim=-1)
|
|
||||||
|
|
||||||
if paged_cache is not None:
|
|
||||||
paged_cache.write(self.layer_id, k, v)
|
|
||||||
k, v = paged_cache.gather(self.layer_id)
|
|
||||||
|
|
||||||
q = q.permute(0, 2, 1, 3)
|
|
||||||
k = k.permute(0, 2, 1, 3)
|
|
||||||
v = v.permute(0, 2, 1, 3)
|
|
||||||
|
|
||||||
attn_out = F.scaled_dot_product_attention(
|
|
||||||
q, k, v, attn_mask, is_causal=is_causal
|
|
||||||
)
|
|
||||||
attn_out = attn_out.permute(0, 2, 1, 3).contiguous().flatten(2)
|
|
||||||
|
|
||||||
if self.use_gated_attention:
|
|
||||||
attn_out = attn_out * F.sigmoid(self.gate(x))
|
|
||||||
|
|
||||||
out = self.o_proj(attn_out)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
class DecoderBlock(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
dim: int,
|
|
||||||
n_heads: int,
|
|
||||||
dim_ffn: int,
|
|
||||||
n_kv_heads: int,
|
|
||||||
norm_eps: int,
|
|
||||||
use_qk_norm: bool,
|
|
||||||
use_gated_attention: bool,
|
|
||||||
layer_id: int,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.attention = GQA(
|
|
||||||
dim,
|
|
||||||
n_heads,
|
|
||||||
n_kv_heads,
|
|
||||||
use_qk_norm,
|
|
||||||
norm_eps,
|
|
||||||
use_gated_attention,
|
|
||||||
layer_id,
|
|
||||||
)
|
|
||||||
self.input_norm = RMSNorm(dim, norm_eps)
|
|
||||||
self.mlp = MLP(dim, dim_ffn)
|
|
||||||
self.post_attention_norm = RMSNorm(dim, norm_eps)
|
|
||||||
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
x: Tensor,
|
|
||||||
rotary_emb: Tensor,
|
|
||||||
attention_mask: Optional[Tensor] = None,
|
|
||||||
paged_cache: Optional[KvcacheView] = None,
|
|
||||||
) -> Tensor:
|
|
||||||
attn_output = self.attention(
|
|
||||||
self.input_norm(x),
|
|
||||||
rotary_emb,
|
|
||||||
attention_mask,
|
|
||||||
paged_cache,
|
|
||||||
)
|
|
||||||
x = attn_output + x
|
|
||||||
x = self.mlp(self.post_attention_norm(x)) + x
|
|
||||||
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
class Embedding(nn.Module):
|
|
||||||
def __init__(self, vocab_size: int, embedding_dim: int):
|
|
||||||
super().__init__()
|
|
||||||
self.weight = nn.Parameter(torch.empty((vocab_size, embedding_dim)))
|
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
|
||||||
return F.embedding(x, self.weight)
|
|
||||||
+51
-69
@@ -1,93 +1,74 @@
|
|||||||
from typing import Any, Mapping, Optional
|
from typing import Any, Dict, Mapping, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.config.model_config import ModelConfig
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
from astrai.inference.core.cache import KvcacheView
|
from astrai.inference.core.cache import KVCache
|
||||||
from astrai.model.automodel import AutoModel
|
from astrai.model.automodel import AutoModel, ModelFactory
|
||||||
from astrai.model.module import (
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
DecoderBlock,
|
from astrai.model.components.embedding import Embedding
|
||||||
Embedding,
|
from astrai.model.components.linear import Linear
|
||||||
Linear,
|
from astrai.model.components.norm import RMSNorm
|
||||||
RMSNorm,
|
from astrai.model.components.rope import RotaryEmbedding
|
||||||
RotaryEmbedding,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def process_attention_mask(
|
def process_attention_mask(
|
||||||
input_tensor: Tensor,
|
input_mask: Optional[Tensor],
|
||||||
position_ids: Optional[Tensor],
|
|
||||||
input_mask: Optional[Tensor] = None,
|
|
||||||
is_causal: bool = False,
|
|
||||||
) -> Optional[Tensor]:
|
) -> Optional[Tensor]:
|
||||||
if position_ids is None:
|
|
||||||
return None
|
|
||||||
if input_mask is not None and input_mask.dim() > 2:
|
|
||||||
return input_mask
|
|
||||||
|
|
||||||
device = input_tensor.device
|
|
||||||
dtype = input_tensor.dtype
|
|
||||||
B, S = input_tensor.size()[:2]
|
|
||||||
T = position_ids.max().item() + 1
|
|
||||||
|
|
||||||
if input_mask is None:
|
if input_mask is None:
|
||||||
if position_ids.min().item() == 0 and is_causal:
|
return None
|
||||||
return None
|
if input_mask.dim() == 2:
|
||||||
pad = torch.ones(B, T, dtype=torch.bool, device=device)
|
return input_mask[:, None, None, :]
|
||||||
else:
|
if input_mask.dim() == 3:
|
||||||
pad = input_mask[:, :T].to(device=device, dtype=torch.bool)
|
return input_mask[:, None, :, :]
|
||||||
|
return input_mask
|
||||||
attend = pad.view(B, 1, T).expand(B, S, T).clone()
|
|
||||||
if is_causal:
|
|
||||||
attend &= position_ids.unsqueeze(-1) >= torch.arange(T, device=device)
|
|
||||||
|
|
||||||
return torch.full(
|
|
||||||
(B, 1, S, T), -torch.finfo(dtype).max / 2, dtype=dtype, device=device
|
|
||||||
).masked_fill_(attend.unsqueeze(1), 0.0)
|
|
||||||
|
|
||||||
|
|
||||||
@AutoModel.register("transformer")
|
@ModelFactory.register("autoregressive_lm")
|
||||||
class Transformer(AutoModel):
|
class AutoRegressiveLM(AutoModel):
|
||||||
"""Transformer language model with paged KV cache."""
|
"""Autoregressive language model with paged KV cache."""
|
||||||
|
|
||||||
def __init__(self, config: ModelConfig):
|
def __init__(self, config: AutoRegressiveLMConfig):
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
self.config = 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(
|
self.rotary_embedding = RotaryEmbedding(
|
||||||
config.dim // config.n_heads, config.max_len
|
rope_dim,
|
||||||
|
config.max_position_embeddings,
|
||||||
|
rope_base,
|
||||||
|
rope_scaling=config.rope_scaling,
|
||||||
|
)
|
||||||
|
self.embed_tokens = Embedding(
|
||||||
|
config.vocab_size,
|
||||||
|
config.hidden_size,
|
||||||
|
neftune_alpha=config.neftune_alpha,
|
||||||
)
|
)
|
||||||
self.embed_tokens = Embedding(config.vocab_size, config.dim)
|
|
||||||
|
|
||||||
self.layers = nn.ModuleList(
|
self.layers = nn.ModuleList(
|
||||||
[
|
[
|
||||||
DecoderBlock(
|
DecoderBlock(config, layer_id)
|
||||||
config.dim,
|
for layer_id in range(config.num_hidden_layers)
|
||||||
config.n_heads,
|
|
||||||
config.dim_ffn,
|
|
||||||
config.n_kv_heads,
|
|
||||||
config.norm_eps,
|
|
||||||
config.use_qk_norm,
|
|
||||||
config.use_gated_attention,
|
|
||||||
layer_id,
|
|
||||||
)
|
|
||||||
for layer_id in range(config.n_layers)
|
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
self.norm = RMSNorm(config.dim, config.norm_eps)
|
self.norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
||||||
self.lm_head = Linear(config.dim, config.vocab_size)
|
self.lm_head = Linear(config.hidden_size, config.vocab_size)
|
||||||
|
|
||||||
if self.config.tie_weight:
|
if self.config.tie_word_embeddings is True:
|
||||||
self.lm_head.weight = self.embed_tokens.weight
|
self.lm_head.weight = self.embed_tokens.weight
|
||||||
|
|
||||||
self._init_weights()
|
self.apply(self._init_weights)
|
||||||
|
|
||||||
def _init_weights(self):
|
def _init_weights(self, module):
|
||||||
for param in self.parameters():
|
if hasattr(module, "reset_parameters"):
|
||||||
if param.dim() > 1:
|
module.reset_parameters()
|
||||||
nn.init.normal_(param, mean=0.0, std=0.006)
|
|
||||||
|
|
||||||
def load_state_dict(self, state_dict: Mapping[str, Any], strict=True, assign=False):
|
def load_state_dict(self, state_dict: Mapping[str, Any], strict=True, assign=False):
|
||||||
lm_head_key = "lm_head.weight"
|
lm_head_key = "lm_head.weight"
|
||||||
@@ -95,7 +76,7 @@ class Transformer(AutoModel):
|
|||||||
|
|
||||||
state_dict = dict(state_dict)
|
state_dict = dict(state_dict)
|
||||||
|
|
||||||
if self.config.tie_weight:
|
if self.config.tie_word_embeddings is True:
|
||||||
# same tensor for embed and lm_head
|
# same tensor for embed and lm_head
|
||||||
if embed_key in state_dict:
|
if embed_key in state_dict:
|
||||||
state_dict[lm_head_key] = state_dict[embed_key]
|
state_dict[lm_head_key] = state_dict[embed_key]
|
||||||
@@ -111,7 +92,7 @@ class Transformer(AutoModel):
|
|||||||
destination=destination, prefix=prefix, keep_vars=keep_vars
|
destination=destination, prefix=prefix, keep_vars=keep_vars
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.config.tie_weight:
|
if self.config.tie_word_embeddings is True:
|
||||||
lm_head_key = prefix + "lm_head.weight"
|
lm_head_key = prefix + "lm_head.weight"
|
||||||
if lm_head_key in state_dict:
|
if lm_head_key in state_dict:
|
||||||
del state_dict[lm_head_key]
|
del state_dict[lm_head_key]
|
||||||
@@ -122,17 +103,18 @@ class Transformer(AutoModel):
|
|||||||
self,
|
self,
|
||||||
input_ids: Tensor,
|
input_ids: Tensor,
|
||||||
input_mask: Optional[Tensor] = None,
|
input_mask: Optional[Tensor] = None,
|
||||||
paged_cache: Optional[KvcacheView] = None,
|
kv_cache: Optional[KVCache] = None,
|
||||||
position_ids: Optional[Tensor] = None,
|
position_ids: Optional[Tensor] = None,
|
||||||
) -> Tensor:
|
) -> Dict[str, Tensor]:
|
||||||
assert input_ids.ndim == 2
|
assert input_ids.ndim == 2
|
||||||
|
|
||||||
x = self.embed_tokens(input_ids)
|
x = self.embed_tokens(input_ids)
|
||||||
rotary_emb = self.rotary_embedding(x, position_ids)
|
rotary_emb = self.rotary_embedding(x, position_ids)
|
||||||
attn_mask = process_attention_mask(x, position_ids, input_mask, is_causal=True)
|
attn_mask = process_attention_mask(input_mask)
|
||||||
|
use_sdpa_causal_mask = attn_mask is None
|
||||||
|
|
||||||
for layer in self.layers:
|
for layer in self.layers:
|
||||||
x = layer(x, rotary_emb, attn_mask, paged_cache)
|
x = layer(x, rotary_emb, attn_mask, kv_cache, use_sdpa_causal_mask)
|
||||||
|
|
||||||
hidden_states = self.norm(x)
|
hidden_states = self.norm(x)
|
||||||
logits = self.lm_head(hidden_states)
|
logits = self.lm_head(hidden_states)
|
||||||
|
|||||||
@@ -0,0 +1,38 @@
|
|||||||
|
"""Optimizer implementations and factory registration."""
|
||||||
|
|
||||||
|
from astrai.optim.composite import (
|
||||||
|
OptimizerFactory,
|
||||||
|
composite_state_dict,
|
||||||
|
composite_step,
|
||||||
|
composite_zero_grad,
|
||||||
|
refresh_param_groups,
|
||||||
|
)
|
||||||
|
from astrai.optim.mano_adamw import Mano, ManoAdamW
|
||||||
|
from astrai.optim.muon_adamw import MuonAdamW
|
||||||
|
from astrai.optim.nora_nadamw import (
|
||||||
|
NAdamW,
|
||||||
|
Nora,
|
||||||
|
NoraNAdamW,
|
||||||
|
OptimizerParameterGroups,
|
||||||
|
nora_direction,
|
||||||
|
nora_lr_scale,
|
||||||
|
partition_optimizer_parameters,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"Mano",
|
||||||
|
"ManoAdamW",
|
||||||
|
"MuonAdamW",
|
||||||
|
"NAdamW",
|
||||||
|
"Nora",
|
||||||
|
"NoraNAdamW",
|
||||||
|
"OptimizerFactory",
|
||||||
|
"OptimizerParameterGroups",
|
||||||
|
"composite_state_dict",
|
||||||
|
"composite_step",
|
||||||
|
"composite_zero_grad",
|
||||||
|
"nora_direction",
|
||||||
|
"nora_lr_scale",
|
||||||
|
"partition_optimizer_parameters",
|
||||||
|
"refresh_param_groups",
|
||||||
|
]
|
||||||
@@ -0,0 +1,71 @@
|
|||||||
|
"""Shared infrastructure for the optim package.
|
||||||
|
|
||||||
|
This module hosts two things:
|
||||||
|
|
||||||
|
* ``OptimizerFactory`` — the registry for built-in optimizers. Defining it
|
||||||
|
here (rather than in ``__init__.py``) lets each optimizer module import it
|
||||||
|
and register itself with a decorator, avoiding circular imports.
|
||||||
|
* Composite-optimizer helpers — ``step``/``zero_grad``/``state_dict``/
|
||||||
|
``param_groups`` delegation shared by every optimizer that routes different
|
||||||
|
parameter groups through distinct sub-optimizers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch.optim import Optimizer
|
||||||
|
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
|
||||||
|
|
||||||
|
class OptimizerFactory(BaseFactory[Optimizer]):
|
||||||
|
"""Factory for built-in training optimizers."""
|
||||||
|
|
||||||
|
|
||||||
|
def composite_step(
|
||||||
|
sub_optimizers: list[Optimizer],
|
||||||
|
closure=None,
|
||||||
|
) -> torch.Tensor | None:
|
||||||
|
"""Run ``step`` on every sub-optimizer, invoking the closure once.
|
||||||
|
|
||||||
|
The closure (if given) is executed inside ``torch.enable_grad`` exactly
|
||||||
|
once before any sub-optimizer steps, matching the contract of a single
|
||||||
|
``Optimizer.step``. Sub-optimizers receive ``None`` so they do not
|
||||||
|
re-execute it.
|
||||||
|
"""
|
||||||
|
loss = None
|
||||||
|
if closure is not None:
|
||||||
|
with torch.enable_grad():
|
||||||
|
loss = closure()
|
||||||
|
for sub in sub_optimizers:
|
||||||
|
sub.step()
|
||||||
|
return loss
|
||||||
|
|
||||||
|
|
||||||
|
def composite_zero_grad(
|
||||||
|
sub_optimizers: list[Optimizer],
|
||||||
|
set_to_none: bool = True,
|
||||||
|
) -> None:
|
||||||
|
for sub in sub_optimizers:
|
||||||
|
sub.zero_grad(set_to_none=set_to_none)
|
||||||
|
|
||||||
|
|
||||||
|
def composite_state_dict(
|
||||||
|
named_sub_optimizers: dict[str, Optimizer | None],
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Serialize sub-optimizers, preserving ``None`` slots."""
|
||||||
|
return {
|
||||||
|
name: sub.state_dict() if sub is not None else None
|
||||||
|
for name, sub in named_sub_optimizers.items()
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def refresh_param_groups(
|
||||||
|
sub_optimizers: list[Optimizer],
|
||||||
|
) -> list[dict]:
|
||||||
|
"""Concatenate param_groups from every non-None sub-optimizer."""
|
||||||
|
groups: list[dict] = []
|
||||||
|
for sub in sub_optimizers:
|
||||||
|
if sub is not None:
|
||||||
|
groups.extend(sub.param_groups)
|
||||||
|
return groups
|
||||||
@@ -0,0 +1,214 @@
|
|||||||
|
"""Mano manifold optimizer combined with AdamW.
|
||||||
|
|
||||||
|
Mano projects the momentum onto the tangent space of the Oblique manifold
|
||||||
|
(axis-wise tangent projection) and normalizes it, replacing the expensive
|
||||||
|
Newton-Schulz iteration in Muon with a cheaper manifold normalization.
|
||||||
|
|
||||||
|
Reference: https://arxiv.org/abs/2601.23000
|
||||||
|
"""
|
||||||
|
|
||||||
|
import math
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import nn, optim
|
||||||
|
from torch.optim import Optimizer
|
||||||
|
|
||||||
|
from astrai.optim.composite import (
|
||||||
|
OptimizerFactory,
|
||||||
|
composite_state_dict,
|
||||||
|
composite_step,
|
||||||
|
composite_zero_grad,
|
||||||
|
refresh_param_groups,
|
||||||
|
)
|
||||||
|
from astrai.optim.nora_nadamw import partition_optimizer_parameters
|
||||||
|
|
||||||
|
|
||||||
|
class Mano(Optimizer):
|
||||||
|
"""Manifold Normalized Optimizer for two-dimensional matrices.
|
||||||
|
|
||||||
|
Each step alternates the projection axis (dim 0 / dim 1) to restrike the
|
||||||
|
manifold along both rows and columns. The tangent momentum is computed
|
||||||
|
without normalizing the parameter itself (v2 simplification) and the
|
||||||
|
epsilon is added (not clamped) to the norm denominator.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
params,
|
||||||
|
lr: float = 1e-3,
|
||||||
|
weight_decay: float = 0.1,
|
||||||
|
momentum: float = 0.95,
|
||||||
|
nesterov: bool = True,
|
||||||
|
eps: float = 1e-8,
|
||||||
|
):
|
||||||
|
if lr < 0:
|
||||||
|
raise ValueError(f"Invalid learning rate: {lr}")
|
||||||
|
if weight_decay < 0:
|
||||||
|
raise ValueError(f"Invalid weight decay: {weight_decay}")
|
||||||
|
if not 0 <= momentum <= 1:
|
||||||
|
raise ValueError(f"Invalid momentum: {momentum}")
|
||||||
|
if eps <= 0:
|
||||||
|
raise ValueError(f"Invalid epsilon: {eps}")
|
||||||
|
|
||||||
|
defaults = {
|
||||||
|
"lr": lr,
|
||||||
|
"weight_decay": weight_decay,
|
||||||
|
"momentum": momentum,
|
||||||
|
"nesterov": nesterov,
|
||||||
|
"eps": eps,
|
||||||
|
"steps": 0,
|
||||||
|
}
|
||||||
|
super().__init__(params, defaults)
|
||||||
|
for group in self.param_groups:
|
||||||
|
for param in group["params"]:
|
||||||
|
if param.ndim != 2:
|
||||||
|
raise ValueError(
|
||||||
|
f"Mano only supports 2D matrices, got shape {tuple(param.shape)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def step(self, closure=None):
|
||||||
|
loss = None
|
||||||
|
if closure is not None:
|
||||||
|
with torch.enable_grad():
|
||||||
|
loss = closure()
|
||||||
|
|
||||||
|
for group in self.param_groups:
|
||||||
|
lr = group["lr"]
|
||||||
|
weight_decay = group["weight_decay"]
|
||||||
|
momentum = group["momentum"]
|
||||||
|
nesterov = group["nesterov"]
|
||||||
|
eps = group["eps"]
|
||||||
|
dim = int(group["steps"] % 2)
|
||||||
|
|
||||||
|
for param in group["params"]:
|
||||||
|
if param.grad is None:
|
||||||
|
continue
|
||||||
|
if param.grad.is_sparse:
|
||||||
|
raise RuntimeError("Mano does not support sparse gradients")
|
||||||
|
|
||||||
|
grad = param.grad
|
||||||
|
state = self.state[param]
|
||||||
|
momentum_buffer = state.get("momentum_buffer")
|
||||||
|
if momentum_buffer is None:
|
||||||
|
momentum_buffer = torch.zeros_like(grad)
|
||||||
|
momentum_buffer.mul_(momentum).add_(grad)
|
||||||
|
update = (
|
||||||
|
grad.add(momentum_buffer, alpha=momentum)
|
||||||
|
if nesterov
|
||||||
|
else momentum_buffer
|
||||||
|
)
|
||||||
|
|
||||||
|
tangent = update - (
|
||||||
|
torch.sum(update * param.data, dim=dim, keepdim=True) * param.data
|
||||||
|
)
|
||||||
|
direction = tangent / (
|
||||||
|
torch.norm(tangent, p=2, dim=dim, keepdim=True) + eps
|
||||||
|
)
|
||||||
|
|
||||||
|
if weight_decay != 0:
|
||||||
|
param.mul_(1 - lr * weight_decay)
|
||||||
|
adjusted_lr = lr * 0.2 * math.sqrt(direction.shape[dim])
|
||||||
|
param.add_(direction, alpha=-adjusted_lr)
|
||||||
|
state["momentum_buffer"] = momentum_buffer
|
||||||
|
|
||||||
|
group["steps"] += 1
|
||||||
|
|
||||||
|
return loss
|
||||||
|
|
||||||
|
|
||||||
|
@OptimizerFactory.register("mano_adamw")
|
||||||
|
class ManoAdamW(Optimizer):
|
||||||
|
"""Mano for internal linear weights and AdamW for remaining parameters."""
|
||||||
|
|
||||||
|
optimizer_name = "mano_adamw"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: nn.Module,
|
||||||
|
lr: float = 3e-4,
|
||||||
|
weight_decay: float = 0.1,
|
||||||
|
momentum: float = 0.95,
|
||||||
|
nesterov: bool = True,
|
||||||
|
):
|
||||||
|
groups = partition_optimizer_parameters(model)
|
||||||
|
all_params = [
|
||||||
|
*groups.nora,
|
||||||
|
*groups.nadamw_decay,
|
||||||
|
*groups.nadamw_no_decay,
|
||||||
|
]
|
||||||
|
if not all_params:
|
||||||
|
raise ValueError(
|
||||||
|
"Cannot build an optimizer for a model with no trainable parameters"
|
||||||
|
)
|
||||||
|
super().__init__(all_params, {})
|
||||||
|
|
||||||
|
self.mano = (
|
||||||
|
Mano(
|
||||||
|
groups.nora,
|
||||||
|
lr=lr,
|
||||||
|
weight_decay=weight_decay,
|
||||||
|
momentum=momentum,
|
||||||
|
nesterov=nesterov,
|
||||||
|
)
|
||||||
|
if groups.nora
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
|
adamw_groups = []
|
||||||
|
if groups.nadamw_decay:
|
||||||
|
adamw_groups.append(
|
||||||
|
{"params": groups.nadamw_decay, "weight_decay": weight_decay}
|
||||||
|
)
|
||||||
|
if groups.nadamw_no_decay:
|
||||||
|
adamw_groups.append({"params": groups.nadamw_no_decay, "weight_decay": 0.0})
|
||||||
|
self.adamw = (
|
||||||
|
optim.AdamW(
|
||||||
|
adamw_groups,
|
||||||
|
lr=lr,
|
||||||
|
betas=(0.9, 0.95),
|
||||||
|
fused=True,
|
||||||
|
)
|
||||||
|
if adamw_groups
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
self.param_groups = refresh_param_groups([self.mano, self.adamw])
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def step(self, closure=None):
|
||||||
|
return composite_step(
|
||||||
|
[opt for opt in (self.mano, self.adamw) if opt is not None],
|
||||||
|
closure,
|
||||||
|
)
|
||||||
|
|
||||||
|
def zero_grad(self, set_to_none: bool = True):
|
||||||
|
composite_zero_grad(
|
||||||
|
[opt for opt in (self.mano, self.adamw) if opt is not None],
|
||||||
|
set_to_none,
|
||||||
|
)
|
||||||
|
|
||||||
|
def state_dict(self) -> dict:
|
||||||
|
return composite_state_dict({"mano": self.mano, "adamw": self.adamw})
|
||||||
|
|
||||||
|
def load_state_dict(self, state_dict: dict):
|
||||||
|
if "muon" in state_dict or "nora" in state_dict:
|
||||||
|
raise ValueError(
|
||||||
|
"Checkpoint uses a different optimizer; select the matching "
|
||||||
|
"--optimizer to resume it"
|
||||||
|
)
|
||||||
|
if "mano" not in state_dict or "adamw" not in state_dict:
|
||||||
|
raise ValueError(
|
||||||
|
"Checkpoint optimizer state is not compatible with mano_adamw"
|
||||||
|
)
|
||||||
|
|
||||||
|
saved_mano = state_dict["mano"]
|
||||||
|
saved_adamw = state_dict["adamw"]
|
||||||
|
if (self.mano is None) != (saved_mano is None):
|
||||||
|
raise ValueError("Checkpoint Mano parameter groups do not match the model")
|
||||||
|
if (self.adamw is None) != (saved_adamw is None):
|
||||||
|
raise ValueError("Checkpoint AdamW parameter groups do not match the model")
|
||||||
|
if self.mano is not None:
|
||||||
|
self.mano.load_state_dict(saved_mano)
|
||||||
|
if self.adamw is not None:
|
||||||
|
self.adamw.load_state_dict(saved_adamw)
|
||||||
|
self.param_groups = refresh_param_groups([self.mano, self.adamw])
|
||||||
@@ -0,0 +1,95 @@
|
|||||||
|
"""Legacy Muon + AdamW combined optimizer."""
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor, nn, optim
|
||||||
|
|
||||||
|
from astrai.optim.composite import (
|
||||||
|
OptimizerFactory,
|
||||||
|
composite_state_dict,
|
||||||
|
composite_step,
|
||||||
|
composite_zero_grad,
|
||||||
|
refresh_param_groups,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@OptimizerFactory.register("muon_adamw")
|
||||||
|
class MuonAdamW(optim.Optimizer):
|
||||||
|
"""Combined Muon (matrix) + AdamW (non-matrix) optimizer."""
|
||||||
|
|
||||||
|
optimizer_name = "muon_adamw"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: nn.Module,
|
||||||
|
lr: float = 3e-4,
|
||||||
|
weight_decay: float = 0.1,
|
||||||
|
momentum: float = 0.95,
|
||||||
|
nesterov: bool = True,
|
||||||
|
ns_steps: int = 5,
|
||||||
|
adjust_lr_fn: str = "match_rms_adamw",
|
||||||
|
):
|
||||||
|
defaults = {
|
||||||
|
"lr": lr,
|
||||||
|
"weight_decay": weight_decay,
|
||||||
|
"momentum": momentum,
|
||||||
|
"nesterov": nesterov,
|
||||||
|
"ns_steps": ns_steps,
|
||||||
|
"adjust_lr_fn": adjust_lr_fn,
|
||||||
|
}
|
||||||
|
params = [param for param in model.parameters() if param.requires_grad]
|
||||||
|
super().__init__(params, defaults)
|
||||||
|
|
||||||
|
matrix_params: list[Tensor] = []
|
||||||
|
other_params: list[Tensor] = []
|
||||||
|
for name, param in model.named_parameters():
|
||||||
|
if not param.requires_grad:
|
||||||
|
continue
|
||||||
|
if (
|
||||||
|
param.dim() >= 2
|
||||||
|
and "norm" not in name
|
||||||
|
and "bias" not in name
|
||||||
|
and "embed" not in name
|
||||||
|
and "lm_head" not in name
|
||||||
|
):
|
||||||
|
matrix_params.append(param)
|
||||||
|
else:
|
||||||
|
other_params.append(param)
|
||||||
|
|
||||||
|
self.muon = optim.Muon(
|
||||||
|
matrix_params,
|
||||||
|
lr=lr,
|
||||||
|
weight_decay=weight_decay,
|
||||||
|
momentum=momentum,
|
||||||
|
nesterov=nesterov,
|
||||||
|
ns_steps=ns_steps,
|
||||||
|
adjust_lr_fn=adjust_lr_fn,
|
||||||
|
)
|
||||||
|
self.adamw = optim.AdamW(
|
||||||
|
[{"params": other_params, "weight_decay": 0.0}],
|
||||||
|
lr=lr,
|
||||||
|
betas=(0.9, 0.95),
|
||||||
|
fused=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.param_groups = refresh_param_groups([self.muon, self.adamw])
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def step(self, closure=None):
|
||||||
|
return composite_step([self.muon, self.adamw], closure)
|
||||||
|
|
||||||
|
def zero_grad(self, set_to_none: bool = True):
|
||||||
|
composite_zero_grad([self.muon, self.adamw], set_to_none)
|
||||||
|
|
||||||
|
def state_dict(self) -> dict[str, Any]:
|
||||||
|
return composite_state_dict({"muon": self.muon, "adamw": self.adamw})
|
||||||
|
|
||||||
|
def load_state_dict(self, state_dict: dict[str, Any]):
|
||||||
|
if "muon" not in state_dict or "adamw" not in state_dict:
|
||||||
|
raise ValueError(
|
||||||
|
"Checkpoint optimizer state is not compatible with muon_adamw"
|
||||||
|
)
|
||||||
|
self.muon.load_state_dict(state_dict["muon"])
|
||||||
|
self.adamw.load_state_dict(state_dict["adamw"])
|
||||||
|
self.param_groups = refresh_param_groups([self.muon, self.adamw])
|
||||||
@@ -0,0 +1,372 @@
|
|||||||
|
"""Nora matrix optimizer combined with Nesterov AdamW."""
|
||||||
|
|
||||||
|
import math
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor, nn
|
||||||
|
from torch.distributed.tensor import DTensor, Shard
|
||||||
|
from torch.optim import Optimizer
|
||||||
|
|
||||||
|
from astrai.model.components.embedding import Embedding
|
||||||
|
from astrai.model.components.linear import Linear
|
||||||
|
from astrai.model.components.lora import LoRALinear
|
||||||
|
from astrai.model.components.norm import RMSNorm
|
||||||
|
from astrai.optim.composite import (
|
||||||
|
OptimizerFactory,
|
||||||
|
composite_state_dict,
|
||||||
|
composite_step,
|
||||||
|
composite_zero_grad,
|
||||||
|
refresh_param_groups,
|
||||||
|
)
|
||||||
|
|
||||||
|
NORA_EPS = 1e-10
|
||||||
|
|
||||||
|
|
||||||
|
def _row_normalize(tensor: Tensor, eps: float) -> Tensor:
|
||||||
|
return tensor / tensor.norm(dim=-1, keepdim=True).clamp(min=eps)
|
||||||
|
|
||||||
|
|
||||||
|
def nora_direction(update: Tensor, param: Tensor, eps: float = NORA_EPS) -> Tensor:
|
||||||
|
"""Project an update onto each parameter row's tangent space and normalize."""
|
||||||
|
theta_hat = _row_normalize(param.to(torch.float32), eps)
|
||||||
|
update_fp32 = update.to(torch.float32)
|
||||||
|
radial = (update_fp32 * theta_hat).sum(dim=-1, keepdim=True) * theta_hat
|
||||||
|
direction = _row_normalize(update_fp32 - radial, eps)
|
||||||
|
return direction.to(update.dtype)
|
||||||
|
|
||||||
|
|
||||||
|
def nora_lr_scale(lr: float, shape: torch.Size) -> float:
|
||||||
|
"""Scale Nora's LR for tall ``[d_out, d_in]`` linear weights."""
|
||||||
|
return lr * math.sqrt(max(1.0, shape[-2] / shape[-1]))
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_complete_rows(param: Tensor) -> None:
|
||||||
|
if not isinstance(param, DTensor):
|
||||||
|
return
|
||||||
|
last_dim = param.ndim - 1
|
||||||
|
for placement in param.placements:
|
||||||
|
if isinstance(placement, Shard) and placement.dim % param.ndim == last_dim:
|
||||||
|
raise ValueError(
|
||||||
|
"Nora requires complete parameter rows, but this DTensor is sharded "
|
||||||
|
"along its last dimension"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class Nora(Optimizer):
|
||||||
|
"""Normalized Orthogonal Row Alignment for two-dimensional matrices."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
params,
|
||||||
|
lr: float = 5e-3,
|
||||||
|
weight_decay: float = 0.0,
|
||||||
|
momentum: float = 0.95,
|
||||||
|
beta: float = 0.95,
|
||||||
|
nesterov: bool = True,
|
||||||
|
eps: float = NORA_EPS,
|
||||||
|
):
|
||||||
|
if lr < 0:
|
||||||
|
raise ValueError(f"Invalid learning rate: {lr}")
|
||||||
|
if weight_decay < 0:
|
||||||
|
raise ValueError(f"Invalid weight decay: {weight_decay}")
|
||||||
|
if not 0 <= momentum <= 1:
|
||||||
|
raise ValueError(f"Invalid momentum: {momentum}")
|
||||||
|
if not 0 <= beta < 1:
|
||||||
|
raise ValueError(f"Invalid beta: {beta}")
|
||||||
|
if eps <= 0:
|
||||||
|
raise ValueError(f"Invalid epsilon: {eps}")
|
||||||
|
|
||||||
|
defaults = {
|
||||||
|
"lr": lr,
|
||||||
|
"weight_decay": weight_decay,
|
||||||
|
"momentum": momentum,
|
||||||
|
"beta": beta,
|
||||||
|
"nesterov": nesterov,
|
||||||
|
"eps": eps,
|
||||||
|
}
|
||||||
|
super().__init__(params, defaults)
|
||||||
|
for group in self.param_groups:
|
||||||
|
for param in group["params"]:
|
||||||
|
if param.ndim != 2:
|
||||||
|
raise ValueError(
|
||||||
|
f"Nora only supports 2D matrices, got shape {tuple(param.shape)}"
|
||||||
|
)
|
||||||
|
_validate_complete_rows(param)
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def step(self, closure=None):
|
||||||
|
loss = None
|
||||||
|
if closure is not None:
|
||||||
|
with torch.enable_grad():
|
||||||
|
loss = closure()
|
||||||
|
|
||||||
|
for group in self.param_groups:
|
||||||
|
lr = group["lr"]
|
||||||
|
weight_decay = group["weight_decay"]
|
||||||
|
momentum = group["momentum"]
|
||||||
|
beta = group["beta"]
|
||||||
|
nesterov = group["nesterov"]
|
||||||
|
eps = group["eps"]
|
||||||
|
for param in group["params"]:
|
||||||
|
if param.grad is None:
|
||||||
|
continue
|
||||||
|
if param.grad.is_sparse:
|
||||||
|
raise RuntimeError("Nora does not support sparse gradients")
|
||||||
|
|
||||||
|
grad = param.grad
|
||||||
|
state = self.state[param]
|
||||||
|
momentum_buffer = state.get("momentum_buffer")
|
||||||
|
if momentum_buffer is None:
|
||||||
|
momentum_buffer = torch.zeros_like(grad)
|
||||||
|
momentum_buffer.lerp_(grad, 1 - beta)
|
||||||
|
update = (
|
||||||
|
grad.lerp(momentum_buffer, momentum)
|
||||||
|
if nesterov
|
||||||
|
else momentum_buffer
|
||||||
|
)
|
||||||
|
direction = nora_direction(update, param, eps)
|
||||||
|
|
||||||
|
if weight_decay != 0:
|
||||||
|
param.mul_(1 - lr * weight_decay)
|
||||||
|
param.add_(direction, alpha=-nora_lr_scale(lr, param.shape))
|
||||||
|
state["momentum_buffer"] = momentum_buffer
|
||||||
|
|
||||||
|
return loss
|
||||||
|
|
||||||
|
|
||||||
|
class NAdamW(Optimizer):
|
||||||
|
"""AdamW using the reference Nesterov first-moment update."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
params,
|
||||||
|
lr: float = 3e-4,
|
||||||
|
betas: tuple[float, float] = (0.9, 0.999),
|
||||||
|
eps: float = 1e-8,
|
||||||
|
weight_decay: float = 0.1,
|
||||||
|
):
|
||||||
|
beta1, beta2 = betas
|
||||||
|
if lr < 0:
|
||||||
|
raise ValueError(f"Invalid learning rate: {lr}")
|
||||||
|
if not 0 <= beta1 < 1 or not 0 <= beta2 < 1:
|
||||||
|
raise ValueError(f"Invalid betas: {betas}")
|
||||||
|
if eps <= 0:
|
||||||
|
raise ValueError(f"Invalid epsilon: {eps}")
|
||||||
|
if weight_decay < 0:
|
||||||
|
raise ValueError(f"Invalid weight decay: {weight_decay}")
|
||||||
|
defaults = {
|
||||||
|
"lr": lr,
|
||||||
|
"betas": betas,
|
||||||
|
"eps": eps,
|
||||||
|
"weight_decay": weight_decay,
|
||||||
|
}
|
||||||
|
super().__init__(params, defaults)
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def step(self, closure=None):
|
||||||
|
loss = None
|
||||||
|
if closure is not None:
|
||||||
|
with torch.enable_grad():
|
||||||
|
loss = closure()
|
||||||
|
|
||||||
|
for group in self.param_groups:
|
||||||
|
beta1, beta2 = group["betas"]
|
||||||
|
eps = group["eps"]
|
||||||
|
lr = group["lr"]
|
||||||
|
weight_decay = group["weight_decay"]
|
||||||
|
for param in group["params"]:
|
||||||
|
if param.grad is None:
|
||||||
|
continue
|
||||||
|
if param.grad.is_sparse:
|
||||||
|
raise RuntimeError("NAdamW does not support sparse gradients")
|
||||||
|
|
||||||
|
grad = param.grad
|
||||||
|
state = self.state[param]
|
||||||
|
if not state:
|
||||||
|
state["step"] = 0
|
||||||
|
state["m"] = torch.zeros_like(param)
|
||||||
|
state["v"] = torch.zeros_like(param)
|
||||||
|
|
||||||
|
state["step"] += 1
|
||||||
|
first_moment = state["m"]
|
||||||
|
second_moment = state["v"]
|
||||||
|
first_moment.mul_(beta1).add_(grad, alpha=1 - beta1)
|
||||||
|
second_moment.mul_(beta2).addcmul_(grad, grad, value=1 - beta2)
|
||||||
|
|
||||||
|
bias_correction1 = 1 - beta1 ** state["step"]
|
||||||
|
bias_correction2 = 1 - beta2 ** state["step"]
|
||||||
|
nesterov_moment = (
|
||||||
|
beta1 * first_moment + (1 - beta1) * grad
|
||||||
|
) / bias_correction1
|
||||||
|
corrected_second_moment = second_moment / bias_correction2
|
||||||
|
|
||||||
|
if weight_decay != 0:
|
||||||
|
param.mul_(1 - lr * weight_decay)
|
||||||
|
param.addcdiv_(
|
||||||
|
nesterov_moment,
|
||||||
|
corrected_second_moment.sqrt().add_(eps),
|
||||||
|
value=-lr,
|
||||||
|
)
|
||||||
|
|
||||||
|
return loss
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class OptimizerParameterGroups:
|
||||||
|
nora: list[Tensor]
|
||||||
|
nadamw_decay: list[Tensor]
|
||||||
|
nadamw_no_decay: list[Tensor]
|
||||||
|
|
||||||
|
|
||||||
|
def partition_optimizer_parameters(model: nn.Module) -> OptimizerParameterGroups:
|
||||||
|
"""Partition trainable parameters by module role and parameter identity."""
|
||||||
|
nora_ids: set[int] = set()
|
||||||
|
no_decay_ids: set[int] = set()
|
||||||
|
|
||||||
|
for module_name, module in model.named_modules():
|
||||||
|
if isinstance(module, LoRALinear):
|
||||||
|
for param in module.parameters(recurse=False):
|
||||||
|
if param.requires_grad:
|
||||||
|
no_decay_ids.add(id(param))
|
||||||
|
continue
|
||||||
|
|
||||||
|
if isinstance(module, (Embedding, RMSNorm)):
|
||||||
|
for param in module.parameters(recurse=False):
|
||||||
|
if param.requires_grad:
|
||||||
|
no_decay_ids.add(id(param))
|
||||||
|
continue
|
||||||
|
|
||||||
|
if not isinstance(module, Linear):
|
||||||
|
continue
|
||||||
|
|
||||||
|
if module.bias is not None and module.bias.requires_grad:
|
||||||
|
no_decay_ids.add(id(module.bias))
|
||||||
|
if not module.weight.requires_grad:
|
||||||
|
continue
|
||||||
|
if module_name.rsplit(".", 1)[-1] == "lm_head":
|
||||||
|
no_decay_ids.add(id(module.weight))
|
||||||
|
elif module.weight.ndim == 2:
|
||||||
|
nora_ids.add(id(module.weight))
|
||||||
|
|
||||||
|
nora: list[Tensor] = []
|
||||||
|
nadamw_decay: list[Tensor] = []
|
||||||
|
nadamw_no_decay: list[Tensor] = []
|
||||||
|
seen: set[int] = set()
|
||||||
|
for param in model.parameters():
|
||||||
|
param_id = id(param)
|
||||||
|
if not param.requires_grad or param_id in seen:
|
||||||
|
continue
|
||||||
|
seen.add(param_id)
|
||||||
|
if param_id in no_decay_ids or param.ndim <= 1:
|
||||||
|
nadamw_no_decay.append(param)
|
||||||
|
elif param_id in nora_ids:
|
||||||
|
nora.append(param)
|
||||||
|
else:
|
||||||
|
nadamw_decay.append(param)
|
||||||
|
|
||||||
|
trainable_ids = {id(param) for param in model.parameters() if param.requires_grad}
|
||||||
|
grouped_ids = {id(param) for param in [*nora, *nadamw_decay, *nadamw_no_decay]}
|
||||||
|
if grouped_ids != trainable_ids:
|
||||||
|
missing = len(trainable_ids - grouped_ids)
|
||||||
|
extra = len(grouped_ids - trainable_ids)
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Optimizer parameter partition is incomplete: missing={missing}, extra={extra}"
|
||||||
|
)
|
||||||
|
|
||||||
|
return OptimizerParameterGroups(nora, nadamw_decay, nadamw_no_decay)
|
||||||
|
|
||||||
|
|
||||||
|
@OptimizerFactory.register("nora_nadamw")
|
||||||
|
class NoraNAdamW(Optimizer):
|
||||||
|
"""Nora for internal linear weights and NAdamW for remaining parameters."""
|
||||||
|
|
||||||
|
optimizer_name = "nora_nadamw"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: nn.Module,
|
||||||
|
lr: float = 3e-4,
|
||||||
|
weight_decay: float = 0.1,
|
||||||
|
nora_lr: float = 5e-3,
|
||||||
|
nora_weight_decay: float = 0.0,
|
||||||
|
nora_beta: float = 0.95,
|
||||||
|
nora_momentum: float = 0.95,
|
||||||
|
):
|
||||||
|
groups = partition_optimizer_parameters(model)
|
||||||
|
all_params = [
|
||||||
|
*groups.nora,
|
||||||
|
*groups.nadamw_decay,
|
||||||
|
*groups.nadamw_no_decay,
|
||||||
|
]
|
||||||
|
if not all_params:
|
||||||
|
raise ValueError(
|
||||||
|
"Cannot build an optimizer for a model with no trainable parameters"
|
||||||
|
)
|
||||||
|
super().__init__(all_params, {})
|
||||||
|
|
||||||
|
self.nora = (
|
||||||
|
Nora(
|
||||||
|
groups.nora,
|
||||||
|
lr=nora_lr,
|
||||||
|
weight_decay=nora_weight_decay,
|
||||||
|
momentum=nora_momentum,
|
||||||
|
beta=nora_beta,
|
||||||
|
)
|
||||||
|
if groups.nora
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
|
nadamw_groups = []
|
||||||
|
if groups.nadamw_decay:
|
||||||
|
nadamw_groups.append(
|
||||||
|
{"params": groups.nadamw_decay, "weight_decay": weight_decay}
|
||||||
|
)
|
||||||
|
if groups.nadamw_no_decay:
|
||||||
|
nadamw_groups.append(
|
||||||
|
{"params": groups.nadamw_no_decay, "weight_decay": 0.0}
|
||||||
|
)
|
||||||
|
self.nadamw = NAdamW(nadamw_groups, lr=lr) if nadamw_groups else None
|
||||||
|
self.param_groups = refresh_param_groups([self.nora, self.nadamw])
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def step(self, closure=None):
|
||||||
|
return composite_step(
|
||||||
|
[opt for opt in (self.nora, self.nadamw) if opt is not None],
|
||||||
|
closure,
|
||||||
|
)
|
||||||
|
|
||||||
|
def zero_grad(self, set_to_none: bool = True):
|
||||||
|
composite_zero_grad(
|
||||||
|
[opt for opt in (self.nora, self.nadamw) if opt is not None],
|
||||||
|
set_to_none,
|
||||||
|
)
|
||||||
|
|
||||||
|
def state_dict(self) -> dict[str, Any]:
|
||||||
|
return composite_state_dict({"nora": self.nora, "nadamw": self.nadamw})
|
||||||
|
|
||||||
|
def load_state_dict(self, state_dict: dict[str, Any]):
|
||||||
|
if "muon" in state_dict or "adamw" in state_dict:
|
||||||
|
raise ValueError(
|
||||||
|
"Checkpoint uses muon_adamw state; select optimizer='muon_adamw' "
|
||||||
|
"to resume it"
|
||||||
|
)
|
||||||
|
if "nora" not in state_dict or "nadamw" not in state_dict:
|
||||||
|
raise ValueError(
|
||||||
|
"Checkpoint optimizer state is not compatible with nora_nadamw"
|
||||||
|
)
|
||||||
|
|
||||||
|
saved_nora = state_dict["nora"]
|
||||||
|
saved_nadamw = state_dict["nadamw"]
|
||||||
|
if (self.nora is None) != (saved_nora is None):
|
||||||
|
raise ValueError("Checkpoint Nora parameter groups do not match the model")
|
||||||
|
if (self.nadamw is None) != (saved_nadamw is None):
|
||||||
|
raise ValueError(
|
||||||
|
"Checkpoint NAdamW parameter groups do not match the model"
|
||||||
|
)
|
||||||
|
if self.nora is not None:
|
||||||
|
self.nora.load_state_dict(saved_nora)
|
||||||
|
if self.nadamw is not None:
|
||||||
|
self.nadamw.load_state_dict(saved_nadamw)
|
||||||
|
self.param_groups = refresh_param_groups([self.nora, self.nadamw])
|
||||||
@@ -1,4 +1,15 @@
|
|||||||
from astrai.parallel.module import ColumnParallelLinear, RowParallelLinear
|
from astrai.parallel.executor import (
|
||||||
|
AccumOptimizer,
|
||||||
|
AccumScheduler,
|
||||||
|
BaseExecutor,
|
||||||
|
DDPExecutor,
|
||||||
|
ExecutorFactory,
|
||||||
|
FSDPExecutor,
|
||||||
|
GradientState,
|
||||||
|
NoneExecutor,
|
||||||
|
broadcast_state_dict,
|
||||||
|
create_ref_model,
|
||||||
|
)
|
||||||
from astrai.parallel.setup import (
|
from astrai.parallel.setup import (
|
||||||
get_current_device,
|
get_current_device,
|
||||||
get_rank,
|
get_rank,
|
||||||
@@ -15,6 +26,14 @@ __all__ = [
|
|||||||
"only_on_rank",
|
"only_on_rank",
|
||||||
"setup_parallel",
|
"setup_parallel",
|
||||||
"spawn_parallel_fn",
|
"spawn_parallel_fn",
|
||||||
"RowParallelLinear",
|
"ExecutorFactory",
|
||||||
"ColumnParallelLinear",
|
"BaseExecutor",
|
||||||
|
"GradientState",
|
||||||
|
"AccumOptimizer",
|
||||||
|
"AccumScheduler",
|
||||||
|
"NoneExecutor",
|
||||||
|
"DDPExecutor",
|
||||||
|
"FSDPExecutor",
|
||||||
|
"create_ref_model",
|
||||||
|
"broadcast_state_dict",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -0,0 +1,428 @@
|
|||||||
|
"""Unified training executor — parallel strategy + gradient accumulation."""
|
||||||
|
|
||||||
|
import contextlib
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from typing import Any, Callable, Dict, Optional, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
|
import torch.nn as nn
|
||||||
|
from torch.distributed.fsdp import (
|
||||||
|
FSDPModule,
|
||||||
|
fully_shard,
|
||||||
|
)
|
||||||
|
from torch.distributed.tensor import DTensor
|
||||||
|
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||||
|
from torch.optim import Optimizer
|
||||||
|
from torch.optim.lr_scheduler import LRScheduler
|
||||||
|
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.parallel.setup import get_rank, get_world_size
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def broadcast_state_dict(
|
||||||
|
state_dict: Optional[Dict[str, torch.Tensor]],
|
||||||
|
src: int = 0,
|
||||||
|
) -> Optional[Dict[str, torch.Tensor]]:
|
||||||
|
"""Broadcast a state_dict from *src* rank to all ranks.
|
||||||
|
|
||||||
|
Tensors stay on their original device (GPU) for the broadcast.
|
||||||
|
All ranks must call this collectively.
|
||||||
|
|
||||||
|
On non-distributed runs, returns *state_dict* unchanged.
|
||||||
|
"""
|
||||||
|
if not dist.is_initialized() or dist.get_world_size() == 1:
|
||||||
|
return state_dict
|
||||||
|
|
||||||
|
rank = dist.get_rank()
|
||||||
|
|
||||||
|
# Broadcast metadata (keys, shapes, dtypes, device) so non-src ranks
|
||||||
|
# can allocate matching empty tensors on the correct device.
|
||||||
|
if rank == src:
|
||||||
|
device = next(iter(state_dict.values())).device
|
||||||
|
metadata = [
|
||||||
|
(k, tuple(v.shape), v.dtype, str(device)) for k, v in state_dict.items()
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
metadata = None
|
||||||
|
metadata_list = [metadata]
|
||||||
|
dist.broadcast_object_list(metadata_list, src=src)
|
||||||
|
metadata = metadata_list[0]
|
||||||
|
|
||||||
|
# Non-src ranks allocate empty tensors with the broadcasted metadata.
|
||||||
|
if rank != src:
|
||||||
|
state_dict = {
|
||||||
|
k: torch.empty(s, dtype=d, device=torch.device(dev))
|
||||||
|
for k, s, d, dev in metadata
|
||||||
|
}
|
||||||
|
|
||||||
|
# Broadcast each tensor in-place.
|
||||||
|
for tensor in state_dict.values():
|
||||||
|
dist.broadcast(tensor, src=src)
|
||||||
|
|
||||||
|
return state_dict
|
||||||
|
|
||||||
|
|
||||||
|
def create_ref_model(
|
||||||
|
model_fn: Callable[[], nn.Module],
|
||||||
|
executor: Optional["BaseExecutor"] = None,
|
||||||
|
model: Optional[nn.Module] = None,
|
||||||
|
state_dict: Optional[Dict[str, torch.Tensor]] = None,
|
||||||
|
device: Optional[str] = None,
|
||||||
|
) -> Optional[nn.Module]:
|
||||||
|
"""Create a frozen reference model from executor or state dict.
|
||||||
|
|
||||||
|
In distributed mode (FSDP), ``unwrap_model`` returns ``None`` on
|
||||||
|
non-rank-0. The state_dict is broadcast from rank-0 to all ranks
|
||||||
|
so every rank gets a complete copy.
|
||||||
|
"""
|
||||||
|
if state_dict is None and executor is not None and model is not None:
|
||||||
|
state_dict = executor.unwrap_model(model)
|
||||||
|
|
||||||
|
# FSDP's unwrap_model returns None on non-rank-0. Broadcast from
|
||||||
|
# rank-0 so every rank receives a complete state_dict.
|
||||||
|
if executor is not None and executor.use_distributed:
|
||||||
|
state_dict = broadcast_state_dict(state_dict)
|
||||||
|
|
||||||
|
if state_dict is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
ref_model = model_fn()
|
||||||
|
ref_model.load_state_dict(state_dict)
|
||||||
|
ref_model.requires_grad_(False)
|
||||||
|
ref_model.eval()
|
||||||
|
if device is not None:
|
||||||
|
ref_model = ref_model.to(device=device)
|
||||||
|
return ref_model
|
||||||
|
|
||||||
|
|
||||||
|
class GradientState:
|
||||||
|
def __init__(self, grad_accum_steps: int = 1):
|
||||||
|
self.num_steps = max(grad_accum_steps, 1)
|
||||||
|
self._step: int = 0
|
||||||
|
self._sync_gradients: bool = True
|
||||||
|
|
||||||
|
@property
|
||||||
|
def sync_gradients(self) -> bool:
|
||||||
|
return self._sync_gradients
|
||||||
|
|
||||||
|
def _do_sync(self):
|
||||||
|
self._step += 1
|
||||||
|
self._sync_gradients = self._step % self.num_steps == 0
|
||||||
|
|
||||||
|
|
||||||
|
class AccumOptimizer:
|
||||||
|
def __init__(self, optimizer: Optimizer, gradient_state: GradientState):
|
||||||
|
self.optimizer = optimizer
|
||||||
|
self.gradient_state = gradient_state
|
||||||
|
|
||||||
|
def step(self, closure=None):
|
||||||
|
if self.gradient_state.sync_gradients:
|
||||||
|
self.optimizer.step(closure)
|
||||||
|
|
||||||
|
def zero_grad(self):
|
||||||
|
if self.gradient_state.sync_gradients:
|
||||||
|
self.optimizer.zero_grad()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def param_groups(self):
|
||||||
|
return self.optimizer.param_groups
|
||||||
|
|
||||||
|
def state_dict(self):
|
||||||
|
return self.optimizer.state_dict()
|
||||||
|
|
||||||
|
def load_state_dict(self, d):
|
||||||
|
self.optimizer.load_state_dict(d)
|
||||||
|
|
||||||
|
|
||||||
|
class AccumScheduler:
|
||||||
|
def __init__(self, scheduler: LRScheduler, gradient_state: GradientState):
|
||||||
|
self.scheduler = scheduler
|
||||||
|
self.gradient_state = gradient_state
|
||||||
|
|
||||||
|
def step(self):
|
||||||
|
if self.gradient_state.sync_gradients:
|
||||||
|
self.scheduler.step()
|
||||||
|
|
||||||
|
def state_dict(self):
|
||||||
|
return self.scheduler.state_dict()
|
||||||
|
|
||||||
|
def load_state_dict(self, d):
|
||||||
|
self.scheduler.load_state_dict(d)
|
||||||
|
|
||||||
|
def get_last_lr(self):
|
||||||
|
return self.scheduler.get_last_lr()
|
||||||
|
|
||||||
|
|
||||||
|
class BaseExecutor:
|
||||||
|
def __init__(self, grad_accum_steps: int = 1):
|
||||||
|
self.gradient_state = GradientState(grad_accum_steps)
|
||||||
|
|
||||||
|
def prepare(
|
||||||
|
self,
|
||||||
|
model_fn: Callable[[], nn.Module],
|
||||||
|
optimizer_fn: Optional[Callable[[nn.Module], Optimizer]] = None,
|
||||||
|
scheduler_fn: Optional[Callable[[Optimizer], LRScheduler]] = None,
|
||||||
|
before_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
|
||||||
|
after_wrap: Optional[Callable[[nn.Module], nn.Module]] = None,
|
||||||
|
) -> Tuple[nn.Module, Optional[Optimizer], Optional[LRScheduler]]:
|
||||||
|
model = model_fn()
|
||||||
|
if before_wrap is not None:
|
||||||
|
model = before_wrap(model)
|
||||||
|
model = self._prepare_model(model)
|
||||||
|
if after_wrap is not None:
|
||||||
|
model = after_wrap(model)
|
||||||
|
optimizer = None
|
||||||
|
scheduler = None
|
||||||
|
if optimizer_fn is not None:
|
||||||
|
optimizer = optimizer_fn(model)
|
||||||
|
if scheduler_fn is not None:
|
||||||
|
scheduler = scheduler_fn(optimizer)
|
||||||
|
optimizer = AccumOptimizer(optimizer, self.gradient_state)
|
||||||
|
if scheduler is not None:
|
||||||
|
scheduler = AccumScheduler(scheduler, self.gradient_state)
|
||||||
|
return model, optimizer, scheduler
|
||||||
|
|
||||||
|
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||||
|
return model
|
||||||
|
|
||||||
|
def _no_sync(self, model: nn.Module):
|
||||||
|
return contextlib.nullcontext()
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def accumulate(self, model: nn.Module):
|
||||||
|
self.gradient_state._do_sync()
|
||||||
|
if not self.gradient_state.sync_gradients:
|
||||||
|
with self._no_sync(model):
|
||||||
|
yield
|
||||||
|
else:
|
||||||
|
yield
|
||||||
|
|
||||||
|
def backward(self, loss: torch.Tensor):
|
||||||
|
loss.backward()
|
||||||
|
|
||||||
|
def unwrap_model(self, model: nn.Module):
|
||||||
|
return model.state_dict()
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def checkpoint_context(self, model: nn.Module):
|
||||||
|
if self.use_distributed:
|
||||||
|
dist.barrier()
|
||||||
|
state_dict = self._gather_state_dict(model)
|
||||||
|
yield state_dict
|
||||||
|
if self.use_distributed:
|
||||||
|
dist.barrier()
|
||||||
|
|
||||||
|
def _gather_state_dict(self, model: nn.Module):
|
||||||
|
state_dict = self.unwrap_model(model)
|
||||||
|
if self.use_distributed and get_rank() != 0:
|
||||||
|
return None
|
||||||
|
return state_dict
|
||||||
|
|
||||||
|
@property
|
||||||
|
def use_distributed(self) -> bool:
|
||||||
|
return get_world_size() > 1
|
||||||
|
|
||||||
|
@property
|
||||||
|
def sync_gradients(self) -> bool:
|
||||||
|
return self.gradient_state.sync_gradients
|
||||||
|
|
||||||
|
@property
|
||||||
|
def grad_accum_steps(self) -> int:
|
||||||
|
return self.gradient_state.num_steps
|
||||||
|
|
||||||
|
def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
|
||||||
|
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
|
||||||
|
if isinstance(total_norm, torch.Tensor):
|
||||||
|
return total_norm.item()
|
||||||
|
return total_norm
|
||||||
|
|
||||||
|
|
||||||
|
class ExecutorFactory(BaseFactory[BaseExecutor]):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@ExecutorFactory.register("none")
|
||||||
|
class NoneExecutor(BaseExecutor):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@ExecutorFactory.register("ddp")
|
||||||
|
class DDPExecutor(BaseExecutor):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
grad_accum_steps: int = 1,
|
||||||
|
dim: int = 0,
|
||||||
|
broadcast_buffers: bool = True,
|
||||||
|
init_sync: bool = True,
|
||||||
|
process_group=None,
|
||||||
|
bucket_cap_mb: int = 25,
|
||||||
|
find_unused_parameters: bool = False,
|
||||||
|
check_reduction: bool = False,
|
||||||
|
gradient_as_bucket_view: bool = False,
|
||||||
|
static_graph: bool = False,
|
||||||
|
delay_all_reduce_named_params=None,
|
||||||
|
param_to_hook_all_reduce=None,
|
||||||
|
mixed_precision=None,
|
||||||
|
device_mesh=None,
|
||||||
|
):
|
||||||
|
super().__init__(grad_accum_steps=grad_accum_steps)
|
||||||
|
self._ddp_kwargs = dict(
|
||||||
|
dim=dim,
|
||||||
|
broadcast_buffers=broadcast_buffers,
|
||||||
|
init_sync=init_sync,
|
||||||
|
process_group=process_group,
|
||||||
|
bucket_cap_mb=bucket_cap_mb,
|
||||||
|
find_unused_parameters=find_unused_parameters,
|
||||||
|
check_reduction=check_reduction,
|
||||||
|
gradient_as_bucket_view=gradient_as_bucket_view,
|
||||||
|
static_graph=static_graph,
|
||||||
|
delay_all_reduce_named_params=delay_all_reduce_named_params,
|
||||||
|
param_to_hook_all_reduce=param_to_hook_all_reduce,
|
||||||
|
mixed_precision=mixed_precision,
|
||||||
|
device_mesh=device_mesh,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||||
|
if not self.use_distributed:
|
||||||
|
logger.warning("DDP backend selected but world_size=1, model not wrapped")
|
||||||
|
return model
|
||||||
|
local_rank = int(os.environ.get("LOCAL_RANK", get_rank()))
|
||||||
|
model = DDP(
|
||||||
|
model,
|
||||||
|
device_ids=[local_rank],
|
||||||
|
output_device=local_rank,
|
||||||
|
**self._ddp_kwargs,
|
||||||
|
)
|
||||||
|
logger.info("Model wrapped with DDP (world_size=%d)", get_world_size())
|
||||||
|
return model
|
||||||
|
|
||||||
|
def _no_sync(self, model: nn.Module):
|
||||||
|
if isinstance(model, DDP):
|
||||||
|
return model.no_sync()
|
||||||
|
return contextlib.nullcontext()
|
||||||
|
|
||||||
|
def unwrap_model(self, model: nn.Module):
|
||||||
|
if isinstance(model, DDP):
|
||||||
|
return model.module.state_dict()
|
||||||
|
return model.state_dict()
|
||||||
|
|
||||||
|
|
||||||
|
@ExecutorFactory.register("fsdp")
|
||||||
|
class FSDPExecutor(BaseExecutor):
|
||||||
|
"""FSDP executor using `torch.distributed.fsdp.fully_shard` (per-module API).
|
||||||
|
|
||||||
|
Wraps each child module individually via ``fully_shard``.
|
||||||
|
Skips the root model because ``ABC + Generic[T]`` in the MRO makes
|
||||||
|
``fully_shard``'s dynamic ``__class__`` assignment fail at the CPython level.
|
||||||
|
Original ``Parameter`` objects are preserved (as DTensors) — no
|
||||||
|
``FlatParameter``, no ``use_orig_params=True`` hack.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
grad_accum_steps: int = 1,
|
||||||
|
mesh: Optional[Any] = None,
|
||||||
|
mp_policy: Optional[Any] = None,
|
||||||
|
reshard_after_forward: bool = False,
|
||||||
|
):
|
||||||
|
super().__init__(grad_accum_steps=grad_accum_steps)
|
||||||
|
self._mesh = mesh
|
||||||
|
self._mp_policy = mp_policy
|
||||||
|
self._reshard_after_forward = reshard_after_forward
|
||||||
|
|
||||||
|
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||||
|
if not self.use_distributed:
|
||||||
|
logger.warning("FSDP backend selected but world_size=1, model not wrapped")
|
||||||
|
return model
|
||||||
|
|
||||||
|
kwargs = dict(
|
||||||
|
mesh=self._mesh,
|
||||||
|
mp_policy=self._mp_policy,
|
||||||
|
reshard_after_forward=self._reshard_after_forward,
|
||||||
|
)
|
||||||
|
kwargs = {k: v for k, v in kwargs.items() if v is not None}
|
||||||
|
|
||||||
|
for child in model.children():
|
||||||
|
if isinstance(child, nn.ModuleList):
|
||||||
|
for sub in child:
|
||||||
|
fully_shard(sub, **kwargs)
|
||||||
|
else:
|
||||||
|
fully_shard(child, **kwargs)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"FSDP wrapping applied to %d direct children (root skipped for ABC compat)",
|
||||||
|
len(list(model.children())),
|
||||||
|
)
|
||||||
|
return model
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def _no_sync(self, model: nn.Module):
|
||||||
|
fsdp_modules = [m for m in model.modules() if isinstance(m, FSDPModule)]
|
||||||
|
if fsdp_modules:
|
||||||
|
for m in fsdp_modules:
|
||||||
|
m.set_requires_gradient_sync(False, recurse=True)
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
for m in fsdp_modules:
|
||||||
|
m.set_requires_gradient_sync(True, recurse=True)
|
||||||
|
else:
|
||||||
|
yield
|
||||||
|
|
||||||
|
def clip_grad_norm(self, model: nn.Module, max_norm: float) -> float:
|
||||||
|
if not self.use_distributed:
|
||||||
|
return super().clip_grad_norm(model, max_norm)
|
||||||
|
|
||||||
|
# FSDP params are DTensors (sharded across ranks).
|
||||||
|
# torch.nn.utils.clip_grad_norm_ computes LOCAL norm per rank,
|
||||||
|
# so we must all-reduce to get the global norm before clipping.
|
||||||
|
local_norm = torch.nn.utils.get_total_norm(
|
||||||
|
[p.grad for p in model.parameters() if p.grad is not None],
|
||||||
|
)
|
||||||
|
if isinstance(local_norm, DTensor):
|
||||||
|
local_norm = local_norm.to_local()
|
||||||
|
total_norm_sq = local_norm**2
|
||||||
|
dist.all_reduce(total_norm_sq, op=dist.ReduceOp.SUM)
|
||||||
|
total_norm = total_norm_sq.sqrt()
|
||||||
|
|
||||||
|
clip_coef = max_norm / (total_norm + 1e-6)
|
||||||
|
clip_coef_clamped = torch.clamp(clip_coef, max=1.0)
|
||||||
|
for p in model.parameters():
|
||||||
|
if p.grad is not None:
|
||||||
|
p.grad.mul_(clip_coef_clamped)
|
||||||
|
|
||||||
|
return total_norm.item()
|
||||||
|
|
||||||
|
def unwrap_model(self, model: nn.Module):
|
||||||
|
if not self.use_distributed:
|
||||||
|
return model.state_dict()
|
||||||
|
|
||||||
|
# unshard() and full_tensor() are collective ops — all ranks must
|
||||||
|
# participate. Non-rank-0 ranks still call them but discard results.
|
||||||
|
for module in model.modules():
|
||||||
|
if isinstance(module, FSDPModule):
|
||||||
|
module.unshard()
|
||||||
|
|
||||||
|
state_dict = model.state_dict()
|
||||||
|
result = {}
|
||||||
|
for k, v in state_dict.items():
|
||||||
|
if isinstance(v, DTensor):
|
||||||
|
full = v.full_tensor()
|
||||||
|
if get_rank() == 0:
|
||||||
|
result[k] = full
|
||||||
|
elif get_rank() == 0:
|
||||||
|
result[k] = v
|
||||||
|
|
||||||
|
for module in model.modules():
|
||||||
|
if isinstance(module, FSDPModule):
|
||||||
|
module.reshard()
|
||||||
|
|
||||||
|
if get_rank() != 0:
|
||||||
|
return None
|
||||||
|
|
||||||
|
return result
|
||||||
@@ -1,115 +0,0 @@
|
|||||||
from typing import Dict
|
|
||||||
|
|
||||||
import torch
|
|
||||||
import torch.distributed as dist
|
|
||||||
import torch.nn as nn
|
|
||||||
import torch.nn.functional as F
|
|
||||||
from torch import Tensor
|
|
||||||
|
|
||||||
|
|
||||||
class ParallelModel(nn.Module):
|
|
||||||
def __init__(self, process_group: dist.ProcessGroup):
|
|
||||||
super().__init__()
|
|
||||||
self.process_group = process_group
|
|
||||||
self.rank = dist.get_rank(self.process_group)
|
|
||||||
self.world_size = dist.get_world_size(self.process_group)
|
|
||||||
|
|
||||||
|
|
||||||
class RowParallelLinear(ParallelModel):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
process_group: dist.ProcessGroup,
|
|
||||||
in_features: int,
|
|
||||||
out_features: int,
|
|
||||||
bias: bool = True,
|
|
||||||
reduce_results: bool = True,
|
|
||||||
):
|
|
||||||
super().__init__(process_group)
|
|
||||||
|
|
||||||
self.in_features = in_features
|
|
||||||
self.out_features = out_features
|
|
||||||
self.in_features_per_rank = in_features // self.world_size
|
|
||||||
self.reduce_results = reduce_results
|
|
||||||
|
|
||||||
if in_features % self.world_size != 0:
|
|
||||||
raise ValueError(
|
|
||||||
f"in_features must be divisible by world_size. Got {in_features} and {self.world_size}"
|
|
||||||
)
|
|
||||||
|
|
||||||
self.weight = nn.Parameter(torch.empty(out_features, self.in_features_per_rank))
|
|
||||||
self.bias = nn.Parameter(torch.zeros(out_features)) if bias else None
|
|
||||||
|
|
||||||
def forward(self, input: Tensor) -> Tensor:
|
|
||||||
output = F.linear(input, self.weight)
|
|
||||||
|
|
||||||
if self.reduce_results:
|
|
||||||
dist.all_reduce(output, op=dist.ReduceOp.SUM, group=self.process_group)
|
|
||||||
|
|
||||||
if self.bias is not None:
|
|
||||||
output += self.bias
|
|
||||||
|
|
||||||
return output
|
|
||||||
|
|
||||||
def load_state_dict(self, state_dict: Dict[str, Tensor]):
|
|
||||||
full_weight = state_dict.get("weight")
|
|
||||||
full_bias = state_dict.get("bias")
|
|
||||||
|
|
||||||
start_idx = self.rank * self.in_features_per_rank
|
|
||||||
end_idx = start_idx + self.in_features_per_rank
|
|
||||||
weight_slice = full_weight[:, start_idx:end_idx]
|
|
||||||
self.weight.data.copy_(weight_slice)
|
|
||||||
|
|
||||||
if self.bias is not None:
|
|
||||||
self.bias.data.copy_(full_bias)
|
|
||||||
|
|
||||||
|
|
||||||
class ColumnParallelLinear(ParallelModel):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
process_group: dist.ProcessGroup,
|
|
||||||
in_features: int,
|
|
||||||
out_features: int,
|
|
||||||
bias: bool = True,
|
|
||||||
gather_results: bool = True,
|
|
||||||
):
|
|
||||||
super().__init__(process_group)
|
|
||||||
|
|
||||||
self.in_features = in_features
|
|
||||||
self.out_features = out_features
|
|
||||||
self.out_features_per_rank = out_features // self.world_size
|
|
||||||
self.gather_results = gather_results
|
|
||||||
|
|
||||||
if out_features % self.world_size != 0:
|
|
||||||
raise ValueError(
|
|
||||||
f"out_features must be divisible by world_size. Got {out_features} and {self.world_size}"
|
|
||||||
)
|
|
||||||
|
|
||||||
self.weight = nn.Parameter(
|
|
||||||
torch.empty(self.out_features_per_rank, self.in_features)
|
|
||||||
)
|
|
||||||
self.bias = (
|
|
||||||
nn.Parameter(torch.zeros(self.out_features_per_rank)) if bias else None
|
|
||||||
)
|
|
||||||
|
|
||||||
def forward(self, input: Tensor) -> Tensor:
|
|
||||||
output = F.linear(input, self.weight, self.bias)
|
|
||||||
|
|
||||||
if self.gather_results:
|
|
||||||
output_list = [torch.empty_like(output) for _ in range(self.world_size)]
|
|
||||||
dist.all_gather(output_list, output, group=self.process_group)
|
|
||||||
output = torch.cat(output_list, dim=-1)
|
|
||||||
|
|
||||||
return output
|
|
||||||
|
|
||||||
def load_state_dict(self, state_dict: Dict[str, Tensor]):
|
|
||||||
full_weight = state_dict.get("weight")
|
|
||||||
full_bias = state_dict.get("bias")
|
|
||||||
|
|
||||||
start_idx = self.rank * self.out_features_per_rank
|
|
||||||
end_idx = start_idx + self.out_features_per_rank
|
|
||||||
weight_slice = full_weight[start_idx:end_idx, :]
|
|
||||||
self.weight.data.copy_(weight_slice)
|
|
||||||
|
|
||||||
if self.bias is not None:
|
|
||||||
bias_slice = full_bias[start_idx:end_idx]
|
|
||||||
self.bias.data.copy_(bias_slice)
|
|
||||||
+174
-50
@@ -1,12 +1,27 @@
|
|||||||
|
import logging
|
||||||
import os
|
import os
|
||||||
|
import signal
|
||||||
|
import socket
|
||||||
|
import threading
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from functools import wraps
|
from functools import wraps
|
||||||
from typing import Callable
|
from typing import Callable, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
import torch.multiprocessing as mp
|
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():
|
def get_current_device():
|
||||||
return os.environ["LOCAL_DEVICE"]
|
return os.environ["LOCAL_DEVICE"]
|
||||||
@@ -30,6 +45,7 @@ def get_rank() -> int:
|
|||||||
def setup_parallel(
|
def setup_parallel(
|
||||||
rank: int,
|
rank: int,
|
||||||
world_size: int,
|
world_size: int,
|
||||||
|
local_rank: int,
|
||||||
backend: str = "nccl",
|
backend: str = "nccl",
|
||||||
master_addr: str = "localhost",
|
master_addr: str = "localhost",
|
||||||
master_port: str = "29500",
|
master_port: str = "29500",
|
||||||
@@ -41,20 +57,26 @@ def setup_parallel(
|
|||||||
return
|
return
|
||||||
|
|
||||||
if world_size <= 1:
|
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
|
yield None
|
||||||
return
|
return
|
||||||
|
|
||||||
device_id = torch.device(device_type, rank)
|
device_id = torch.device(device_type, local_rank)
|
||||||
|
|
||||||
os.environ["MASTER_ADDR"] = master_addr
|
os.environ["MASTER_ADDR"] = master_addr
|
||||||
os.environ["MASTER_PORT"] = master_port
|
os.environ["MASTER_PORT"] = master_port
|
||||||
os.environ["LOCAL_RANK"] = str(rank)
|
os.environ["LOCAL_RANK"] = str(local_rank)
|
||||||
os.environ["WORLD_SIZE"] = str(world_size)
|
os.environ["WORLD_SIZE"] = str(world_size)
|
||||||
os.environ["LOCAL_DEVICE"] = str(device_id)
|
os.environ["LOCAL_DEVICE"] = str(device_id)
|
||||||
|
|
||||||
dist.init_process_group(
|
pg_kwargs = dict(rank=rank, world_size=world_size, backend=backend)
|
||||||
rank=rank, world_size=world_size, backend=backend, device_id=device_id
|
if backend in ("nccl", "ccl"):
|
||||||
)
|
pg_kwargs["device_id"] = device_id
|
||||||
|
|
||||||
|
dist.init_process_group(**pg_kwargs)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if backend == "nccl" and torch.cuda.is_available():
|
if backend == "nccl" and torch.cuda.is_available():
|
||||||
@@ -90,7 +112,7 @@ def only_on_rank(rank, sync=False):
|
|||||||
return decorator
|
return decorator
|
||||||
|
|
||||||
|
|
||||||
def wrapper_spawn_func(
|
def _run_single_rank(
|
||||||
rank: int,
|
rank: int,
|
||||||
world_size: int,
|
world_size: int,
|
||||||
backend: str,
|
backend: str,
|
||||||
@@ -100,20 +122,143 @@ def wrapper_spawn_func(
|
|||||||
func: Callable,
|
func: Callable,
|
||||||
kwargs: dict,
|
kwargs: dict,
|
||||||
):
|
):
|
||||||
try:
|
install_early_signal_handlers()
|
||||||
|
with setup_parallel(
|
||||||
|
rank=rank,
|
||||||
|
world_size=world_size,
|
||||||
|
local_rank=rank,
|
||||||
|
backend=backend,
|
||||||
|
master_addr=master_addr,
|
||||||
|
master_port=master_port,
|
||||||
|
device_type=device_type,
|
||||||
|
):
|
||||||
|
func(**kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
class LaunchStrategy(ABC):
|
||||||
|
"""Strategy for launching a function in a distributed context."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
world_size: int,
|
||||||
|
backend: str,
|
||||||
|
master_addr: str,
|
||||||
|
master_port: str,
|
||||||
|
device_type: str,
|
||||||
|
start_method: str,
|
||||||
|
):
|
||||||
|
self.world_size = world_size
|
||||||
|
self.backend = backend
|
||||||
|
self.master_addr = master_addr
|
||||||
|
self.master_port = master_port
|
||||||
|
self.device_type = device_type
|
||||||
|
self.start_method = start_method
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def launch(self, func: Callable, **kwargs):
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
|
||||||
|
class TorchrunStrategy(LaunchStrategy):
|
||||||
|
"""External orchestrator (torchrun, SLURM, K8s) — env vars pre-set."""
|
||||||
|
|
||||||
|
def launch(self, func: Callable, **kwargs):
|
||||||
|
install_early_signal_handlers()
|
||||||
|
rank = int(os.environ["RANK"])
|
||||||
|
world_size = int(os.environ["WORLD_SIZE"])
|
||||||
|
local_rank = int(os.environ.get("LOCAL_RANK", rank))
|
||||||
with setup_parallel(
|
with setup_parallel(
|
||||||
rank=rank,
|
rank=rank,
|
||||||
world_size=world_size,
|
world_size=world_size,
|
||||||
backend=backend,
|
local_rank=local_rank,
|
||||||
master_addr=master_addr,
|
backend=self.backend,
|
||||||
master_port=master_port,
|
master_addr=os.environ.get("MASTER_ADDR", self.master_addr),
|
||||||
device_type=device_type,
|
master_port=os.environ.get("MASTER_PORT", self.master_port),
|
||||||
|
device_type=self.device_type,
|
||||||
):
|
):
|
||||||
func(**kwargs)
|
func(**kwargs)
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
print(f"Error in rank {rank}: {e}")
|
class LocalStrategy(LaunchStrategy):
|
||||||
raise
|
"""Local launcher — single-process or mp.start_processes."""
|
||||||
|
|
||||||
|
def launch(self, func: Callable, **kwargs):
|
||||||
|
args = (
|
||||||
|
self.world_size,
|
||||||
|
self.backend,
|
||||||
|
self.master_addr,
|
||||||
|
self.master_port,
|
||||||
|
self.device_type,
|
||||||
|
func,
|
||||||
|
kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.world_size == 1:
|
||||||
|
_run_single_rank(0, *args)
|
||||||
|
return
|
||||||
|
|
||||||
|
install_early_signal_handlers()
|
||||||
|
ctx = mp.start_processes(
|
||||||
|
_run_single_rank,
|
||||||
|
args=args,
|
||||||
|
nprocs=self.world_size,
|
||||||
|
start_method=self.start_method,
|
||||||
|
join=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
parent_stop = threading.Event()
|
||||||
|
original_handlers = {}
|
||||||
|
|
||||||
|
def _parent_handler(signum, frame):
|
||||||
|
sig = signal.Signals(signum)
|
||||||
|
logger.warning(
|
||||||
|
"Parent (pid=%d) received %s, forwarding to children...",
|
||||||
|
os.getpid(),
|
||||||
|
sig.name,
|
||||||
|
)
|
||||||
|
parent_stop.set()
|
||||||
|
for p in ctx.processes:
|
||||||
|
if p.is_alive():
|
||||||
|
p.terminate()
|
||||||
|
|
||||||
|
for sig in (signal.SIGTERM, signal.SIGINT):
|
||||||
|
prev = signal.signal(sig, _parent_handler)
|
||||||
|
if prev not in (signal.SIG_DFL, signal.SIG_IGN, None, _parent_handler):
|
||||||
|
original_handlers[sig] = prev
|
||||||
|
|
||||||
|
try:
|
||||||
|
while not ctx.join() and not parent_stop.is_set():
|
||||||
|
pass
|
||||||
|
except BaseException:
|
||||||
|
logger.warning(
|
||||||
|
"Parent received unexpected exception, terminating children..."
|
||||||
|
)
|
||||||
|
for p in ctx.processes:
|
||||||
|
if p.is_alive():
|
||||||
|
p.terminate()
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
for sig, handler in original_handlers.items():
|
||||||
|
signal.signal(sig, handler)
|
||||||
|
|
||||||
|
for p in ctx.processes:
|
||||||
|
p.join()
|
||||||
|
|
||||||
|
ctx.join()
|
||||||
|
|
||||||
|
|
||||||
|
def _detect_launcher() -> str:
|
||||||
|
"""Detect the distributed launcher from environment.
|
||||||
|
|
||||||
|
Returns one of: "torchelastic", "torchrun", "external", "local".
|
||||||
|
"""
|
||||||
|
if dist.is_torchelastic_launched():
|
||||||
|
return "torchelastic"
|
||||||
|
if "LOCAL_WORLD_SIZE" in os.environ:
|
||||||
|
return "torchrun"
|
||||||
|
if "RANK" in os.environ and "WORLD_SIZE" in os.environ:
|
||||||
|
return "external"
|
||||||
|
return "local"
|
||||||
|
|
||||||
|
|
||||||
def spawn_parallel_fn(
|
def spawn_parallel_fn(
|
||||||
@@ -121,41 +266,20 @@ def spawn_parallel_fn(
|
|||||||
world_size: int,
|
world_size: int,
|
||||||
backend: str = "nccl",
|
backend: str = "nccl",
|
||||||
master_addr: str = "localhost",
|
master_addr: str = "localhost",
|
||||||
master_port: str = "29500",
|
master_port: Optional[str] = None,
|
||||||
device_type: str = "cuda",
|
device_type: str = "cuda",
|
||||||
|
start_method: str = "spawn",
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
# clear environment variables
|
if master_port is None:
|
||||||
for key in [
|
master_port = find_free_port()
|
||||||
"MASTER_ADDR",
|
launcher = _detect_launcher()
|
||||||
"MASTER_PORT",
|
if launcher in ("torchelastic", "torchrun", "external"):
|
||||||
"RANK",
|
strategy = TorchrunStrategy(
|
||||||
"WORLD_SIZE",
|
world_size, backend, master_addr, master_port, device_type, start_method
|
||||||
"LOCAL_RANK",
|
)
|
||||||
"LOCAL_DEVICE",
|
else:
|
||||||
]:
|
strategy = LocalStrategy(
|
||||||
if key in os.environ:
|
world_size, backend, master_addr, master_port, device_type, start_method
|
||||||
del os.environ[key]
|
)
|
||||||
|
strategy.launch(func, **kwargs)
|
||||||
if world_size == 1:
|
|
||||||
device_id = torch.device(device_type, 0)
|
|
||||||
os.environ["LOCAL_RANK"] = "0"
|
|
||||||
os.environ["WORLD_SIZE"] = "1"
|
|
||||||
os.environ["LOCAL_DEVICE"] = str(device_id)
|
|
||||||
|
|
||||||
func(**kwargs)
|
|
||||||
return
|
|
||||||
|
|
||||||
wrapper_spawn_func_args = (
|
|
||||||
world_size,
|
|
||||||
backend,
|
|
||||||
master_addr,
|
|
||||||
master_port,
|
|
||||||
device_type,
|
|
||||||
func,
|
|
||||||
kwargs,
|
|
||||||
)
|
|
||||||
|
|
||||||
mp.spawn(
|
|
||||||
wrapper_spawn_func, nprocs=world_size, args=wrapper_spawn_func_args, join=True
|
|
||||||
)
|
|
||||||
|
|||||||
@@ -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,537 @@
|
|||||||
|
"""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 = {}
|
||||||
|
any_output = False
|
||||||
|
|
||||||
|
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
|
||||||
|
any_output = True
|
||||||
|
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
|
||||||
|
any_output = True
|
||||||
|
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
|
||||||
|
|
||||||
|
any_output = True
|
||||||
|
|
||||||
|
if not any_output:
|
||||||
|
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]
|
||||||
|
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 result
|
||||||
|
else None
|
||||||
|
for item, result in zip(items, results)
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@MaskBuilderFactory.register("sectioned")
|
||||||
|
class SectionedMaskBuilder(BaseMaskBuilder):
|
||||||
|
"""Façade that dispatches to SingleOutputMaskBuilder or MultiOutputMaskBuilder.
|
||||||
|
|
||||||
|
Preserves backward compatibility for existing configs and code that rely
|
||||||
|
on the ``"sectioned"`` factory name.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._single = SingleOutputMaskBuilder()
|
||||||
|
self._multi = MultiOutputMaskBuilder()
|
||||||
|
|
||||||
|
def build(self, item: dict, config, tokenizer) -> Optional[dict]:
|
||||||
|
sources_spec = getattr(config.input, "sources", None)
|
||||||
|
if sources_spec:
|
||||||
|
return self._multi.build(item, config, tokenizer)
|
||||||
|
return self._single.build(item, config, tokenizer)
|
||||||
|
|
||||||
|
def build_batch(self, items, config, tokenizer):
|
||||||
|
sources_spec = getattr(config.input, "sources", None)
|
||||||
|
if sources_spec:
|
||||||
|
return self._multi.build_batch(items, config, tokenizer)
|
||||||
|
return self._single.build_batch(items, config, tokenizer)
|
||||||
@@ -0,0 +1,124 @@
|
|||||||
|
"""Shared preprocessing kernel used by both :class:`Pipeline` and
|
||||||
|
:class:`TokenizeTransform`.
|
||||||
|
|
||||||
|
The two entry points previously duplicated ~60 % of their logic:
|
||||||
|
record iteration, mask-builder invocation, primary-id extraction,
|
||||||
|
per-key accumulation, dtype inference and position-id generation.
|
||||||
|
This module factors out the common core as pure functions so that
|
||||||
|
the online (``TokenizeTransform``) and offline (``Pipeline``) paths
|
||||||
|
stay in lockstep.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from itertools import chain
|
||||||
|
from typing import Dict, Iterator, List, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
|
from astrai.preprocessing.builder import MaskBuilderFactory
|
||||||
|
from astrai.preprocessing.position_id import PositionIdStrategyFactory
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
|
def build_preprocessing_components(config: PipelineConfig, tokenizer_path: str):
|
||||||
|
"""Load tokenizer, mask builder and position-id strategy together.
|
||||||
|
|
||||||
|
Both ``Pipeline`` and ``TokenizeTransform`` need the same triple;
|
||||||
|
centralising the construction avoids drift (e.g. one path forgetting
|
||||||
|
to create the position-id strategy).
|
||||||
|
"""
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
||||||
|
mask_builder = MaskBuilderFactory.create("sectioned")
|
||||||
|
position_strategy = PositionIdStrategyFactory.create(
|
||||||
|
config.output.position_ids_mode
|
||||||
|
)
|
||||||
|
return tokenizer, mask_builder, position_strategy
|
||||||
|
|
||||||
|
|
||||||
|
def primary_ids(result: dict) -> List[int]:
|
||||||
|
"""Return the first flat int-list value in *result*.
|
||||||
|
|
||||||
|
Used for token counting and position-id generation when the
|
||||||
|
primary key name is not known (DPO uses ``chosen``, GRPO uses
|
||||||
|
``prompts``, SFT uses ``sequence``).
|
||||||
|
"""
|
||||||
|
for val in result.values():
|
||||||
|
if isinstance(val, list) and val and isinstance(val[0], int):
|
||||||
|
return val
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
def infer_dtype(ids: List) -> torch.dtype:
|
||||||
|
"""Float values become float32, everything else int32."""
|
||||||
|
if ids and isinstance(ids[0], float):
|
||||||
|
return torch.float32
|
||||||
|
return torch.int32
|
||||||
|
|
||||||
|
|
||||||
|
def iter_raw_records(
|
||||||
|
records: List[dict],
|
||||||
|
mask_builder,
|
||||||
|
config: PipelineConfig,
|
||||||
|
tokenizer,
|
||||||
|
) -> Iterator[dict]:
|
||||||
|
"""Yield mask-builder output dicts for each record, skipping failures.
|
||||||
|
|
||||||
|
Drops ``domain`` from the result (callers that need it should read
|
||||||
|
it before calling this). Each yielded dict maps a key
|
||||||
|
(``sequence``, ``chosen``, ``responses``…) to either a flat
|
||||||
|
``List[int]`` or a nested ``List[List[int]]`` (GRPO responses/masks).
|
||||||
|
"""
|
||||||
|
for item in records:
|
||||||
|
result = mask_builder.build(item, config, tokenizer)
|
||||||
|
if result is None:
|
||||||
|
continue
|
||||||
|
result.pop("domain", None)
|
||||||
|
if not primary_ids(result):
|
||||||
|
continue
|
||||||
|
yield result
|
||||||
|
|
||||||
|
|
||||||
|
def to_per_record_tensors(
|
||||||
|
raw: Dict[str, list],
|
||||||
|
) -> Dict[str, List[torch.Tensor]]:
|
||||||
|
"""Convert an accumulated ``{key: [per-record ids]}`` dict to tensors.
|
||||||
|
|
||||||
|
Handles three shapes transparently:
|
||||||
|
|
||||||
|
- ``List[int]`` per record (``sequence``, ``chosen``…) → one tensor per record.
|
||||||
|
- ``List[List[int]]`` per record (GRPO ``responses``/``masks``) → one
|
||||||
|
``List[Tensor]`` per record (nested), preserving the per-response
|
||||||
|
boundary so downstream code can index responses individually.
|
||||||
|
- ``List[int]`` for the whole shard (pre-packed keys) → single tensor.
|
||||||
|
|
||||||
|
The detection mirrors the previous inline logic in
|
||||||
|
``Pipeline._flush`` and ``TokenizeTransform.apply``.
|
||||||
|
"""
|
||||||
|
tensors: Dict[str, List[torch.Tensor]] = {}
|
||||||
|
for key, ids_list in raw.items():
|
||||||
|
if ids_list and isinstance(ids_list[0], list):
|
||||||
|
tensors[key] = [
|
||||||
|
[torch.tensor(sub, dtype=infer_dtype(sub)) for sub in ids]
|
||||||
|
if ids and isinstance(ids[0], list)
|
||||||
|
else torch.tensor(ids, dtype=infer_dtype(ids))
|
||||||
|
for ids in ids_list
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
tensors[key] = [
|
||||||
|
torch.tensor(list(chain.from_iterable(ids_list)), dtype=torch.int32)
|
||||||
|
]
|
||||||
|
return tensors
|
||||||
|
|
||||||
|
|
||||||
|
def build_position_ids(
|
||||||
|
sequences: List[List[int]],
|
||||||
|
strategy,
|
||||||
|
) -> Optional[List[int]]:
|
||||||
|
"""Generate position ids for *sequences* using *strategy*.
|
||||||
|
|
||||||
|
Returns ``None`` when the strategy produces no ids (e.g. ``none``
|
||||||
|
mode), so callers can skip attaching the key instead of storing
|
||||||
|
an empty list.
|
||||||
|
"""
|
||||||
|
pos_ids = strategy.generate(sequences)
|
||||||
|
return pos_ids or None
|
||||||
@@ -0,0 +1,176 @@
|
|||||||
|
"""Sequence packing strategies for shard-level reordering and truncation.
|
||||||
|
|
||||||
|
Each strategy receives the accumulated ``{key: [list of token lists]}``
|
||||||
|
dict for a shard and returns a reordered / truncated version. The
|
||||||
|
pipeline later flattens the result into contiguous tensors.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from typing import Dict, List
|
||||||
|
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
|
||||||
|
|
||||||
|
def _truncate(seq: List[int], max_len: int, mode: str) -> List[int]:
|
||||||
|
if len(seq) <= max_len:
|
||||||
|
return seq
|
||||||
|
if mode == "keep_end":
|
||||||
|
return seq[-max_len:]
|
||||||
|
return seq[:max_len]
|
||||||
|
|
||||||
|
|
||||||
|
def plan_bfd(
|
||||||
|
sequences: List[List[int]], max_packed_len: int, truncation_mode: str = "keep_start"
|
||||||
|
) -> List[List[int]]:
|
||||||
|
"""Best-Fit Decreasing bin packing of *sequences* into bins.
|
||||||
|
|
||||||
|
Returns a list of bins, each bin a list of original indices into
|
||||||
|
*sequences*. Bin capacities are respected on the *truncated*
|
||||||
|
length of each sequence (so a sequence longer than
|
||||||
|
*max_packed_len* counts at *max_packed_len*).
|
||||||
|
|
||||||
|
Pure index-based so callers can apply the same plan to any
|
||||||
|
aligned key (``loss_mask``, ``position_ids``…).
|
||||||
|
"""
|
||||||
|
n = len(sequences)
|
||||||
|
order = sorted(range(n), key=lambda i: len(sequences[i]), reverse=True)
|
||||||
|
bins: List[List[int]] = []
|
||||||
|
bin_lengths: List[int] = []
|
||||||
|
|
||||||
|
for orig_idx in order:
|
||||||
|
seq_len = len(_truncate(sequences[orig_idx], max_packed_len, truncation_mode))
|
||||||
|
best_bin = None
|
||||||
|
best_remain = max_packed_len + 1
|
||||||
|
for i, bl in enumerate(bin_lengths):
|
||||||
|
remain = max_packed_len - bl
|
||||||
|
if seq_len <= remain < best_remain:
|
||||||
|
best_remain = remain
|
||||||
|
best_bin = i
|
||||||
|
if best_bin is not None:
|
||||||
|
bins[best_bin].append(orig_idx)
|
||||||
|
bin_lengths[best_bin] += seq_len
|
||||||
|
else:
|
||||||
|
bins.append([orig_idx])
|
||||||
|
bin_lengths.append(seq_len)
|
||||||
|
|
||||||
|
return bins
|
||||||
|
|
||||||
|
|
||||||
|
class PackingStrategy(ABC):
|
||||||
|
"""Reorder and truncate sequences within a shard."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def apply(
|
||||||
|
self,
|
||||||
|
keys: Dict[str, List[List[int]]],
|
||||||
|
max_packed_len: int,
|
||||||
|
truncation_mode: str,
|
||||||
|
) -> Dict[str, List[List[int]]]:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
|
||||||
|
class PackingStrategyFactory(BaseFactory["PackingStrategy"]):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@PackingStrategyFactory.register("simple")
|
||||||
|
class SimplePacking(PackingStrategy):
|
||||||
|
def apply(
|
||||||
|
self,
|
||||||
|
keys: Dict[str, List[List[int]]],
|
||||||
|
max_packed_len: int,
|
||||||
|
truncation_mode: str,
|
||||||
|
) -> Dict[str, List[List[int]]]:
|
||||||
|
return {
|
||||||
|
k: [_truncate(v, max_packed_len, truncation_mode) for v in vals]
|
||||||
|
for k, vals in keys.items()
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@PackingStrategyFactory.register("bfd")
|
||||||
|
class BFDPacking(PackingStrategy):
|
||||||
|
"""Best-Fit Decreasing bin packing.
|
||||||
|
|
||||||
|
Assigns sequences to bins using a best-fit heuristic (sorted by
|
||||||
|
decreasing length) and concatenates sequences within each bin into
|
||||||
|
a single packed sequence. Packed sequences are truncated to
|
||||||
|
*max_packed_len* so that each packed bin fits within one context
|
||||||
|
window during training.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def apply(
|
||||||
|
self,
|
||||||
|
keys: Dict[str, List[List[int]]],
|
||||||
|
max_packed_len: int,
|
||||||
|
truncation_mode: str,
|
||||||
|
) -> Dict[str, List[List[int]]]:
|
||||||
|
sequences = keys.get("sequence", [])
|
||||||
|
if not sequences:
|
||||||
|
return keys
|
||||||
|
bins = plan_bfd(sequences, max_packed_len, truncation_mode)
|
||||||
|
|
||||||
|
packed: Dict[str, List[List[int]]] = {}
|
||||||
|
for k, vals in keys.items():
|
||||||
|
packed[k] = [
|
||||||
|
_truncate(
|
||||||
|
self._concat_bin(vals, bin_indices),
|
||||||
|
max_packed_len,
|
||||||
|
truncation_mode,
|
||||||
|
)
|
||||||
|
for bin_indices in bins
|
||||||
|
]
|
||||||
|
return packed
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _concat_bin(vals: List[List[int]], indices: List[int]) -> List[int]:
|
||||||
|
result: List[int] = []
|
||||||
|
for i in indices:
|
||||||
|
result.extend(vals[i])
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
@PackingStrategyFactory.register("bfd_split")
|
||||||
|
class BFDSplitPacking(BFDPacking):
|
||||||
|
"""BFD packing with over-length sequences split into chunks.
|
||||||
|
|
||||||
|
Sequences longer than *max_packed_len* are split into consecutive
|
||||||
|
chunks of at most *max_packed_len* tokens instead of being
|
||||||
|
truncated. Each chunk becomes an independent sequence that enters
|
||||||
|
BFD planning. All keys (``loss_mask``, ``position_ids``, …) are
|
||||||
|
split in lockstep so per-token alignment is preserved.
|
||||||
|
|
||||||
|
Note: because each chunk is treated as a separate document, the
|
||||||
|
second chunk of a split sequence loses the preceding context.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def apply(
|
||||||
|
self,
|
||||||
|
keys: Dict[str, List[List[int]]],
|
||||||
|
max_packed_len: int,
|
||||||
|
truncation_mode: str,
|
||||||
|
) -> Dict[str, List[List[int]]]:
|
||||||
|
sequences = keys.get("sequence", [])
|
||||||
|
if not sequences:
|
||||||
|
return keys
|
||||||
|
if max_packed_len <= 0:
|
||||||
|
return super().apply(keys, max_packed_len, truncation_mode)
|
||||||
|
|
||||||
|
split_keys = self._split_all(keys, max_packed_len)
|
||||||
|
return super().apply(split_keys, max_packed_len, truncation_mode)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _split_all(
|
||||||
|
keys: Dict[str, List[List[int]]], max_packed_len: int
|
||||||
|
) -> Dict[str, List[List[int]]]:
|
||||||
|
"""Split every sequence exceeding *max_packed_len* into chunks,
|
||||||
|
applying the same chunk boundaries to all keys."""
|
||||||
|
sequences = keys["sequence"]
|
||||||
|
chunk_bounds = [list(range(0, len(s), max_packed_len)) for s in sequences]
|
||||||
|
result: Dict[str, List[List[int]]] = {}
|
||||||
|
for key, vals in keys.items():
|
||||||
|
split_vals: List[List[int]] = []
|
||||||
|
for val, starts in zip(vals, chunk_bounds):
|
||||||
|
for start in starts:
|
||||||
|
split_vals.append(val[start : start + max_packed_len])
|
||||||
|
result[key] = split_vals
|
||||||
|
return result
|
||||||
@@ -0,0 +1,277 @@
|
|||||||
|
"""Config-driven JSONL preprocessing pipeline.
|
||||||
|
|
||||||
|
Composes a :class:`BaseMaskBuilder` (selected by ``input.type``) with
|
||||||
|
sharding and flush to ``.bin`` storage. Packing, position-id
|
||||||
|
generation and storage writing are each delegated to pluggable strategies,
|
||||||
|
dispatched by configuration keys.
|
||||||
|
|
||||||
|
Record iteration, mask building, primary-id extraction and per-key
|
||||||
|
accumulation are shared with :class:`TokenizeTransform` via the
|
||||||
|
:mod:`astrai.preprocessing.core` helpers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
from collections import defaultdict
|
||||||
|
from itertools import chain
|
||||||
|
from typing import Dict, List, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import tqdm
|
||||||
|
|
||||||
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
|
from astrai.preprocessing.core import (
|
||||||
|
build_preprocessing_components,
|
||||||
|
primary_ids,
|
||||||
|
)
|
||||||
|
from astrai.preprocessing.packing import PackingStrategyFactory
|
||||||
|
from astrai.preprocessing.writer import StoreWriterFactory
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_STR_TO_DTYPE: dict[str, torch.dtype] = {
|
||||||
|
"bool": torch.bool,
|
||||||
|
"uint8": torch.uint8,
|
||||||
|
"int8": torch.int8,
|
||||||
|
"int16": torch.int16,
|
||||||
|
"int32": torch.int32,
|
||||||
|
"int64": torch.int64,
|
||||||
|
"float16": torch.float16,
|
||||||
|
"float32": torch.float32,
|
||||||
|
"float64": torch.float64,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def filter_by_length(text: str, min_len: int = 50, max_len: int = 2_000_000) -> bool:
|
||||||
|
return min_len <= len(text) <= max_len
|
||||||
|
|
||||||
|
|
||||||
|
class Pipeline:
|
||||||
|
"""Tokenization pipeline driven by a declarative :class:`PipelineConfig`.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
config = PipelineConfig.from_file("sft_pipeline.json")
|
||||||
|
Pipeline(config, ["data.jsonl"], output_dir="out", tokenizer_path="params").run()
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: PipelineConfig,
|
||||||
|
input_paths: list[str],
|
||||||
|
output_dir: str,
|
||||||
|
tokenizer_path: str,
|
||||||
|
):
|
||||||
|
os.makedirs(output_dir, exist_ok=True)
|
||||||
|
self.config = config
|
||||||
|
self.paths = input_paths
|
||||||
|
self.output_dir = output_dir
|
||||||
|
self.tokenizer_path = tokenizer_path
|
||||||
|
|
||||||
|
self.tokenizer, self.mask_builder, self._position_id = (
|
||||||
|
build_preprocessing_components(config, tokenizer_path)
|
||||||
|
)
|
||||||
|
self._packer = PackingStrategyFactory.create(
|
||||||
|
config.preprocessing.packing_strategy
|
||||||
|
)
|
||||||
|
self._writer = StoreWriterFactory.create(config.output.storage_format)
|
||||||
|
|
||||||
|
def transform(self, item: dict) -> Optional[dict]:
|
||||||
|
return self.mask_builder.build(item, self.config, self.tokenizer)
|
||||||
|
|
||||||
|
def transform_batch(self, items: list[dict]) -> list[Optional[dict]]:
|
||||||
|
return self.mask_builder.build_batch(items, self.config, self.tokenizer)
|
||||||
|
|
||||||
|
def run(self):
|
||||||
|
domains: dict = defaultdict(lambda: defaultdict(list))
|
||||||
|
total_tokens = 0
|
||||||
|
shard_idx: dict[str, int] = defaultdict(int)
|
||||||
|
count = 0
|
||||||
|
|
||||||
|
pp = self.config.preprocessing
|
||||||
|
|
||||||
|
progress = tqdm.tqdm(desc="Tokenizing", unit="docs", mininterval=0.5)
|
||||||
|
stop = False
|
||||||
|
for items in self._iter_batches(pp.batch_size):
|
||||||
|
progress.update(len(items))
|
||||||
|
try:
|
||||||
|
results = self.transform_batch(items)
|
||||||
|
except Exception:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to process batch, retrying records individually",
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
results = []
|
||||||
|
for item in items:
|
||||||
|
try:
|
||||||
|
results.append(self.transform(item))
|
||||||
|
except Exception:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to process item, skipping", exc_info=True
|
||||||
|
)
|
||||||
|
results.append(None)
|
||||||
|
|
||||||
|
for result in results:
|
||||||
|
if pp.max_items and count >= pp.max_items:
|
||||||
|
stop = True
|
||||||
|
break
|
||||||
|
if result is None:
|
||||||
|
continue
|
||||||
|
|
||||||
|
domain = result.pop("domain", "__default__")
|
||||||
|
ids = primary_ids(result)
|
||||||
|
if not ids:
|
||||||
|
continue
|
||||||
|
|
||||||
|
bucket = domains[domain]
|
||||||
|
self._align_bucket(bucket, result, ids)
|
||||||
|
for key, val in result.items():
|
||||||
|
bucket[key].append(val)
|
||||||
|
|
||||||
|
count += 1
|
||||||
|
total_tokens += len(ids)
|
||||||
|
|
||||||
|
if total_tokens >= self.config.output.max_tokens_per_shard:
|
||||||
|
self._flush(domains, shard_idx)
|
||||||
|
domains.clear()
|
||||||
|
total_tokens = 0
|
||||||
|
if stop:
|
||||||
|
break
|
||||||
|
|
||||||
|
progress.close()
|
||||||
|
|
||||||
|
if total_tokens > 0:
|
||||||
|
self._flush(domains, shard_idx)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _align_bucket(bucket: dict, result: dict, ids: list):
|
||||||
|
"""Pad previously-accumulated keys that are missing from *result*."""
|
||||||
|
for key in list(bucket.keys()):
|
||||||
|
if key in result:
|
||||||
|
continue
|
||||||
|
bucket[key].append([0] * len(ids))
|
||||||
|
|
||||||
|
def _iter_items(self):
|
||||||
|
for path in self.paths:
|
||||||
|
with open(path, "r", encoding="utf-8") as f:
|
||||||
|
if path.endswith(".json"):
|
||||||
|
data = json.load(f)
|
||||||
|
if isinstance(data, dict):
|
||||||
|
yield data
|
||||||
|
elif isinstance(data, list):
|
||||||
|
yield from data
|
||||||
|
else:
|
||||||
|
for line in f:
|
||||||
|
line = line.strip()
|
||||||
|
if not line:
|
||||||
|
continue
|
||||||
|
yield json.loads(line)
|
||||||
|
|
||||||
|
def _iter_batches(self, batch_size: int):
|
||||||
|
batch_size = max(1, batch_size)
|
||||||
|
batch = []
|
||||||
|
for item in self._iter_items():
|
||||||
|
batch.append(item)
|
||||||
|
if len(batch) >= batch_size:
|
||||||
|
yield batch
|
||||||
|
batch = []
|
||||||
|
if batch:
|
||||||
|
yield batch
|
||||||
|
|
||||||
|
def _flush(self, domains, shard_idx):
|
||||||
|
for domain, keys in domains.items():
|
||||||
|
idx = shard_idx[domain]
|
||||||
|
|
||||||
|
pp = self.config.preprocessing
|
||||||
|
original_sequences = keys.get("sequence", [])
|
||||||
|
mode = self.config.output.position_ids_mode
|
||||||
|
|
||||||
|
keys = self._inject_doc_reset_position_ids(keys, mode, original_sequences)
|
||||||
|
keys = self._packer.apply(dict(keys), pp.max_packed_len, pp.truncation_mode)
|
||||||
|
tensors = self._to_tensors(keys)
|
||||||
|
tensors = self._inject_continuous_position_ids(
|
||||||
|
tensors, mode, keys.get("sequence", [])
|
||||||
|
)
|
||||||
|
|
||||||
|
self._writer.save(self.output_dir, domain, idx, tensors)
|
||||||
|
shard_idx[domain] = idx + 1
|
||||||
|
|
||||||
|
first_key = "sequence" if "sequence" in tensors else next(iter(tensors))
|
||||||
|
tqdm.tqdm.write(
|
||||||
|
f" saved {domain}/shard_{idx:04d} "
|
||||||
|
f"({tensors[first_key][0].numel():,} tokens)"
|
||||||
|
)
|
||||||
|
|
||||||
|
def _inject_doc_reset_position_ids(
|
||||||
|
self,
|
||||||
|
keys: Dict[str, list],
|
||||||
|
mode: str,
|
||||||
|
original_sequences: List[List[int]],
|
||||||
|
) -> Dict[str, list]:
|
||||||
|
"""Attach per-document position_ids before packing (``doc_reset``).
|
||||||
|
|
||||||
|
``doc_reset`` position ids must enter the packer so that each
|
||||||
|
packed bin concatenates the per-doc ranges in bin order. The
|
||||||
|
per-record structure ``[range(len(s)) for s in seqs]`` is required
|
||||||
|
by the packer (it concatenates per-record lists per bin); the
|
||||||
|
``PositionIdStrategy.generate`` flattens, so it cannot be used
|
||||||
|
directly here — it is only consulted for the ``continuous``
|
||||||
|
post-packing path.
|
||||||
|
"""
|
||||||
|
if mode != "doc_reset" or not original_sequences:
|
||||||
|
return keys
|
||||||
|
keys["position_ids"] = [list(range(len(s))) for s in original_sequences]
|
||||||
|
return keys
|
||||||
|
|
||||||
|
def _inject_continuous_position_ids(
|
||||||
|
self,
|
||||||
|
tensors: Dict[str, List[torch.Tensor]],
|
||||||
|
mode: str,
|
||||||
|
packed_sequences: List[List[int]],
|
||||||
|
) -> Dict[str, List[torch.Tensor]]:
|
||||||
|
"""Attach a single continuous position_ids tensor after packing.
|
||||||
|
|
||||||
|
``continuous`` mode spans the whole shard (post-packing), so it
|
||||||
|
cannot participate in bin packing — it is computed from the
|
||||||
|
packed sequences and appended directly to the tensor dict.
|
||||||
|
"""
|
||||||
|
if mode != "continuous" or not packed_sequences:
|
||||||
|
return tensors
|
||||||
|
pos_ids = self._position_id.generate(packed_sequences)
|
||||||
|
if pos_ids:
|
||||||
|
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
||||||
|
return tensors
|
||||||
|
|
||||||
|
def _to_tensors(self, keys: Dict[str, list]) -> Dict[str, List[torch.Tensor]]:
|
||||||
|
"""Convert packed per-key id lists to tensors.
|
||||||
|
|
||||||
|
Honours ``config.output.dtype`` overrides per key; falls back to
|
||||||
|
``int32``. Handles three shapes (see
|
||||||
|
:func:`astrai.preprocessing.core.to_per_record_tensors` for the
|
||||||
|
equivalent online-path helper):
|
||||||
|
- ``List[int]`` per record → one tensor per record.
|
||||||
|
- ``List[List[int]]`` per record (GRPO responses/masks) → one tensor
|
||||||
|
per record, inner lists flattened.
|
||||||
|
- ``List[int]`` for the whole shard (pre-packed keys) → single tensor.
|
||||||
|
"""
|
||||||
|
tensors: Dict[str, List[torch.Tensor]] = {}
|
||||||
|
for key, ids_list in keys.items():
|
||||||
|
dt = _STR_TO_DTYPE.get(
|
||||||
|
self.config.output.dtype.get(key, "int32"), torch.int32
|
||||||
|
)
|
||||||
|
if ids_list and isinstance(ids_list[0], list):
|
||||||
|
tensors[key] = [
|
||||||
|
torch.tensor(
|
||||||
|
list(chain.from_iterable(ids))
|
||||||
|
if ids and isinstance(ids[0], list)
|
||||||
|
else ids,
|
||||||
|
dtype=dt,
|
||||||
|
)
|
||||||
|
for ids in ids_list
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
tensors[key] = [
|
||||||
|
torch.tensor(list(chain.from_iterable(ids_list)), dtype=dt)
|
||||||
|
]
|
||||||
|
return tensors
|
||||||
@@ -0,0 +1,46 @@
|
|||||||
|
"""Position-id generation strategies for packed sequences.
|
||||||
|
|
||||||
|
Each strategy takes the list of per-document token sequences after packing
|
||||||
|
and returns a flat list of position ids (same total length as all
|
||||||
|
sequences combined). The pipeline wraps the result into a tensor and
|
||||||
|
attaches it as ``position_ids``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from typing import List
|
||||||
|
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
|
||||||
|
|
||||||
|
class PositionIdStrategy(ABC):
|
||||||
|
"""Generate ``position_ids`` for packed sequences."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def generate(self, sequences: List[List[int]]) -> List[int]:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
|
||||||
|
class PositionIdStrategyFactory(BaseFactory["PositionIdStrategy"]):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@PositionIdStrategyFactory.register("none")
|
||||||
|
class NoPositionId(PositionIdStrategy):
|
||||||
|
def generate(self, sequences: List[List[int]]) -> List[int]:
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
@PositionIdStrategyFactory.register("doc_reset")
|
||||||
|
class DocResetPositionId(PositionIdStrategy):
|
||||||
|
def generate(self, sequences: List[List[int]]) -> List[int]:
|
||||||
|
pos_ids = []
|
||||||
|
for seq in sequences:
|
||||||
|
pos_ids.extend(range(len(seq)))
|
||||||
|
return pos_ids
|
||||||
|
|
||||||
|
|
||||||
|
@PositionIdStrategyFactory.register("continuous")
|
||||||
|
class ContinuousPositionId(PositionIdStrategy):
|
||||||
|
def generate(self, sequences: List[List[int]]) -> List[int]:
|
||||||
|
total = sum(len(seq) for seq in sequences)
|
||||||
|
return list(range(total))
|
||||||
@@ -0,0 +1,92 @@
|
|||||||
|
"""Tokenization transform for JSONL record streams.
|
||||||
|
|
||||||
|
Bridges the Reader layer (``JsonlStore`` reads raw JSON records) and the
|
||||||
|
Dataset layer (expects per-record tensors). Holds the tokenizer,
|
||||||
|
mask-builder and position-id strategy together so that I/O code stays
|
||||||
|
free of model dependencies.
|
||||||
|
|
||||||
|
The record-processing core (mask building, primary-id extraction,
|
||||||
|
per-key tensorisation, position-id generation) is shared with
|
||||||
|
:class:`astrai.preprocessing.pipeline.Pipeline` via the
|
||||||
|
:mod:`astrai.preprocessing.core` helpers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Dict, List
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
|
from astrai.preprocessing.core import (
|
||||||
|
build_position_ids,
|
||||||
|
build_preprocessing_components,
|
||||||
|
iter_raw_records,
|
||||||
|
to_per_record_tensors,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TokenizeTransform:
|
||||||
|
"""Tokenize raw JSONL record dicts into per-key tensor lists.
|
||||||
|
|
||||||
|
Owns the three preprocessing concerns that were previously inlined in
|
||||||
|
``JsonlStore``: tokenization, loss-mask construction and position-id
|
||||||
|
generation. Constructing it loads the tokenizer, so it is intentionally
|
||||||
|
cheap to pass around once built.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config: Pipeline config describing sections / masks / position mode.
|
||||||
|
tokenizer_path: Path passed to ``AutoTokenizer.from_pretrained``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, config: PipelineConfig, tokenizer_path: str):
|
||||||
|
self.config = config
|
||||||
|
self.tokenizer, self.mask_builder, self.position_strategy = (
|
||||||
|
build_preprocessing_components(config, tokenizer_path)
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_config_file(cls, config_path: str) -> "TokenizeTransform":
|
||||||
|
"""Build from a ``dataset_config.json`` file path.
|
||||||
|
|
||||||
|
The config file follows :class:`PipelineConfig` schema with an
|
||||||
|
extra ``tokenizer_path`` field. When omitted, the config's
|
||||||
|
parent directory is used as the tokenizer path.
|
||||||
|
"""
|
||||||
|
root = Path(config_path).parent
|
||||||
|
with open(config_path, "r", encoding="utf-8") as f:
|
||||||
|
raw_config = json.load(f)
|
||||||
|
tokenizer_path = raw_config.pop("tokenizer_path", None) or str(root)
|
||||||
|
config = PipelineConfig.from_dict(raw_config)
|
||||||
|
return cls(config, tokenizer_path)
|
||||||
|
|
||||||
|
def apply(self, records: List[dict]) -> Dict[str, list]:
|
||||||
|
"""Tokenize a list of raw record dicts.
|
||||||
|
|
||||||
|
Returns a dict mapping key (``sequence``, ``chosen``, ``responses``,
|
||||||
|
…) to a list of per-record tensors (or nested tensor lists for
|
||||||
|
multi-response keys such as GRPO ``responses``).
|
||||||
|
"""
|
||||||
|
raw: Dict[str, list] = {}
|
||||||
|
doc_sequences: List[List[int]] = []
|
||||||
|
|
||||||
|
for result in iter_raw_records(
|
||||||
|
records, self.mask_builder, self.config, self.tokenizer
|
||||||
|
):
|
||||||
|
primary = None
|
||||||
|
for val in result.values():
|
||||||
|
if isinstance(val, list) and val and isinstance(val[0], int):
|
||||||
|
primary = val
|
||||||
|
break
|
||||||
|
if primary is not None:
|
||||||
|
doc_sequences.append(primary)
|
||||||
|
for key, ids in result.items():
|
||||||
|
raw.setdefault(key, []).append(ids)
|
||||||
|
|
||||||
|
tensors = to_per_record_tensors(raw)
|
||||||
|
|
||||||
|
pos_ids = build_position_ids(doc_sequences, self.position_strategy)
|
||||||
|
if pos_ids is not None:
|
||||||
|
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
||||||
|
|
||||||
|
return tensors
|
||||||
@@ -0,0 +1,56 @@
|
|||||||
|
"""Storage writer strategies for pipeline output.
|
||||||
|
|
||||||
|
The :class:`StoreWriter` abstraction decouples the pipeline from the
|
||||||
|
concrete storage format (bin). The pipeline builds a ``{key:
|
||||||
|
List[Tensor]}`` dict and delegates the write to the writer selected
|
||||||
|
by ``output.storage_format``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import shutil
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from typing import Dict, List
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.serialization import save_bin
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class StoreWriter(ABC):
|
||||||
|
"""Write pre-tokenized tensors to disk in a format-specific way."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def save(
|
||||||
|
self,
|
||||||
|
output_dir: str,
|
||||||
|
domain: str,
|
||||||
|
shard_idx: int,
|
||||||
|
tensors: Dict[str, List[torch.Tensor]],
|
||||||
|
) -> None: ...
|
||||||
|
|
||||||
|
|
||||||
|
class StoreWriterFactory(BaseFactory["StoreWriter"]):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@StoreWriterFactory.register("bin")
|
||||||
|
class BinWriter(StoreWriter):
|
||||||
|
def save(self, output_dir, domain, shard_idx, tensors):
|
||||||
|
shard_path = os.path.join(output_dir, domain, f"shard_{shard_idx:04d}")
|
||||||
|
try:
|
||||||
|
save_bin(shard_path, tensors)
|
||||||
|
except Exception:
|
||||||
|
if os.path.exists(shard_path):
|
||||||
|
shutil.rmtree(shard_path, ignore_errors=True)
|
||||||
|
logger.error(
|
||||||
|
"Failed to write shard %s/%s_%04d, cleaned up partial output",
|
||||||
|
domain,
|
||||||
|
"shard",
|
||||||
|
shard_idx,
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
raise
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
"""Training component protocols — structural subtyping for optimizer/scheduler wrappers."""
|
||||||
|
|
||||||
|
from typing import Any, Protocol, runtime_checkable
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class OptimizerProtocol(Protocol):
|
||||||
|
def step(self, closure=None): ...
|
||||||
|
def zero_grad(self): ...
|
||||||
|
@property
|
||||||
|
def param_groups(self) -> Any: ...
|
||||||
|
def state_dict(self) -> dict: ...
|
||||||
|
def load_state_dict(self, d: dict): ...
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class SchedulerProtocol(Protocol):
|
||||||
|
def step(self): ...
|
||||||
|
def state_dict(self) -> dict: ...
|
||||||
|
def load_state_dict(self, d: dict): ...
|
||||||
|
def get_last_lr(self): ...
|
||||||
@@ -1,77 +0,0 @@
|
|||||||
import json
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any, Dict, Optional
|
|
||||||
|
|
||||||
import safetensors.torch as st
|
|
||||||
import torch
|
|
||||||
import torch.distributed as dist
|
|
||||||
|
|
||||||
from astrai.parallel.setup import get_rank
|
|
||||||
|
|
||||||
|
|
||||||
class Checkpoint:
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
state_dict: Dict[str, Any],
|
|
||||||
epoch: int = 0,
|
|
||||||
iteration: int = 0,
|
|
||||||
extra: Optional[Dict[str, Any]] = None,
|
|
||||||
):
|
|
||||||
self.state_dict = state_dict
|
|
||||||
self.epoch = epoch
|
|
||||||
self.iteration = iteration
|
|
||||||
self.extra = extra or {}
|
|
||||||
|
|
||||||
def save(
|
|
||||||
self,
|
|
||||||
save_dir: str,
|
|
||||||
) -> None:
|
|
||||||
|
|
||||||
save_path = Path(save_dir)
|
|
||||||
save_path.mkdir(parents=True, exist_ok=True)
|
|
||||||
|
|
||||||
rank = get_rank()
|
|
||||||
if rank == 0:
|
|
||||||
meta = {
|
|
||||||
"epoch": self.epoch,
|
|
||||||
"iteration": self.iteration,
|
|
||||||
}
|
|
||||||
with open(save_path / "meta.json", "w") as f:
|
|
||||||
json.dump(meta, f, indent=2)
|
|
||||||
|
|
||||||
st.save_file(self.state_dict, save_path / "state_dict.safetensors")
|
|
||||||
if self.extra:
|
|
||||||
torch.save(self.extra, save_path / "extra.pt")
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def load(
|
|
||||||
cls,
|
|
||||||
save_dir: str,
|
|
||||||
) -> "Checkpoint":
|
|
||||||
|
|
||||||
rank = get_rank()
|
|
||||||
save_path = Path(save_dir)
|
|
||||||
|
|
||||||
meta = {}
|
|
||||||
if rank == 0:
|
|
||||||
with open(Path(save_dir) / "meta.json", "r") as f:
|
|
||||||
meta = json.load(f)
|
|
||||||
|
|
||||||
if dist.is_initialized():
|
|
||||||
meta_list = [meta]
|
|
||||||
dist.broadcast_object_list(meta_list, src=0)
|
|
||||||
meta = meta_list[0]
|
|
||||||
|
|
||||||
state_dict = st.load_file(save_path / "state_dict.safetensors")
|
|
||||||
|
|
||||||
extra = None
|
|
||||||
extra_path = save_path / "extra.pt"
|
|
||||||
if extra_path.exists():
|
|
||||||
extra = torch.load(extra_path, map_location="cpu", weights_only=False)
|
|
||||||
|
|
||||||
return cls(
|
|
||||||
state_dict=state_dict,
|
|
||||||
epoch=meta["epoch"],
|
|
||||||
iteration=meta["iteration"],
|
|
||||||
extra=extra,
|
|
||||||
)
|
|
||||||
@@ -0,0 +1,41 @@
|
|||||||
|
"""Serialization utilities for models and datasets.
|
||||||
|
|
||||||
|
This package re-exports checkpoint helpers and dataset storage helpers so
|
||||||
|
that existing imports from ``astrai.serialization`` continue to work.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from astrai.serialization.checkpoint import (
|
||||||
|
Checkpoint,
|
||||||
|
load_json,
|
||||||
|
load_model_config,
|
||||||
|
load_model_weights,
|
||||||
|
load_safetensors,
|
||||||
|
load_state_dict,
|
||||||
|
load_torch,
|
||||||
|
save_json,
|
||||||
|
save_model,
|
||||||
|
save_safetensors,
|
||||||
|
save_torch,
|
||||||
|
)
|
||||||
|
from astrai.serialization.dataset import (
|
||||||
|
load_bin,
|
||||||
|
load_bin_offsets,
|
||||||
|
save_bin,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"Checkpoint",
|
||||||
|
"load_json",
|
||||||
|
"load_model_config",
|
||||||
|
"load_model_weights",
|
||||||
|
"load_safetensors",
|
||||||
|
"load_state_dict",
|
||||||
|
"load_torch",
|
||||||
|
"save_json",
|
||||||
|
"save_model",
|
||||||
|
"save_safetensors",
|
||||||
|
"save_torch",
|
||||||
|
"load_bin",
|
||||||
|
"load_bin_offsets",
|
||||||
|
"save_bin",
|
||||||
|
]
|
||||||
@@ -0,0 +1,201 @@
|
|||||||
|
"""Model checkpoint serialization helpers."""
|
||||||
|
|
||||||
|
import io
|
||||||
|
import json
|
||||||
|
import time
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Dict, Optional, Union
|
||||||
|
|
||||||
|
import safetensors.torch as st
|
||||||
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
|
|
||||||
|
from astrai.parallel.setup import get_rank
|
||||||
|
|
||||||
|
_META_FILE = "meta.json"
|
||||||
|
_CONFIG_FILE = "config.json"
|
||||||
|
_WEIGHTS_FILE = "model.safetensors"
|
||||||
|
|
||||||
|
|
||||||
|
def save_safetensors(state_dict: dict, path: Union[str, Path]):
|
||||||
|
st.save_file(state_dict, str(path))
|
||||||
|
|
||||||
|
|
||||||
|
def load_safetensors(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||||
|
if not broadcast or not dist.is_initialized():
|
||||||
|
return st.load_file(str(path))
|
||||||
|
|
||||||
|
rank = get_rank()
|
||||||
|
if rank == 0:
|
||||||
|
state_dict = st.load_file(str(path))
|
||||||
|
else:
|
||||||
|
state_dict = {}
|
||||||
|
tmp = [state_dict]
|
||||||
|
dist.broadcast_object_list(tmp, src=0)
|
||||||
|
return tmp[0]
|
||||||
|
|
||||||
|
|
||||||
|
def save_json(data: dict, path: Union[str, Path]):
|
||||||
|
with open(str(path), "w") as f:
|
||||||
|
json.dump(data, f, indent=2)
|
||||||
|
|
||||||
|
|
||||||
|
def load_json(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||||
|
if not broadcast or not dist.is_initialized():
|
||||||
|
with open(str(path), "r") as f:
|
||||||
|
return json.load(f)
|
||||||
|
|
||||||
|
rank = get_rank()
|
||||||
|
if rank == 0:
|
||||||
|
with open(str(path), "r") as f:
|
||||||
|
data = json.load(f)
|
||||||
|
else:
|
||||||
|
data = {}
|
||||||
|
tmp = [data]
|
||||||
|
dist.broadcast_object_list(tmp, src=0)
|
||||||
|
return tmp[0]
|
||||||
|
|
||||||
|
|
||||||
|
def save_torch(obj: Any, path: Union[str, Path]):
|
||||||
|
torch.save(obj, str(path))
|
||||||
|
|
||||||
|
|
||||||
|
def load_torch(path: Union[str, Path], broadcast: bool = False) -> Any:
|
||||||
|
if not broadcast or not dist.is_initialized():
|
||||||
|
return torch.load(str(path), map_location="cpu", weights_only=False)
|
||||||
|
|
||||||
|
path = Path(path)
|
||||||
|
rank = get_rank()
|
||||||
|
|
||||||
|
if rank == 0:
|
||||||
|
with open(path, "rb") as f:
|
||||||
|
raw = f.read()
|
||||||
|
data_tensor = torch.frombuffer(bytearray(raw), dtype=torch.uint8)
|
||||||
|
num_bytes = torch.tensor([len(raw)], dtype=torch.long)
|
||||||
|
else:
|
||||||
|
num_bytes = torch.tensor([0], dtype=torch.long)
|
||||||
|
|
||||||
|
dist.broadcast(num_bytes, src=0)
|
||||||
|
|
||||||
|
if rank != 0:
|
||||||
|
data_tensor = torch.empty(num_bytes.item(), dtype=torch.uint8)
|
||||||
|
|
||||||
|
dist.broadcast(data_tensor, src=0)
|
||||||
|
|
||||||
|
buf = io.BytesIO(data_tensor.numpy().tobytes())
|
||||||
|
return torch.load(buf, map_location="cpu", weights_only=False)
|
||||||
|
|
||||||
|
|
||||||
|
def save_model(config: dict, state_dict: dict, save_directory: str):
|
||||||
|
save_path = Path(save_directory)
|
||||||
|
save_path.mkdir(parents=True, exist_ok=True)
|
||||||
|
save_json(config, save_path / _CONFIG_FILE)
|
||||||
|
save_safetensors(state_dict, save_path / _WEIGHTS_FILE)
|
||||||
|
|
||||||
|
|
||||||
|
def load_model_config(save_directory: str) -> dict:
|
||||||
|
return load_json(Path(save_directory) / _CONFIG_FILE)
|
||||||
|
|
||||||
|
|
||||||
|
def load_model_weights(save_directory: str) -> dict:
|
||||||
|
return load_state_dict(Path(save_directory) / _WEIGHTS_FILE)
|
||||||
|
|
||||||
|
|
||||||
|
def load_state_dict(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||||
|
path = Path(path)
|
||||||
|
if not broadcast or not dist.is_initialized():
|
||||||
|
return load_safetensors(path)
|
||||||
|
|
||||||
|
rank = get_rank()
|
||||||
|
if rank == 0:
|
||||||
|
state_dict = load_safetensors(path)
|
||||||
|
specs = [
|
||||||
|
(k, list(state_dict[k].shape), str(state_dict[k].dtype).split(".")[-1])
|
||||||
|
for k in sorted(state_dict)
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
state_dict = {}
|
||||||
|
specs = []
|
||||||
|
|
||||||
|
specs_list = [specs]
|
||||||
|
dist.broadcast_object_list(specs_list, src=0)
|
||||||
|
specs = specs_list[0]
|
||||||
|
|
||||||
|
for key, shape, dtype_name in specs:
|
||||||
|
dtype = getattr(torch, dtype_name)
|
||||||
|
if rank != 0:
|
||||||
|
tensor = torch.empty(shape, dtype=dtype, device="cpu")
|
||||||
|
else:
|
||||||
|
tensor = state_dict[key].contiguous().cpu()
|
||||||
|
dist.broadcast(tensor, src=0)
|
||||||
|
if rank != 0:
|
||||||
|
state_dict[key] = tensor
|
||||||
|
return state_dict
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class Checkpoint:
|
||||||
|
state_dict: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
epoch: int = 0
|
||||||
|
consumed_samples: int = 0
|
||||||
|
extra: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
meta: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
config: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
def save(self, save_dir: str):
|
||||||
|
save_path = Path(save_dir)
|
||||||
|
save_path.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
meta = {
|
||||||
|
"epoch": self.epoch,
|
||||||
|
"consumed_samples": self.consumed_samples,
|
||||||
|
"timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"),
|
||||||
|
**self.meta,
|
||||||
|
}
|
||||||
|
save_json(meta, save_path / _META_FILE)
|
||||||
|
save_json(self.config, save_path / _CONFIG_FILE)
|
||||||
|
save_safetensors(self.state_dict, save_path / _WEIGHTS_FILE)
|
||||||
|
for key, value in self.extra.items():
|
||||||
|
save_torch(value, save_path / f"{key}.pt")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def load(cls, save_dir: str, broadcast: bool = False) -> "Checkpoint":
|
||||||
|
save_path = Path(save_dir)
|
||||||
|
|
||||||
|
meta = load_json(save_path / _META_FILE, broadcast)
|
||||||
|
config = load_json(save_path / _CONFIG_FILE, broadcast)
|
||||||
|
state_dict = load_state_dict(save_path / _WEIGHTS_FILE, broadcast=broadcast)
|
||||||
|
|
||||||
|
extra = {}
|
||||||
|
for f in sorted(save_path.iterdir()):
|
||||||
|
if f.suffix == ".pt":
|
||||||
|
extra[f.stem] = load_torch(f, broadcast=broadcast)
|
||||||
|
|
||||||
|
return cls(
|
||||||
|
state_dict=state_dict,
|
||||||
|
epoch=meta.get("epoch", 0),
|
||||||
|
consumed_samples=meta.get("consumed_samples", 0),
|
||||||
|
extra=extra,
|
||||||
|
meta=meta,
|
||||||
|
config=config,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def load_any(cls, save_dir: str, broadcast: bool = False) -> Optional["Checkpoint"]:
|
||||||
|
save_path = Path(save_dir)
|
||||||
|
meta_path = save_path / _META_FILE
|
||||||
|
weights_path = save_path / _WEIGHTS_FILE
|
||||||
|
|
||||||
|
if meta_path.exists():
|
||||||
|
return cls.load(save_dir, broadcast=broadcast)
|
||||||
|
|
||||||
|
if weights_path.exists():
|
||||||
|
state_dict = load_state_dict(weights_path, broadcast=broadcast)
|
||||||
|
config = {}
|
||||||
|
config_path = save_path / _CONFIG_FILE
|
||||||
|
if config_path.exists():
|
||||||
|
config = load_json(config_path, broadcast)
|
||||||
|
return cls(state_dict=state_dict, config=config)
|
||||||
|
|
||||||
|
return None
|
||||||
@@ -0,0 +1,82 @@
|
|||||||
|
"""Dataset storage serialization helpers (memory-mapped binary)."""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
def save_bin(
|
||||||
|
file_path: str,
|
||||||
|
tensor_group: Dict[str, List[Tensor]],
|
||||||
|
record_keys: Optional[List[str]] = None,
|
||||||
|
):
|
||||||
|
"""Save tensors as memory-mapped binary files.
|
||||||
|
|
||||||
|
When *record_keys* is provided, those keys are written with per-record
|
||||||
|
cumulative offsets in ``meta.json`` so that ``MmapStore.fetch_record``
|
||||||
|
can slice individual records from the concatenated binary without
|
||||||
|
cross-record concatenation. Keys not in *record_keys* (e.g. SEQ
|
||||||
|
``sequence``) are written as a single contiguous stream without
|
||||||
|
offsets, preserving backward compatibility.
|
||||||
|
|
||||||
|
Nested keys (``List[List[Tensor]]`` such as GRPO ``responses``) are
|
||||||
|
not supported in bin format — use JSONL for those.
|
||||||
|
"""
|
||||||
|
os.makedirs(file_path, exist_ok=True)
|
||||||
|
record_keys = set(record_keys or [])
|
||||||
|
meta = {}
|
||||||
|
for key, tensors in tensor_group.items():
|
||||||
|
if tensors and isinstance(tensors[0], list):
|
||||||
|
raise ValueError(
|
||||||
|
f"Nested key '{key}' (List[List[Tensor]]) is not supported "
|
||||||
|
f"in bin format. Use JSONL storage instead."
|
||||||
|
)
|
||||||
|
cat = torch.cat(tensors, dim=0)
|
||||||
|
entry: Dict[str, Any] = {
|
||||||
|
"shape": list(cat.shape),
|
||||||
|
"dtype": str(cat.dtype).split(".")[-1],
|
||||||
|
}
|
||||||
|
if key in record_keys:
|
||||||
|
offsets = [0]
|
||||||
|
for t in tensors:
|
||||||
|
offsets.append(offsets[-1] + t.shape[0])
|
||||||
|
entry["offsets"] = offsets
|
||||||
|
meta[key] = entry
|
||||||
|
np.asarray(cat.cpu().numpy()).tofile(os.path.join(file_path, f"{key}.bin"))
|
||||||
|
with open(os.path.join(file_path, "meta.json"), "w") as f:
|
||||||
|
json.dump(meta, f)
|
||||||
|
|
||||||
|
|
||||||
|
def load_bin(file_path: str) -> Dict[str, List[Tensor]]:
|
||||||
|
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
||||||
|
meta = json.load(f)
|
||||||
|
segments: Dict[str, List[Tensor]] = {}
|
||||||
|
for key, info in meta.items():
|
||||||
|
arr = np.memmap(
|
||||||
|
os.path.join(file_path, f"{key}.bin"),
|
||||||
|
dtype=info["dtype"],
|
||||||
|
mode="c",
|
||||||
|
shape=tuple(info["shape"]),
|
||||||
|
)
|
||||||
|
segments[key] = [torch.from_numpy(arr)]
|
||||||
|
return segments
|
||||||
|
|
||||||
|
|
||||||
|
def load_bin_offsets(file_path: str) -> Dict[str, List[int]]:
|
||||||
|
"""Read per-record cumulative offsets from ``meta.json``.
|
||||||
|
|
||||||
|
Returns an empty dict when no key has offsets (legacy bin files),
|
||||||
|
in which case record-mode access falls back to per-record segment
|
||||||
|
indexing (JSONL layout).
|
||||||
|
"""
|
||||||
|
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
||||||
|
meta = json.load(f)
|
||||||
|
offsets: Dict[str, List[int]] = {}
|
||||||
|
for key, info in meta.items():
|
||||||
|
if "offsets" in info:
|
||||||
|
offsets[key] = info["offsets"]
|
||||||
|
return offsets
|
||||||
@@ -0,0 +1,53 @@
|
|||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import signal
|
||||||
|
import threading
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_early_stop = threading.Event()
|
||||||
|
_active_context = None
|
||||||
|
|
||||||
|
|
||||||
|
def _early_handler(signum: int, frame):
|
||||||
|
sig = signal.Signals(signum)
|
||||||
|
logger.warning(
|
||||||
|
"Received %s (pid=%d), requesting graceful training stop...",
|
||||||
|
sig.name,
|
||||||
|
os.getpid(),
|
||||||
|
)
|
||||||
|
_early_stop.set()
|
||||||
|
if _active_context is not None:
|
||||||
|
_active_context.request_stop()
|
||||||
|
|
||||||
|
|
||||||
|
def install_early_signal_handlers():
|
||||||
|
for sig in (signal.SIGTERM, signal.SIGINT):
|
||||||
|
signal.signal(sig, _early_handler)
|
||||||
|
_unblock_signals()
|
||||||
|
|
||||||
|
|
||||||
|
def _unblock_signals():
|
||||||
|
try:
|
||||||
|
mask = signal.pthread_sigmask(signal.SIG_BLOCK, set())
|
||||||
|
blocked = {signal.SIGTERM, signal.SIGINT} & mask
|
||||||
|
if blocked:
|
||||||
|
signal.pthread_sigmask(signal.SIG_UNBLOCK, blocked)
|
||||||
|
except (AttributeError, OSError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def register_signal_handlers(context):
|
||||||
|
global _active_context
|
||||||
|
_active_context = context
|
||||||
|
for sig in (signal.SIGTERM, signal.SIGINT):
|
||||||
|
signal.signal(sig, _early_handler)
|
||||||
|
if _early_stop.is_set():
|
||||||
|
context.request_stop()
|
||||||
|
logger.warning("Signal was received during initialization, stopping...")
|
||||||
|
|
||||||
|
|
||||||
|
def unregister_signal_handlers():
|
||||||
|
global _active_context
|
||||||
|
_active_context = None
|
||||||
|
_early_stop.clear()
|
||||||
@@ -1,8 +1,10 @@
|
|||||||
from astrai.tokenize.chat_template import ChatTemplate, MessageType
|
from astrai.tokenize.chat_template import ChatTemplate, MessageType
|
||||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
from astrai.tokenize.tokenizer import AutoTokenizer, Message, Messages
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"AutoTokenizer",
|
"AutoTokenizer",
|
||||||
"ChatTemplate",
|
"ChatTemplate",
|
||||||
"MessageType",
|
"MessageType",
|
||||||
|
"Message",
|
||||||
|
"Messages",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -1,13 +1,11 @@
|
|||||||
from dataclasses import dataclass
|
from functools import cached_property
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
from jinja2 import Template
|
from jinja2 import Template
|
||||||
|
|
||||||
# Message type for chat messages
|
|
||||||
type MessageType = Dict[str, Any]
|
type MessageType = Dict[str, Any]
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class ChatTemplate:
|
class ChatTemplate:
|
||||||
"""A chat template with Jinja2 rendering support.
|
"""A chat template with Jinja2 rendering support.
|
||||||
|
|
||||||
@@ -15,23 +13,51 @@ class ChatTemplate:
|
|||||||
name: Unique identifier for the template.
|
name: Unique identifier for the template.
|
||||||
template_str: Jinja2 template string.
|
template_str: Jinja2 template string.
|
||||||
description: Optional description.
|
description: Optional description.
|
||||||
default_variables: Optional dictionary of default variable values
|
default_variables: Optional dictionary of default variable values.
|
||||||
that will be passed to the template if not overridden during rendering.
|
|
||||||
special_tokens: Optional dictionary mapping token names to their string values.
|
special_tokens: Optional dictionary mapping token names to their string values.
|
||||||
These tokens are automatically added to the template variables.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
name: str
|
def __init__(
|
||||||
template_str: str
|
self,
|
||||||
description: str = ""
|
name: str = "",
|
||||||
default_variables: Dict[str, Any] = None
|
template_str: str = "",
|
||||||
special_tokens: Dict[str, str] = None
|
description: str = "",
|
||||||
|
default_variables: Optional[Dict[str, Any]] = None,
|
||||||
|
special_tokens: Optional[Dict[str, str]] = None,
|
||||||
|
):
|
||||||
|
self.name = name
|
||||||
|
self.template_str = template_str
|
||||||
|
self.description = description
|
||||||
|
self.default_variables = default_variables or {}
|
||||||
|
self.special_tokens = special_tokens or {}
|
||||||
|
|
||||||
def __post_init__(self):
|
@cached_property
|
||||||
if self.default_variables is None:
|
def _compiled(self) -> Template:
|
||||||
self.default_variables = {}
|
"""Lazy-compiled Jinja2 template, cached on first access.
|
||||||
if self.special_tokens is None:
|
|
||||||
self.special_tokens = {}
|
The compiled :class:`~jinja2.Template` holds a dynamically-generated
|
||||||
|
``root`` render function whose ``__module__`` is ``None``; under
|
||||||
|
``pickle`` it falls back to ``__main__`` and breaks ``spawn``-based
|
||||||
|
multiprocessing. :meth:`__getstate__` drops the cached template so
|
||||||
|
that pickle serialises only ``template_str``; each worker rebuilds
|
||||||
|
the cache on first render.
|
||||||
|
"""
|
||||||
|
return Template(self.template_str)
|
||||||
|
|
||||||
|
def __getstate__(self) -> Dict[str, Any]:
|
||||||
|
"""Exclude the cached Jinja2 template from pickling.
|
||||||
|
|
||||||
|
``Template.root_render_func`` is a dynamically generated closure
|
||||||
|
that cannot be pickled by reference. Dropping ``_compiled`` here
|
||||||
|
lets :class:`cached_property` rebuild it on first access after
|
||||||
|
unpickle.
|
||||||
|
"""
|
||||||
|
state = self.__dict__.copy()
|
||||||
|
state.pop("_compiled", None)
|
||||||
|
return state
|
||||||
|
|
||||||
|
def __setstate__(self, state: Dict[str, Any]) -> None:
|
||||||
|
self.__dict__.update(state)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_string(
|
def from_string(
|
||||||
@@ -43,7 +69,7 @@ class ChatTemplate:
|
|||||||
) -> "ChatTemplate":
|
) -> "ChatTemplate":
|
||||||
"""Create a ChatTemplate instance directly from a template string."""
|
"""Create a ChatTemplate instance directly from a template string."""
|
||||||
return cls(
|
return cls(
|
||||||
name="", # empty name for ad‑hoc templates
|
name="",
|
||||||
template_str=template_str,
|
template_str=template_str,
|
||||||
description=description,
|
description=description,
|
||||||
default_variables=default_variables,
|
default_variables=default_variables,
|
||||||
@@ -73,5 +99,4 @@ class ChatTemplate:
|
|||||||
if system_prompt is not None:
|
if system_prompt is not None:
|
||||||
variables["system_prompt"] = system_prompt
|
variables["system_prompt"] = system_prompt
|
||||||
|
|
||||||
jinja_template = Template(self.template_str)
|
return self._compiled.render(**variables)
|
||||||
return jinja_template.render(**variables)
|
|
||||||
|
|||||||
@@ -10,12 +10,16 @@ from tokenizers import Tokenizer
|
|||||||
|
|
||||||
from astrai.tokenize.chat_template import ChatTemplate
|
from astrai.tokenize.chat_template import ChatTemplate
|
||||||
|
|
||||||
|
Message = Dict[str, str]
|
||||||
|
"""Single chat message with ``role`` and ``content`` keys."""
|
||||||
|
|
||||||
|
Messages = List[Message]
|
||||||
|
"""Single conversation — a list of messages."""
|
||||||
|
|
||||||
|
|
||||||
class AutoTokenizer:
|
class AutoTokenizer:
|
||||||
"""Base tokenizer class with automatic loading support"""
|
"""Base tokenizer class with automatic loading support"""
|
||||||
|
|
||||||
TOKENIZER_CLASSES = {} # Registry for auto-loading
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
path: Optional[Union[str, Path]] = None,
|
path: Optional[Union[str, Path]] = None,
|
||||||
@@ -51,9 +55,26 @@ class AutoTokenizer:
|
|||||||
self.set_chat_template(config["chat_template"])
|
self.set_chat_template(config["chat_template"])
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_pretrained(cls, path: Union[str, Path], **kwargs) -> "AutoTokenizer":
|
def from_pretrained(cls, path: Union[str, Path]) -> "AutoTokenizer":
|
||||||
"""Load tokenizer from pretrained directory."""
|
"""Load tokenizer from pretrained directory.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
FileNotFoundError: If tokenizer.json is missing.
|
||||||
|
RuntimeError: If tokenizer failed to initialize.
|
||||||
|
"""
|
||||||
|
path = Path(path)
|
||||||
|
tokenizer_file = path / "tokenizer.json"
|
||||||
|
if not tokenizer_file.exists():
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"Tokenizer file not found: {tokenizer_file}. "
|
||||||
|
"A valid tokenizer.json is required."
|
||||||
|
)
|
||||||
instance = cls(path)
|
instance = cls(path)
|
||||||
|
if instance._tokenizer is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Failed to load tokenizer from {path}. "
|
||||||
|
"The tokenizer.json may be corrupted or incompatible."
|
||||||
|
)
|
||||||
return instance
|
return instance
|
||||||
|
|
||||||
def save_pretrained(self, save_path: str):
|
def save_pretrained(self, save_path: str):
|
||||||
@@ -85,17 +106,6 @@ class AutoTokenizer:
|
|||||||
with open(save_path / "tokenizer_config.json", "w", encoding="utf-8") as f:
|
with open(save_path / "tokenizer_config.json", "w", encoding="utf-8") as f:
|
||||||
json.dump(config, f, ensure_ascii=False, indent=2)
|
json.dump(config, f, ensure_ascii=False, indent=2)
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def register_tokenizer(cls, name: str, tokenizer_class: type):
|
|
||||||
"""
|
|
||||||
Register a new tokenizer class.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
name: Name to register the tokenizer class under
|
|
||||||
tokenizer_class: The tokenizer class to register
|
|
||||||
"""
|
|
||||||
cls.TOKENIZER_CLASSES[name] = tokenizer_class
|
|
||||||
|
|
||||||
def encode(
|
def encode(
|
||||||
self,
|
self,
|
||||||
tokens: Union[str, List[str]],
|
tokens: Union[str, List[str]],
|
||||||
@@ -103,7 +113,16 @@ class AutoTokenizer:
|
|||||||
is_pretokenized: bool = False,
|
is_pretokenized: bool = False,
|
||||||
add_special_tokens: bool = True,
|
add_special_tokens: bool = True,
|
||||||
) -> List:
|
) -> List:
|
||||||
"""Encode text to tokens or token IDs."""
|
"""Encode text to token IDs.
|
||||||
|
|
||||||
|
Accepts both single strings and batches:
|
||||||
|
|
||||||
|
- ``encode("hello")`` → ``[123, 456]``
|
||||||
|
- ``encode(["hello", "world"])`` → ``[[123, 456], [789]]``
|
||||||
|
|
||||||
|
Batches are tokenised in parallel via the Rust backend's
|
||||||
|
``encode_batch`` (uses all available CPU cores).
|
||||||
|
"""
|
||||||
if self._tokenizer is None:
|
if self._tokenizer is None:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"Tokenizer not initialized. Load or create a tokenizer first."
|
"Tokenizer not initialized. Load or create a tokenizer first."
|
||||||
@@ -116,15 +135,13 @@ class AutoTokenizer:
|
|||||||
add_special_tokens=add_special_tokens,
|
add_special_tokens=add_special_tokens,
|
||||||
)
|
)
|
||||||
return encoded.ids if out_ids else encoded.tokens
|
return encoded.ids if out_ids else encoded.tokens
|
||||||
else:
|
|
||||||
encoded_list = self._tokenizer.encode_batch(
|
encoded_list = self._tokenizer.encode_batch(
|
||||||
tokens,
|
tokens,
|
||||||
is_pretokenized=is_pretokenized,
|
is_pretokenized=is_pretokenized,
|
||||||
add_special_tokens=add_special_tokens,
|
add_special_tokens=add_special_tokens,
|
||||||
)
|
)
|
||||||
return [
|
return [encoded.ids if out_ids else encoded.tokens for encoded in encoded_list]
|
||||||
encoded.ids if out_ids else encoded.tokens for encoded in encoded_list
|
|
||||||
]
|
|
||||||
|
|
||||||
def decode(self, tokens: List[int], skip_special_tokens: bool = True) -> str:
|
def decode(self, tokens: List[int], skip_special_tokens: bool = True) -> str:
|
||||||
"""Decode token IDs to text."""
|
"""Decode token IDs to text."""
|
||||||
@@ -147,7 +164,14 @@ class AutoTokenizer:
|
|||||||
- tokenizer.bos_token → returns string
|
- tokenizer.bos_token → returns string
|
||||||
- tokenizer.bos_token_id → returns corresponding integer ID
|
- tokenizer.bos_token_id → returns corresponding integer ID
|
||||||
- tokenizer.stop_ids → returns list of corresponding integer IDs for all special tokens
|
- tokenizer.stop_ids → returns list of corresponding integer IDs for all special tokens
|
||||||
|
|
||||||
|
Internal/private attrs are not intercepted: during unpickle
|
||||||
|
``__dict__`` is empty, so probing ``self._special_token_map``
|
||||||
|
would recurse infinitely.
|
||||||
"""
|
"""
|
||||||
|
if key.startswith("_"):
|
||||||
|
raise AttributeError(key)
|
||||||
|
|
||||||
# Handle stop_ids - return IDs for all special tokens
|
# Handle stop_ids - return IDs for all special tokens
|
||||||
if key == "stop_ids":
|
if key == "stop_ids":
|
||||||
stop_ids = []
|
stop_ids = []
|
||||||
@@ -203,45 +227,63 @@ class AutoTokenizer:
|
|||||||
|
|
||||||
def apply_chat_template(
|
def apply_chat_template(
|
||||||
self,
|
self,
|
||||||
messages: List[Dict[str, str]],
|
messages: Union[Messages, List[Messages]],
|
||||||
system_prompt: Optional[str] = None,
|
system_prompt: Optional[str] = None,
|
||||||
tokenize: bool = True,
|
tokenize: bool = True,
|
||||||
add_generation_prompt: bool = True,
|
add_generation_prompt: bool = True,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> Union[str, List[int]]:
|
) -> Union[str, List[int], List[str], List[List[int]]]:
|
||||||
"""
|
"""Apply the chat template and optionally tokenize.
|
||||||
Apply the chat template to messages and optionally tokenize the result.
|
|
||||||
|
Accepts both single conversations and batches:
|
||||||
|
|
||||||
|
- ``apply_chat_template([msg1, msg2])`` → ``"..."`` or ``[ids]``
|
||||||
|
- ``apply_chat_template([[msg1, msg2], [msg3]])`` → ``["..", ".."]``
|
||||||
|
or ``[[ids], [ids]]``
|
||||||
|
|
||||||
|
Batches render each conversation list and tokenise all at once via
|
||||||
|
:meth:`encode` (``List[str]`` → Rust parallel ``encode_batch``).
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
messages: List of message dicts with 'role' and 'content'.
|
messages: Single conversation (``Messages``) or batch of
|
||||||
system_prompt: Optional system prompt string (auto-converted to first message).
|
conversations (``BatchMessages``).
|
||||||
|
system_prompt: Optional system prompt prepended (single mode only).
|
||||||
tokenize: Whether to return token IDs (True) or raw string (False).
|
tokenize: Whether to return token IDs (True) or raw string (False).
|
||||||
add_generation_prompt: Whether to add the generation prompt (default: True).
|
add_generation_prompt: Whether to add the generation prompt.
|
||||||
**kwargs: Additional variables to pass to the template.
|
**kwargs: Additional template variables.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Either the rendered string or list of token IDs.
|
Single mode: ``str`` or ``List[int]``.
|
||||||
|
Batch mode: ``List[str]`` or ``List[List[int]]``.
|
||||||
Raises:
|
|
||||||
RuntimeError: If chat template is not set.
|
|
||||||
"""
|
"""
|
||||||
if self._chat_template is None:
|
if self._chat_template is None:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"Chat template not set. Use set_chat_template() to set a template first."
|
"Chat template not set. Use set_chat_template() to set a template first."
|
||||||
)
|
)
|
||||||
|
|
||||||
# Auto-convert system_prompt to first message if provided
|
is_batch = bool(messages) and isinstance(messages[0], list)
|
||||||
|
|
||||||
|
if is_batch:
|
||||||
|
rendered = [
|
||||||
|
self._chat_template.render(
|
||||||
|
messages=msgs,
|
||||||
|
add_generation_prompt=add_generation_prompt,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
for msgs in messages
|
||||||
|
]
|
||||||
|
if tokenize:
|
||||||
|
return self.encode(rendered) # List[str] → batch encode
|
||||||
|
return rendered
|
||||||
|
|
||||||
|
# Single conversation
|
||||||
if system_prompt:
|
if system_prompt:
|
||||||
messages = [{"role": "system", "content": system_prompt}] + list(messages)
|
messages = [{"role": "system", "content": system_prompt}] + list(messages)
|
||||||
|
|
||||||
# Render the template
|
|
||||||
rendered = self._chat_template.render(
|
rendered = self._chat_template.render(
|
||||||
messages=messages,
|
messages=messages,
|
||||||
add_generation_prompt=add_generation_prompt,
|
add_generation_prompt=add_generation_prompt,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
if tokenize:
|
if tokenize:
|
||||||
return self.encode(rendered)
|
return self.encode(rendered)
|
||||||
|
|
||||||
return rendered
|
return rendered
|
||||||
|
|||||||
@@ -1,42 +1,70 @@
|
|||||||
from typing import Any, Callable, Dict
|
from typing import Dict
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
|
|
||||||
def _grad_stat(
|
def grad_norm(model: nn.Module, per_param: bool = False) -> float | Dict[str, float]:
|
||||||
model: nn.Module, fn: Callable[[torch.Tensor], Any], default: Any
|
grads = [p.grad.detach() for p in model.parameters() if p.grad is not None]
|
||||||
) -> dict:
|
if not grads:
|
||||||
results = {}
|
return 0.0
|
||||||
for name, param in model.named_parameters():
|
|
||||||
results[name] = default
|
total_sq = torch.stack([g.pow(2).sum() for g in grads]).sum()
|
||||||
if param.grad is not None:
|
if per_param:
|
||||||
results[name] = fn(param.grad.data)
|
norms = {}
|
||||||
return results
|
for name, param in model.named_parameters():
|
||||||
|
if param.grad is not None:
|
||||||
|
norms[name] = param.grad.norm(2).item()
|
||||||
|
else:
|
||||||
|
norms[name] = 0.0
|
||||||
|
norms["total"] = total_sq.sqrt().item()
|
||||||
|
return norms
|
||||||
|
return total_sq.sqrt().item()
|
||||||
|
|
||||||
|
|
||||||
def grad_norm(model: nn.Module, norm_type: int = 2) -> Dict[str, float]:
|
class GradSNRTracker:
|
||||||
return _grad_stat(model, lambda g: g.norm(norm_type).item(), 0.0)
|
"""Track gradient signal-to-noise ratio via EMA of first/second moments.
|
||||||
|
|
||||||
|
SNR = E[g]^2 / Var(g) = E[g]^2 / (E[g^2] - E[g]^2)
|
||||||
|
|
||||||
def grad_std(model: nn.Module) -> Dict[str, float]:
|
The tracker accumulates per-parameter EMA moments across optimizer steps.
|
||||||
return _grad_stat(model, lambda g: g.std().item(), 0.0)
|
Call ``update`` after backward (before ``optimizer.step``) and read
|
||||||
|
``snr`` to get the aggregate SNR across all parameters.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, beta: float = 0.999, eps: float = 1e-8):
|
||||||
|
self.beta = beta
|
||||||
|
self.eps = eps
|
||||||
|
self._first: Dict[int, torch.Tensor] = {}
|
||||||
|
self._second: Dict[int, torch.Tensor] = {}
|
||||||
|
|
||||||
def grad_max(model: nn.Module) -> Dict[str, float]:
|
@torch.no_grad()
|
||||||
return _grad_stat(model, lambda g: g.max().item(), -float("inf"))
|
def update(self, model: nn.Module) -> None:
|
||||||
|
beta = self.beta
|
||||||
|
for param in model.parameters():
|
||||||
|
if param.grad is None:
|
||||||
|
continue
|
||||||
|
pid = id(param)
|
||||||
|
g = param.grad.detach()
|
||||||
|
if pid not in self._first:
|
||||||
|
self._first[pid] = g.clone()
|
||||||
|
self._second[pid] = g.pow(2).clone()
|
||||||
|
else:
|
||||||
|
self._first[pid].mul_(beta).add_(g, alpha=1 - beta)
|
||||||
|
self._second[pid].mul_(beta).addcmul_(g, g, value=1 - beta)
|
||||||
|
|
||||||
|
@property
|
||||||
def grad_min(model: nn.Module) -> Dict[str, float]:
|
def snr(self) -> float:
|
||||||
return _grad_stat(model, lambda g: g.min().item(), float("inf"))
|
if not self._first:
|
||||||
|
return 0.0
|
||||||
|
total_signal = 0.0
|
||||||
def grad_mean(model: nn.Module) -> Dict[str, float]:
|
total_noise = 0.0
|
||||||
return _grad_stat(model, lambda g: g.mean().item(), 0.0)
|
for m, v in zip(self._first.values(), self._second.values()):
|
||||||
|
signal = m.pow(2).sum().item()
|
||||||
|
noise = (v - m.pow(2)).clamp(min=0).sum().item()
|
||||||
def grad_nan_num(model: nn.Module) -> Dict[str, int]:
|
total_signal += signal
|
||||||
return _grad_stat(model, lambda g: g.isnan().sum().item(), 0)
|
total_noise += noise
|
||||||
|
return total_signal / (total_noise + self.eps)
|
||||||
|
|
||||||
|
|
||||||
def ctx_get_loss(ctx):
|
def ctx_get_loss(ctx):
|
||||||
@@ -47,25 +75,16 @@ def ctx_get_lr(ctx):
|
|||||||
return ctx.optimizer.param_groups[-1]["lr"]
|
return ctx.optimizer.param_groups[-1]["lr"]
|
||||||
|
|
||||||
|
|
||||||
|
def ctx_get_val_loss(ctx):
|
||||||
|
return ctx.val_loss
|
||||||
|
|
||||||
|
|
||||||
def ctx_get_grad_norm(ctx):
|
def ctx_get_grad_norm(ctx):
|
||||||
return grad_norm(ctx.model)
|
return ctx.grad_norm
|
||||||
|
|
||||||
|
|
||||||
def ctx_get_grad_std(ctx):
|
def ctx_get_grad_snr(ctx):
|
||||||
return grad_std(ctx.model)
|
tracker = getattr(ctx, "grad_snr_tracker", None)
|
||||||
|
if tracker is None:
|
||||||
|
return None
|
||||||
def ctx_get_grad_max(ctx):
|
return tracker.snr
|
||||||
return grad_max(ctx.model)
|
|
||||||
|
|
||||||
|
|
||||||
def ctx_get_grad_min(ctx):
|
|
||||||
return grad_min(ctx.model)
|
|
||||||
|
|
||||||
|
|
||||||
def ctx_get_grad_mean(ctx):
|
|
||||||
return grad_mean(ctx.model)
|
|
||||||
|
|
||||||
|
|
||||||
def ctx_get_grad_nan_num(ctx):
|
|
||||||
return grad_nan_num(ctx.model)
|
|
||||||
|
|||||||
@@ -0,0 +1,421 @@
|
|||||||
|
"""Online rollout runner for RL training.
|
||||||
|
|
||||||
|
Provides:
|
||||||
|
- :class:`RawRollout` — generation output container (no reward yet)
|
||||||
|
- :class:`RolloutResult` — a :class:`RawRollout` with rewards attached
|
||||||
|
- :class:`BaseRewardModel` — pluggable reward interface
|
||||||
|
- :class:`RolloutGenerator` — KV-cache-backed generation of grouped
|
||||||
|
responses + decoding (no reward); delegates the generation loop to
|
||||||
|
:class:`~astrai.inference.core.scheduler.InferenceScheduler.run_batch`
|
||||||
|
so rollout and the production inference server share one code path
|
||||||
|
- :class:`RolloutRunner` — orchestrates generation + scoring with a
|
||||||
|
step-driven cache; its ``__call__`` returns ``(RolloutResult, is_fresh)``
|
||||||
|
so callers do not need to rely on object identity to detect refreshes.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Dict, List, Optional, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.inference.core.scheduler import InferenceScheduler
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(kw_only=True)
|
||||||
|
class RawRollout:
|
||||||
|
"""Generation output before reward scoring.
|
||||||
|
|
||||||
|
Produced by :class:`RolloutGenerator`; consumed by :class:`RolloutRunner`
|
||||||
|
to assemble a :class:`RolloutResult` once rewards are attached.
|
||||||
|
|
||||||
|
Fields are designed to cover all common RL algorithms:
|
||||||
|
GRPO, PPO, Online DPO, Rejection Sampling, etc.
|
||||||
|
|
||||||
|
Fields:
|
||||||
|
prompts: Tokenized prompts, shape ``[B, P_len]``.
|
||||||
|
prompt_mask: Boolean mask for real prompt tokens, shape ``[B, P_len]``.
|
||||||
|
responses: Generated response token IDs, shape ``[B, G, R_max]``.
|
||||||
|
response_mask: Boolean mask for real (non-pad) response tokens,
|
||||||
|
shape ``[B, G, R_max]``.
|
||||||
|
logprobs_old: Per-token log-probs under the behaviour policy,
|
||||||
|
shape ``[B, G, R_max]``.
|
||||||
|
prompt_texts: Decoded prompt strings (for reward models that
|
||||||
|
need text).
|
||||||
|
response_texts: Decoded response strings, shape ``[B, G]``
|
||||||
|
(for reward models).
|
||||||
|
"""
|
||||||
|
|
||||||
|
prompts: Tensor
|
||||||
|
prompt_mask: Tensor
|
||||||
|
responses: Tensor
|
||||||
|
response_mask: Tensor
|
||||||
|
logprobs_old: Tensor
|
||||||
|
prompt_texts: List[str] = field(default_factory=list)
|
||||||
|
response_texts: List[List[str]] = field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(kw_only=True)
|
||||||
|
class RolloutResult(RawRollout):
|
||||||
|
"""A :class:`RawRollout` with reward scoring attached.
|
||||||
|
|
||||||
|
Produced by :class:`RolloutRunner` once the :class:`BaseRewardModel`
|
||||||
|
has scored the decoded responses.
|
||||||
|
|
||||||
|
Fields:
|
||||||
|
rewards: Reward per response, shape ``[B, G]``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
rewards: Tensor
|
||||||
|
|
||||||
|
|
||||||
|
class BaseRewardModel(ABC):
|
||||||
|
"""Pluggable reward model interface.
|
||||||
|
|
||||||
|
Subclasses should implement ``score()`` to return a ``[B, G]`` float
|
||||||
|
tensor of rewards. Implementations can be:
|
||||||
|
* A loaded reward model (e.g. ArmoRM, Skywork-Reward)
|
||||||
|
* An external API call
|
||||||
|
* A rule-based function (format, length, keyword matching)
|
||||||
|
"""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def score(self, prompts: List[str], responses: List[List[str]]) -> Tensor:
|
||||||
|
"""Score each generated response.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prompts: Raw prompt strings, length ``B``.
|
||||||
|
responses: Generated response strings, shape ``[B, G]``.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Float tensor of shape ``[B, G]``.
|
||||||
|
"""
|
||||||
|
...
|
||||||
|
|
||||||
|
|
||||||
|
_PAD = 0
|
||||||
|
|
||||||
|
|
||||||
|
class RolloutGenerator:
|
||||||
|
"""Pure generation + decoding for a group of responses per prompt.
|
||||||
|
|
||||||
|
Delegates the prefill/decode loop to
|
||||||
|
:meth:`~astrai.inference.core.scheduler.InferenceScheduler.run_batch`,
|
||||||
|
which uses a real KV cache (no O(n²) recompute). Has no dependency
|
||||||
|
on any reward model; can be reused in isolation for offline
|
||||||
|
generation, qualitative sampling, or eval pipelines.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
scheduler: InferenceScheduler,
|
||||||
|
tokenizer,
|
||||||
|
max_tokens: int = 1024,
|
||||||
|
group_size: int = 8,
|
||||||
|
temperature: float = 1.0,
|
||||||
|
top_k: int = 0,
|
||||||
|
top_p: float = 1.0,
|
||||||
|
frequency_penalty: float = 0.0,
|
||||||
|
rep_window: int = 64,
|
||||||
|
):
|
||||||
|
self.scheduler = scheduler
|
||||||
|
self.tokenizer = tokenizer
|
||||||
|
self.max_tokens = max_tokens
|
||||||
|
self.group_size = group_size
|
||||||
|
self.temperature = temperature
|
||||||
|
self.top_k = top_k
|
||||||
|
self.top_p = top_p
|
||||||
|
self.frequency_penalty = frequency_penalty
|
||||||
|
self.rep_window = rep_window
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def generate(self, batch: Dict) -> RawRollout:
|
||||||
|
"""Expand prompts by ``group_size`` and generate one response each.
|
||||||
|
|
||||||
|
Accepted batch formats (per sample, repeated B times):
|
||||||
|
|
||||||
|
- **messages**: ``{"messages": [{"role": "user", "content": "..."}, ...]}``
|
||||||
|
- **instruction + input + output**: ``{"instruction": "...",
|
||||||
|
"input": "...", "output": "..."}`` — mapped to ``system`` /
|
||||||
|
``user`` / ``assistant`` messages; ``input`` and ``output``
|
||||||
|
are optional and skipped when empty.
|
||||||
|
|
||||||
|
Both are rendered through the tokenizer's chat template with
|
||||||
|
``add_generation_prompt=True`` so rollout prompts match the
|
||||||
|
format the policy was SFT-trained on.
|
||||||
|
"""
|
||||||
|
model = self.scheduler._executor.model
|
||||||
|
was_training = model.training
|
||||||
|
model.eval()
|
||||||
|
try:
|
||||||
|
return self._generate_eval(batch)
|
||||||
|
finally:
|
||||||
|
model.train(was_training)
|
||||||
|
|
||||||
|
def _generate_eval(self, batch: Dict) -> RawRollout:
|
||||||
|
prompt_texts, flat_prompt_ids = self._prepare_prompts(batch)
|
||||||
|
B = len(prompt_texts)
|
||||||
|
G = self.group_size
|
||||||
|
# Re-expand flat list to G copies per prompt for run_batch.
|
||||||
|
expanded_prompt_ids: List[List[int]] = []
|
||||||
|
for ids in flat_prompt_ids:
|
||||||
|
expanded_prompt_ids.extend([list(ids)] * G)
|
||||||
|
|
||||||
|
results = self.scheduler.run_batch(
|
||||||
|
expanded_prompt_ids,
|
||||||
|
max_tokens=self.max_tokens,
|
||||||
|
temperature=self.temperature,
|
||||||
|
top_k=self.top_k,
|
||||||
|
top_p=self.top_p,
|
||||||
|
frequency_penalty=self.frequency_penalty,
|
||||||
|
rep_window=self.rep_window,
|
||||||
|
return_logprobs=True,
|
||||||
|
)
|
||||||
|
if len(results) != B * G:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Rollout scheduler returned {len(results)} results, expected {B * G}"
|
||||||
|
)
|
||||||
|
for token_ids, logprobs in results:
|
||||||
|
if len(token_ids) != len(logprobs):
|
||||||
|
raise RuntimeError(
|
||||||
|
"Rollout scheduler returned misaligned token IDs and logprobs"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Each element is (token_ids, logprobs); pad to max length.
|
||||||
|
max_len = 0
|
||||||
|
for token_ids, _lp in results:
|
||||||
|
max_len = max(max_len, len(token_ids))
|
||||||
|
max_len = max(max_len, 1)
|
||||||
|
|
||||||
|
device = self.scheduler.device
|
||||||
|
P_len = max(len(ids) for ids in flat_prompt_ids)
|
||||||
|
prompts_tensor = torch.zeros(B, P_len, dtype=torch.long, device=device)
|
||||||
|
prompt_mask = torch.zeros(B, P_len, dtype=torch.bool, device=device)
|
||||||
|
for i, ids in enumerate(flat_prompt_ids):
|
||||||
|
prompts_tensor[i, -len(ids) :] = torch.tensor(
|
||||||
|
ids, dtype=torch.long, device=device
|
||||||
|
)
|
||||||
|
prompt_mask[i, -len(ids) :] = True
|
||||||
|
|
||||||
|
responses = torch.full((B, G, max_len), _PAD, dtype=torch.long, device=device)
|
||||||
|
response_mask = torch.zeros((B, G, max_len), dtype=torch.bool, device=device)
|
||||||
|
logprobs_old = torch.zeros((B, G, max_len), dtype=torch.float, device=device)
|
||||||
|
|
||||||
|
flat_idx = 0
|
||||||
|
response_texts: List[List[str]] = [[] for _ in range(B)]
|
||||||
|
for i in range(B):
|
||||||
|
for g in range(G):
|
||||||
|
token_ids, lps = results[flat_idx]
|
||||||
|
flat_idx += 1
|
||||||
|
n = len(token_ids)
|
||||||
|
if n:
|
||||||
|
responses[i, g, :n] = torch.tensor(
|
||||||
|
token_ids, dtype=torch.long, device=device
|
||||||
|
)
|
||||||
|
response_mask[i, g, :n] = True
|
||||||
|
logprobs_old[i, g, :n] = torch.tensor(
|
||||||
|
lps, dtype=torch.float, device=device
|
||||||
|
)
|
||||||
|
response_texts[i].append(
|
||||||
|
self.tokenizer.decode(token_ids, skip_special_tokens=True)
|
||||||
|
)
|
||||||
|
|
||||||
|
return RawRollout(
|
||||||
|
prompts=prompts_tensor,
|
||||||
|
prompt_mask=prompt_mask,
|
||||||
|
responses=responses,
|
||||||
|
response_mask=response_mask,
|
||||||
|
logprobs_old=logprobs_old,
|
||||||
|
prompt_texts=prompt_texts,
|
||||||
|
response_texts=response_texts,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _prepare_prompts(self, batch: Dict) -> Tuple[List[str], List[List[int]]]:
|
||||||
|
"""Render batch prompts to ``(texts, token_id_lists)``.
|
||||||
|
|
||||||
|
Returns two parallel lists of length B (number of prompts in
|
||||||
|
the batch). Dispatches by batch keys:
|
||||||
|
|
||||||
|
- ``"messages"``: treated as a pre-built message list per sample.
|
||||||
|
- ``"instruction"`` (optionally ``"input"`` and ``"output"``): mapped
|
||||||
|
to ``system`` / ``user`` / ``assistant`` messages respectively.
|
||||||
|
|
||||||
|
Both paths go through the tokenizer's chat template with
|
||||||
|
``add_generation_prompt=True``.
|
||||||
|
"""
|
||||||
|
if "messages" in batch:
|
||||||
|
messages_list = batch["messages"]
|
||||||
|
elif "instruction" in batch:
|
||||||
|
instructions = batch["instruction"]
|
||||||
|
B = len(instructions)
|
||||||
|
inputs = batch.get("input") or [""] * B
|
||||||
|
outputs = batch.get("output") or [""] * B
|
||||||
|
messages_list = [
|
||||||
|
self._instruction_to_messages(i, u, o)
|
||||||
|
for i, u, o in zip(instructions, inputs, outputs)
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
"Rollout batch must contain either 'messages' or "
|
||||||
|
"'instruction' (optionally 'input'/'output'); got keys: "
|
||||||
|
f"{list(batch.keys())}"
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
prompt_texts = self.tokenizer.apply_chat_template(
|
||||||
|
messages_list, tokenize=False, add_generation_prompt=True
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
not isinstance(prompt_texts, list)
|
||||||
|
or len(prompt_texts) != len(messages_list)
|
||||||
|
or not all(isinstance(text, str) for text in prompt_texts)
|
||||||
|
):
|
||||||
|
raise TypeError("Tokenizer does not support batched chat templates")
|
||||||
|
flat_prompt_ids = self.tokenizer.encode(prompt_texts)
|
||||||
|
if len(flat_prompt_ids) != len(messages_list) or not all(
|
||||||
|
isinstance(ids, list) for ids in flat_prompt_ids
|
||||||
|
):
|
||||||
|
raise TypeError("Tokenizer does not support batched encoding")
|
||||||
|
except (TypeError, IndexError, KeyError):
|
||||||
|
# Keep compatibility with lightweight tokenizer adapters that only
|
||||||
|
# implement the single-conversation template API.
|
||||||
|
prompt_texts = []
|
||||||
|
flat_prompt_ids = []
|
||||||
|
for messages in messages_list:
|
||||||
|
text = self.tokenizer.apply_chat_template(
|
||||||
|
messages, tokenize=False, add_generation_prompt=True
|
||||||
|
)
|
||||||
|
ids = self.tokenizer.apply_chat_template(
|
||||||
|
messages, tokenize=True, add_generation_prompt=True
|
||||||
|
)
|
||||||
|
prompt_texts.append(text)
|
||||||
|
flat_prompt_ids.append(list(ids))
|
||||||
|
return prompt_texts, flat_prompt_ids
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _instruction_to_messages(
|
||||||
|
instruction: str, inp: str = "", output: str = ""
|
||||||
|
) -> List[Dict[str, str]]:
|
||||||
|
"""Map instruction/input/output to chat messages.
|
||||||
|
|
||||||
|
Role mapping follows the convention used throughout the
|
||||||
|
preprocessing pipeline: ``instruction`` → system, ``input`` →
|
||||||
|
user, ``output`` → assistant. Empty fields are skipped so a
|
||||||
|
bare instruction produces a ``[system]`` list and the chat
|
||||||
|
template's ``add_generation_prompt`` adds the assistant header
|
||||||
|
for sampling.
|
||||||
|
"""
|
||||||
|
messages: List[Dict[str, str]] = []
|
||||||
|
if instruction:
|
||||||
|
messages.append({"role": "system", "content": instruction})
|
||||||
|
if inp:
|
||||||
|
messages.append({"role": "user", "content": inp})
|
||||||
|
if output:
|
||||||
|
messages.append({"role": "assistant", "content": output})
|
||||||
|
return messages
|
||||||
|
|
||||||
|
|
||||||
|
class RolloutRunner:
|
||||||
|
"""Produces :class:`RolloutResult` from a prompt batch.
|
||||||
|
|
||||||
|
Composes a :class:`RolloutGenerator` (generation + decoding) with a
|
||||||
|
:class:`BaseRewardModel` (scoring). Maintains an internal cache so
|
||||||
|
the same batch prompt can be replayed for multiple gradient steps.
|
||||||
|
A new rollout is triggered every ``rollout_interval`` calls to
|
||||||
|
:meth:`step` (or after :meth:`clear_cache`).
|
||||||
|
|
||||||
|
The ``__call__`` contract returns a ``(RolloutResult, is_fresh)``
|
||||||
|
tuple — callers must use the boolean to detect a refreshed rollout
|
||||||
|
rather than relying on object identity.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
generator = RolloutGenerator(policy, tokenizer, pipeline, ...)
|
||||||
|
runner = RolloutRunner(generator, reward_model, rollout_interval=512)
|
||||||
|
result, is_fresh = runner(prompt_batch)
|
||||||
|
if is_fresh:
|
||||||
|
... # e.g. sync behaviour policy
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
generator: RolloutGenerator,
|
||||||
|
reward_model: BaseRewardModel,
|
||||||
|
rollout_interval: int = 512,
|
||||||
|
):
|
||||||
|
self.generator = generator
|
||||||
|
self.reward_model = reward_model
|
||||||
|
self.rollout_interval = rollout_interval
|
||||||
|
|
||||||
|
self._cache: Optional[RolloutResult] = None
|
||||||
|
self._cache_key = None
|
||||||
|
self._steps_since_rollout: int = 0
|
||||||
|
|
||||||
|
def step(self):
|
||||||
|
"""Advance the internal counter (call once per optimizer step)."""
|
||||||
|
self._steps_since_rollout += 1
|
||||||
|
|
||||||
|
def clear_cache(self):
|
||||||
|
"""Force next call to re-run rollout."""
|
||||||
|
self._cache = None
|
||||||
|
self._cache_key = None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _batch_key(batch: Dict):
|
||||||
|
"""Build a stable key for the prompt fields accepted by the generator."""
|
||||||
|
|
||||||
|
def freeze(value):
|
||||||
|
if isinstance(value, dict):
|
||||||
|
return tuple(sorted((key, freeze(val)) for key, val in value.items()))
|
||||||
|
if isinstance(value, (list, tuple)):
|
||||||
|
return tuple(freeze(item) for item in value)
|
||||||
|
return value
|
||||||
|
|
||||||
|
fields = ("messages", "instruction", "input", "output")
|
||||||
|
return tuple(
|
||||||
|
(field, freeze(batch[field])) for field in fields if field in batch
|
||||||
|
)
|
||||||
|
|
||||||
|
def _score(self, raw: RawRollout) -> RolloutResult:
|
||||||
|
rewards = self.reward_model.score(raw.prompt_texts, raw.response_texts)
|
||||||
|
if not isinstance(rewards, Tensor):
|
||||||
|
rewards = torch.as_tensor(rewards, dtype=torch.float32)
|
||||||
|
expected_shape = raw.responses.shape[:2]
|
||||||
|
if rewards.shape != expected_shape:
|
||||||
|
raise ValueError(
|
||||||
|
f"Reward model returned shape {tuple(rewards.shape)}, "
|
||||||
|
f"expected {tuple(expected_shape)}"
|
||||||
|
)
|
||||||
|
if not torch.isfinite(rewards).all():
|
||||||
|
raise ValueError("Reward model returned non-finite values")
|
||||||
|
device = raw.prompts.device
|
||||||
|
return RolloutResult(
|
||||||
|
prompts=raw.prompts,
|
||||||
|
prompt_mask=raw.prompt_mask,
|
||||||
|
responses=raw.responses,
|
||||||
|
response_mask=raw.response_mask,
|
||||||
|
rewards=rewards.to(device=device),
|
||||||
|
logprobs_old=raw.logprobs_old,
|
||||||
|
prompt_texts=raw.prompt_texts,
|
||||||
|
response_texts=raw.response_texts,
|
||||||
|
)
|
||||||
|
|
||||||
|
def __call__(self, batch: Dict[str, Tensor]) -> Tuple[RolloutResult, bool]:
|
||||||
|
"""Return ``(cached or fresh) RolloutResult`` plus an ``is_fresh`` flag.
|
||||||
|
|
||||||
|
Triggers a new rollout when ``_steps_since_rollout >= rollout_interval``
|
||||||
|
or when the cache is empty.
|
||||||
|
"""
|
||||||
|
cache_key = self._batch_key(batch)
|
||||||
|
if (
|
||||||
|
self._cache is None
|
||||||
|
or cache_key != self._cache_key
|
||||||
|
or self._steps_since_rollout >= self.rollout_interval
|
||||||
|
):
|
||||||
|
raw = self.generator.generate(batch)
|
||||||
|
self._cache = self._score(raw)
|
||||||
|
self._cache_key = cache_key
|
||||||
|
self._steps_since_rollout = 0
|
||||||
|
return self._cache, True
|
||||||
|
return self._cache, False
|
||||||
+75
-34
@@ -2,7 +2,7 @@
|
|||||||
|
|
||||||
import math
|
import math
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import Any, Dict, List, Type
|
from typing import Any, Dict, List
|
||||||
|
|
||||||
from torch.optim.lr_scheduler import LRScheduler
|
from torch.optim.lr_scheduler import LRScheduler
|
||||||
|
|
||||||
@@ -31,7 +31,6 @@ class SchedulerFactory(BaseFactory["BaseScheduler"]):
|
|||||||
"""Factory class for creating learning rate schedulers.
|
"""Factory class for creating learning rate schedulers.
|
||||||
|
|
||||||
Supports decorator-based registration for extensible scheduler types.
|
Supports decorator-based registration for extensible scheduler types.
|
||||||
Also supports creation from ScheduleConfig objects.
|
|
||||||
|
|
||||||
Example usage:
|
Example usage:
|
||||||
@SchedulerFactory.register("custom")
|
@SchedulerFactory.register("custom")
|
||||||
@@ -41,33 +40,6 @@ class SchedulerFactory(BaseFactory["BaseScheduler"]):
|
|||||||
scheduler = SchedulerFactory.create("custom", optimizer, **kwargs)
|
scheduler = SchedulerFactory.create("custom", optimizer, **kwargs)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def _validate_component(cls, scheduler_cls: Type[BaseScheduler]) -> None:
|
|
||||||
"""Validate that the scheduler class inherits from BaseScheduler."""
|
|
||||||
if not issubclass(scheduler_cls, BaseScheduler):
|
|
||||||
raise TypeError(f"{scheduler_cls.__name__} must inherit from BaseScheduler")
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def create(
|
|
||||||
cls, optimizer, schedule_type: str = "none", **kwargs
|
|
||||||
) -> "BaseScheduler":
|
|
||||||
"""Create a scheduler instance by type name.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
optimizer: PyTorch optimizer
|
|
||||||
schedule_type: Type of scheduler ("cosine", "sgdr")
|
|
||||||
**kwargs: Arguments passed to the scheduler constructor
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Scheduler instance
|
|
||||||
"""
|
|
||||||
return super().create(schedule_type, optimizer, **kwargs)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def available_types(cls) -> list:
|
|
||||||
"""Return list of registered scheduler type names."""
|
|
||||||
return cls.list_registered()
|
|
||||||
|
|
||||||
|
|
||||||
# ----------- Scheduler implementations -----------
|
# ----------- Scheduler implementations -----------
|
||||||
|
|
||||||
@@ -81,7 +53,7 @@ class CosineScheduler(BaseScheduler):
|
|||||||
optimizer,
|
optimizer,
|
||||||
warmup_steps: int,
|
warmup_steps: int,
|
||||||
lr_decay_steps: int,
|
lr_decay_steps: int,
|
||||||
min_rate: float = 0.05,
|
min_rate: float = 0.01,
|
||||||
last_epoch: int = -1,
|
last_epoch: int = -1,
|
||||||
):
|
):
|
||||||
self.warmup_steps = warmup_steps
|
self.warmup_steps = warmup_steps
|
||||||
@@ -93,11 +65,15 @@ class CosineScheduler(BaseScheduler):
|
|||||||
def get_lr(self) -> List[float]:
|
def get_lr(self) -> List[float]:
|
||||||
# warmup
|
# warmup
|
||||||
if self.last_epoch < self.warmup_steps:
|
if self.last_epoch < self.warmup_steps:
|
||||||
warmup_factor = max(self.min_rate, self.last_epoch / self.warmup_steps)
|
warmup_factor = max(
|
||||||
|
self.min_rate, self.last_epoch / max(self.warmup_steps, 1)
|
||||||
|
)
|
||||||
return [base_lr * warmup_factor for base_lr in self.base_lrs]
|
return [base_lr * warmup_factor for base_lr in self.base_lrs]
|
||||||
|
|
||||||
# cosine decay
|
# cosine decay
|
||||||
decay_progress = (self.last_epoch - self.warmup_steps) / self.lr_decay_steps
|
decay_progress = (self.last_epoch - self.warmup_steps) / max(
|
||||||
|
self.lr_decay_steps, 1
|
||||||
|
)
|
||||||
decay_progress = min(decay_progress, 1.0)
|
decay_progress = min(decay_progress, 1.0)
|
||||||
cosine_decay = 0.5 * (1.0 + math.cos(math.pi * decay_progress))
|
cosine_decay = 0.5 * (1.0 + math.cos(math.pi * decay_progress))
|
||||||
decay_factor = max(self.min_rate, cosine_decay)
|
decay_factor = max(self.min_rate, cosine_decay)
|
||||||
@@ -132,7 +108,7 @@ class SGDRScheduler(BaseScheduler):
|
|||||||
optimizer,
|
optimizer,
|
||||||
warmup_steps: int,
|
warmup_steps: int,
|
||||||
cycle_length: int,
|
cycle_length: int,
|
||||||
min_rate: float = 0.05,
|
min_rate: float = 0.01,
|
||||||
t_mult: int = 2,
|
t_mult: int = 2,
|
||||||
last_epoch: int = -1,
|
last_epoch: int = -1,
|
||||||
):
|
):
|
||||||
@@ -146,7 +122,9 @@ class SGDRScheduler(BaseScheduler):
|
|||||||
def get_lr(self):
|
def get_lr(self):
|
||||||
# warmup
|
# warmup
|
||||||
if self.last_epoch < self.warmup_steps:
|
if self.last_epoch < self.warmup_steps:
|
||||||
warmup_factor = max(self.min_rate, self.last_epoch / self.warmup_steps)
|
warmup_factor = max(
|
||||||
|
self.min_rate, self.last_epoch / max(self.warmup_steps, 1)
|
||||||
|
)
|
||||||
return [base_lr * warmup_factor for base_lr in self.base_lrs]
|
return [base_lr * warmup_factor for base_lr in self.base_lrs]
|
||||||
|
|
||||||
# SGDR
|
# SGDR
|
||||||
@@ -192,3 +170,66 @@ class SGDRScheduler(BaseScheduler):
|
|||||||
self.min_rate = state_dict.pop("min_rate")
|
self.min_rate = state_dict.pop("min_rate")
|
||||||
self.t_mult = state_dict.pop("t_mult")
|
self.t_mult = state_dict.pop("t_mult")
|
||||||
super().load_state_dict(state_dict)
|
super().load_state_dict(state_dict)
|
||||||
|
|
||||||
|
|
||||||
|
@SchedulerFactory.register("wsd")
|
||||||
|
class WSDScheduler(BaseScheduler):
|
||||||
|
"""WSD (Warmup-Stable-Decay) scheduler with sqrt cooldown.
|
||||||
|
|
||||||
|
warmup_steps: linear warmup from min_rate to 1.0
|
||||||
|
stable_steps: constant at base_lr
|
||||||
|
decay_steps: sqrt decay from base_lr to min_rate
|
||||||
|
min_rate: minimum lr as fraction of base_lr (default 0.0)
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
optimizer,
|
||||||
|
warmup_steps: int,
|
||||||
|
stable_steps: int,
|
||||||
|
decay_steps: int,
|
||||||
|
min_rate: float = 0.01,
|
||||||
|
last_epoch: int = -1,
|
||||||
|
):
|
||||||
|
self.warmup_steps = warmup_steps
|
||||||
|
self.stable_steps = stable_steps
|
||||||
|
self.decay_steps = decay_steps
|
||||||
|
self.min_rate = min_rate
|
||||||
|
self.total_steps = warmup_steps + stable_steps + decay_steps
|
||||||
|
super().__init__(optimizer, last_epoch)
|
||||||
|
|
||||||
|
def get_lr(self) -> List[float]:
|
||||||
|
if self.last_epoch < self.warmup_steps:
|
||||||
|
factor = max(self.min_rate, self.last_epoch / max(self.warmup_steps, 1))
|
||||||
|
return [base_lr * factor for base_lr in self.base_lrs]
|
||||||
|
|
||||||
|
offset = self.last_epoch - self.warmup_steps
|
||||||
|
|
||||||
|
if offset < self.stable_steps:
|
||||||
|
return list(self.base_lrs)
|
||||||
|
|
||||||
|
decay_ratio = (offset - self.stable_steps) / max(self.decay_steps, 1)
|
||||||
|
decay_ratio = min(decay_ratio, 1.0)
|
||||||
|
factor = (1.0 - self.min_rate) * (1.0 - decay_ratio) ** 2 + self.min_rate
|
||||||
|
return [base_lr * factor for base_lr in self.base_lrs]
|
||||||
|
|
||||||
|
def state_dict(self):
|
||||||
|
state = super().state_dict()
|
||||||
|
state.update(
|
||||||
|
{
|
||||||
|
"warmup_steps": self.warmup_steps,
|
||||||
|
"stable_steps": self.stable_steps,
|
||||||
|
"decay_steps": self.decay_steps,
|
||||||
|
"min_rate": self.min_rate,
|
||||||
|
"total_steps": self.total_steps,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return state
|
||||||
|
|
||||||
|
def load_state_dict(self, state_dict):
|
||||||
|
self.warmup_steps = state_dict.pop("warmup_steps")
|
||||||
|
self.stable_steps = state_dict.pop("stable_steps")
|
||||||
|
self.decay_steps = state_dict.pop("decay_steps")
|
||||||
|
self.min_rate = state_dict.pop("min_rate")
|
||||||
|
self.total_steps = state_dict.pop("total_steps")
|
||||||
|
super().load_state_dict(state_dict)
|
||||||
|
|||||||
+270
-104
@@ -1,55 +1,37 @@
|
|||||||
"""Training strategy implementations with factory pattern."""
|
"""Training strategy implementations with factory pattern."""
|
||||||
|
|
||||||
import copy
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import Any, Callable, Dict, Union
|
from typing import Callable, Dict, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
|
||||||
|
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.parallel.executor import broadcast_state_dict
|
||||||
|
from astrai.trainer.rollout import RolloutResult
|
||||||
|
|
||||||
|
|
||||||
def unwrap_model(model: nn.Module) -> nn.Module:
|
def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
|
||||||
"""Unwrap DDP wrapper if present to get the original model."""
|
|
||||||
if isinstance(model, DDP):
|
|
||||||
return model.module
|
|
||||||
return model
|
|
||||||
|
|
||||||
|
|
||||||
def create_ref_model(model: nn.Module) -> nn.Module:
|
|
||||||
"""Create a reference model for DPO/GRPO training.
|
|
||||||
|
|
||||||
Handles DDP-wrapped models safely by unwrapping first,
|
|
||||||
then creating a deep copy with frozen gradients.
|
|
||||||
"""
|
|
||||||
original_model = unwrap_model(model)
|
|
||||||
ref_model = copy.deepcopy(original_model)
|
|
||||||
ref_model.requires_grad_(False)
|
|
||||||
ref_model.eval()
|
|
||||||
return ref_model
|
|
||||||
|
|
||||||
|
|
||||||
def move_to_device(batch: Dict[str, Tensor], device: str) -> Any:
|
|
||||||
"""Move batch tensors to specified device with non-blocking transfer."""
|
"""Move batch tensors to specified device with non-blocking transfer."""
|
||||||
return {key: value.to(device, non_blocking=True) for key, value in batch.items()}
|
return {key: value.to(device, non_blocking=True) for key, value in batch.items()}
|
||||||
|
|
||||||
|
|
||||||
def get_logprobs(
|
def get_logprobs(
|
||||||
model: Union[nn.Module, Callable[..., Dict[str, Tensor]]],
|
model: nn.Module,
|
||||||
input_ids: Tensor,
|
input_ids: Tensor,
|
||||||
mask: Tensor,
|
attn_mask: Tensor,
|
||||||
|
loss_mask: Tensor,
|
||||||
reduction: str,
|
reduction: str,
|
||||||
):
|
) -> Tensor:
|
||||||
"""Compute token-wise log probabilities from model outputs.
|
"""Compute token-wise log probabilities from model outputs.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
model: The language model
|
model: The language model
|
||||||
input_ids: Input token IDs of shape [batch_size, seq_len]
|
input_ids: Input token IDs of shape [batch_size, seq_len]
|
||||||
mask: Attention mask of shape [batch_size, seq_len]
|
attn_mask: Attention mask passed to the model (may include causal).
|
||||||
|
loss_mask: Per-token mask for loss reduction.
|
||||||
reduction: How to reduce over sequence dimension ("mean", "sum", "none")
|
reduction: How to reduce over sequence dimension ("mean", "sum", "none")
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
@@ -62,9 +44,12 @@ def get_logprobs(
|
|||||||
)
|
)
|
||||||
|
|
||||||
shifted_input_ids = input_ids[:, 1:]
|
shifted_input_ids = input_ids[:, 1:]
|
||||||
shifted_mask = mask[:, 1:]
|
shifted_loss_mask = loss_mask[:, 1:]
|
||||||
|
|
||||||
logits = model(input_ids[:, :-1], mask[:, :-1])["logits"]
|
logits = model(
|
||||||
|
input_ids[:, :-1],
|
||||||
|
attn_mask[:, :, :-1, :-1] if attn_mask.dim() == 4 else attn_mask[:, :-1],
|
||||||
|
)["logits"]
|
||||||
log_probs = torch.log_softmax(logits.float(), dim=-1)
|
log_probs = torch.log_softmax(logits.float(), dim=-1)
|
||||||
|
|
||||||
token_logprobs = torch.gather(
|
token_logprobs = torch.gather(
|
||||||
@@ -72,24 +57,53 @@ def get_logprobs(
|
|||||||
).squeeze(-1)
|
).squeeze(-1)
|
||||||
|
|
||||||
if reduction == "mean":
|
if reduction == "mean":
|
||||||
return (token_logprobs * shifted_mask).sum(dim=-1) / shifted_mask.sum(
|
return (token_logprobs * shifted_loss_mask).sum(dim=-1) / shifted_loss_mask.sum(
|
||||||
dim=-1
|
dim=-1
|
||||||
).clamp(min=1.0)
|
).clamp(min=1.0)
|
||||||
elif reduction == "sum":
|
elif reduction == "sum":
|
||||||
return (token_logprobs * shifted_mask).sum(dim=-1)
|
return (token_logprobs * shifted_loss_mask).sum(dim=-1)
|
||||||
else:
|
else:
|
||||||
return token_logprobs * shifted_mask
|
return token_logprobs * shifted_loss_mask
|
||||||
|
|
||||||
|
|
||||||
|
def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
|
||||||
|
S = position_ids.size(1)
|
||||||
|
device = position_ids.device
|
||||||
|
boundaries = position_ids[:, 1:] <= position_ids[:, :-1]
|
||||||
|
doc_ids = torch.cat(
|
||||||
|
[
|
||||||
|
torch.zeros(position_ids.size(0), 1, dtype=torch.long, device=device),
|
||||||
|
boundaries.long().cumsum(dim=1),
|
||||||
|
],
|
||||||
|
dim=1,
|
||||||
|
)
|
||||||
|
same_doc = doc_ids.unsqueeze(-1) == doc_ids.unsqueeze(-2)
|
||||||
|
causal = torch.tril(torch.ones(S, S, dtype=torch.bool, device=device))
|
||||||
|
return (same_doc & causal).unsqueeze(1)
|
||||||
|
|
||||||
|
|
||||||
class BaseStrategy(ABC):
|
class BaseStrategy(ABC):
|
||||||
"""Abstract base class for training strategies."""
|
"""Abstract base class for training strategies.
|
||||||
|
|
||||||
|
When a :class:`~astrai.trainer.rollout.RolloutRunner` is injected via
|
||||||
|
:meth:`set_rollout_runner`, the strategy transparently switches to
|
||||||
|
online mode: each ``__call__`` produces a :class:`RolloutResult`,
|
||||||
|
converts it to a training batch via :meth:`prepare_from_rollout`, and
|
||||||
|
then computes the loss. Without a runner the strategy runs in
|
||||||
|
offline mode and consumes the batch directly.
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self, model: Union[Callable[..., Dict[str, Tensor]]], device: str, **kwargs
|
self,
|
||||||
|
model: Union[nn.Module, Callable[..., Dict[str, Tensor]]],
|
||||||
|
device: str,
|
||||||
|
**kwargs,
|
||||||
):
|
):
|
||||||
self.model = model
|
self.model = model
|
||||||
self.device = device
|
self.device = device
|
||||||
|
self.executor = kwargs.pop("executor", None)
|
||||||
self.extra_kwargs = kwargs
|
self.extra_kwargs = kwargs
|
||||||
|
self._rollout_runner = None
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
@@ -103,9 +117,53 @@ class BaseStrategy(ABC):
|
|||||||
"""
|
"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def supports_online(self) -> bool:
|
||||||
|
"""Whether this strategy can operate with a rollout runner.
|
||||||
|
|
||||||
|
Base implementation returns ``False``; strategies that implement
|
||||||
|
:meth:`prepare_from_rollout` should override to return ``True``.
|
||||||
|
"""
|
||||||
|
return False
|
||||||
|
|
||||||
|
def set_rollout_runner(self, runner):
|
||||||
|
"""Inject a :class:`RolloutRunner` to enable online rollout mode."""
|
||||||
|
self._rollout_runner = runner
|
||||||
|
|
||||||
|
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
|
||||||
|
"""Map a :class:`RolloutResult` to the batch layout expected by
|
||||||
|
:meth:`compute_loss`.
|
||||||
|
|
||||||
|
Strategies that return ``True`` from :meth:`supports_online` must
|
||||||
|
override this. Default raises :class:`NotImplementedError`.
|
||||||
|
"""
|
||||||
|
raise NotImplementedError(
|
||||||
|
f"{type(self).__name__} does not support online rollout"
|
||||||
|
)
|
||||||
|
|
||||||
|
def _on_rollout_refresh(self):
|
||||||
|
"""Hook fired when a fresh rollout result is produced.
|
||||||
|
|
||||||
|
Override to refresh stale state (e.g. syncing the behaviour
|
||||||
|
policy). Default is a no-op.
|
||||||
|
"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
def on_optimizer_step(self):
|
||||||
|
"""Advance online rollout state after a successful optimizer step."""
|
||||||
|
if self._rollout_runner is not None:
|
||||||
|
self._rollout_runner.step()
|
||||||
|
|
||||||
def __call__(self, batch: Dict[str, Tensor]) -> Tensor:
|
def __call__(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
"""Allow calling strategy directly as a callable."""
|
"""Run offline or online forward depending on runner injection."""
|
||||||
return self.compute_loss(batch)
|
if self._rollout_runner is None:
|
||||||
|
return self.compute_loss(batch)
|
||||||
|
|
||||||
|
result, is_fresh = self._rollout_runner(batch)
|
||||||
|
if is_fresh:
|
||||||
|
self._on_rollout_refresh()
|
||||||
|
|
||||||
|
train_batch = self.prepare_from_rollout(result)
|
||||||
|
return self.compute_loss(train_batch)
|
||||||
|
|
||||||
|
|
||||||
class StrategyFactory(BaseFactory["BaseStrategy"]):
|
class StrategyFactory(BaseFactory["BaseStrategy"]):
|
||||||
@@ -122,32 +180,6 @@ class StrategyFactory(BaseFactory["BaseStrategy"]):
|
|||||||
strategy = StrategyFactory.create("custom", model, device)
|
strategy = StrategyFactory.create("custom", model, device)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def _validate_component(cls, strategy_cls: type) -> None:
|
|
||||||
"""Validate that the strategy class inherits from BaseStrategy."""
|
|
||||||
if not issubclass(strategy_cls, BaseStrategy):
|
|
||||||
raise TypeError(f"{strategy_cls.__name__} must inherit from BaseStrategy")
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def create(cls, train_type: str, model, device: str, **kwargs) -> "BaseStrategy":
|
|
||||||
"""Create a strategy instance based on training type.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
train_type: Type of training ("seq", "sft", "dpo", "grpo")
|
|
||||||
model: Model instance for the strategy
|
|
||||||
device: Device to run the strategy on
|
|
||||||
**kwargs: Additional arguments passed to strategy constructor
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Strategy instance
|
|
||||||
"""
|
|
||||||
return super().create(train_type, model, device, **kwargs)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def available_strategies(cls) -> list:
|
|
||||||
"""Return list of registered strategy names."""
|
|
||||||
return cls.list_registered()
|
|
||||||
|
|
||||||
|
|
||||||
# ============== Strategy Classes ==============
|
# ============== Strategy Classes ==============
|
||||||
# All strategies are registered at class definition time using the decorator
|
# All strategies are registered at class definition time using the decorator
|
||||||
@@ -160,7 +192,13 @@ class SEQStrategy(BaseStrategy):
|
|||||||
Computes cross-entropy loss for next token prediction.
|
Computes cross-entropy loss for next token prediction.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, model, device, label_smoothing: float = 0.0, **kwargs):
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: Union[nn.Module, Callable[..., Dict[str, Tensor]]],
|
||||||
|
device: str,
|
||||||
|
label_smoothing: float = 0.0,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
super().__init__(model, device, **kwargs)
|
super().__init__(model, device, **kwargs)
|
||||||
self.label_smoothing = label_smoothing
|
self.label_smoothing = label_smoothing
|
||||||
|
|
||||||
@@ -185,21 +223,31 @@ class SFTStrategy(BaseStrategy):
|
|||||||
Applies cross-entropy loss only to tokens where loss_mask is True.
|
Applies cross-entropy loss only to tokens where loss_mask is True.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, model, device, label_smoothing: float = 0.0, **kwargs):
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: Union[nn.Module, Callable[..., Dict[str, Tensor]]],
|
||||||
|
device: str,
|
||||||
|
label_smoothing: float = 0.0,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
super().__init__(model, device, **kwargs)
|
super().__init__(model, device, **kwargs)
|
||||||
self.label_smoothing = label_smoothing
|
self.label_smoothing = label_smoothing
|
||||||
|
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
batch = move_to_device(batch, self.device)
|
batch = move_to_device(batch, self.device)
|
||||||
input_ids, target_ids, loss_mask = (
|
input_ids, target_ids, position_ids, loss_mask = (
|
||||||
batch["input_ids"],
|
batch["input_ids"],
|
||||||
batch["target_ids"],
|
batch["target_ids"],
|
||||||
|
batch["position_ids"],
|
||||||
batch["loss_mask"],
|
batch["loss_mask"],
|
||||||
)
|
)
|
||||||
|
|
||||||
ignore_index = -100
|
ignore_index = -100
|
||||||
logits = self.model(input_ids=input_ids)["logits"]
|
input_mask = make_doc_boundary_mask(position_ids)
|
||||||
target_ids = target_ids.masked_fill(loss_mask == 0, ignore_index)
|
target_ids = target_ids.masked_fill(~loss_mask, ignore_index)
|
||||||
|
logits = self.model(
|
||||||
|
input_ids=input_ids, position_ids=position_ids, input_mask=input_mask
|
||||||
|
)["logits"]
|
||||||
|
|
||||||
loss = F.cross_entropy(
|
loss = F.cross_entropy(
|
||||||
input=logits.flatten(0, 1).float(),
|
input=logits.flatten(0, 1).float(),
|
||||||
@@ -223,12 +271,13 @@ class DPOStrategy(BaseStrategy):
|
|||||||
self,
|
self,
|
||||||
model: nn.Module,
|
model: nn.Module,
|
||||||
device: str,
|
device: str,
|
||||||
|
ref_model: nn.Module,
|
||||||
beta: float = 0.1,
|
beta: float = 0.1,
|
||||||
reduction: str = "mean",
|
reduction: str = "sum",
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__(model, device, **kwargs)
|
super().__init__(model, device, **kwargs)
|
||||||
self.ref_model = create_ref_model(model)
|
self.ref_model = ref_model
|
||||||
self.beta = beta
|
self.beta = beta
|
||||||
self.reduction = reduction
|
self.reduction = reduction
|
||||||
|
|
||||||
@@ -238,13 +287,31 @@ class DPOStrategy(BaseStrategy):
|
|||||||
chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"]
|
chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"]
|
||||||
|
|
||||||
concat_ids = torch.cat([chosen_ids, rejected_ids], dim=0)
|
concat_ids = torch.cat([chosen_ids, rejected_ids], dim=0)
|
||||||
concat_mask = torch.cat([chosen_mask, rejected_mask], dim=0)
|
concat_loss_mask = torch.cat([chosen_mask, rejected_mask], dim=0)
|
||||||
|
|
||||||
log_pi = get_logprobs(self.model, concat_ids, concat_mask, self.reduction)
|
# Build full attention mask: key-padding + causal
|
||||||
|
key_pad = concat_ids.bool()[:, None, None, :] # [B*2, 1, 1, S]
|
||||||
|
S = key_pad.shape[-1]
|
||||||
|
causal = torch.tril(
|
||||||
|
torch.ones(S, S, dtype=torch.bool, device=concat_ids.device)
|
||||||
|
)[None, None, :, :] # [1, 1, S, S]
|
||||||
|
full_mask = key_pad & causal # [B*2, 1, S, S] — composed
|
||||||
|
|
||||||
|
log_pi = get_logprobs(
|
||||||
|
self.model,
|
||||||
|
concat_ids,
|
||||||
|
full_mask,
|
||||||
|
concat_loss_mask,
|
||||||
|
self.reduction,
|
||||||
|
)
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
log_ref = get_logprobs(
|
log_ref = get_logprobs(
|
||||||
self.ref_model, concat_ids, concat_mask, self.reduction
|
self.ref_model,
|
||||||
|
concat_ids,
|
||||||
|
full_mask,
|
||||||
|
concat_loss_mask,
|
||||||
|
self.reduction,
|
||||||
)
|
)
|
||||||
|
|
||||||
log_pi_chosen = log_pi[: chosen_ids.shape[0]]
|
log_pi_chosen = log_pi[: chosen_ids.shape[0]]
|
||||||
@@ -260,46 +327,77 @@ class DPOStrategy(BaseStrategy):
|
|||||||
|
|
||||||
return dpo_loss
|
return dpo_loss
|
||||||
|
|
||||||
|
def supports_online(self) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
|
||||||
|
"""Pick best/worst response per prompt by reward as chosen/rejected."""
|
||||||
|
rewards = result.rewards
|
||||||
|
responses = result.responses
|
||||||
|
masks = result.response_mask
|
||||||
|
best = rewards.argmax(dim=-1)
|
||||||
|
worst = rewards.argmin(dim=-1)
|
||||||
|
B = responses.shape[0]
|
||||||
|
idx = torch.arange(B, device=responses.device)
|
||||||
|
chosen = responses[idx, best]
|
||||||
|
chosen_mask = masks[idx, best].float()
|
||||||
|
rejected = responses[idx, worst]
|
||||||
|
rejected_mask = masks[idx, worst].float()
|
||||||
|
return {
|
||||||
|
"chosen": chosen,
|
||||||
|
"chosen_mask": chosen_mask,
|
||||||
|
"rejected": rejected,
|
||||||
|
"rejected_mask": rejected_mask,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@StrategyFactory.register("grpo")
|
@StrategyFactory.register("grpo")
|
||||||
class GRPOStrategy(BaseStrategy):
|
class GRPOStrategy(BaseStrategy):
|
||||||
"""Group Relative Policy Optimization strategy.
|
"""Group Relative Policy Optimization strategy.
|
||||||
|
|
||||||
On-policy GRPO following DeepSeek-R1: the policy model is updated while
|
Implements GRPO following DeepSeek-R1 with token-level PPO clipping.
|
||||||
a frozen ref_model stores the old-policy log-probs. ratio = exp(logπ_θ - logπ_ref),
|
Advantages are group-normalized from scalar per-response rewards and
|
||||||
clipped PPO objective. Call ``sync_ref_model()`` after each data-generation round.
|
broadcast across all response tokens. The loss is computed **only on
|
||||||
|
response tokens** — prompt tokens are masked out.
|
||||||
|
|
||||||
|
Three model roles are distinguished:
|
||||||
|
|
||||||
|
* **Policy** ``self.model`` — the model being trained.
|
||||||
|
* **Old policy** ``self.old_model`` — the behaviour policy that generated
|
||||||
|
the responses. Used for the importance sampling ratio
|
||||||
|
``ρ = π_θ / π_old``. Synced externally after each data-generation round.
|
||||||
|
* **Reference model** ``self.ref_model`` — a frozen copy of the initial
|
||||||
|
policy (typically the SFT checkpoint) used **only** for the KL
|
||||||
|
regularisation term. It is never updated during training.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
model: nn.Module,
|
model: nn.Module,
|
||||||
device: str,
|
device: str,
|
||||||
|
old_model: nn.Module,
|
||||||
|
ref_model: nn.Module,
|
||||||
clip_eps: float = 0.2,
|
clip_eps: float = 0.2,
|
||||||
kl_coef: float = 0.01,
|
kl_coef: float = 0.01,
|
||||||
group_size: int = 4,
|
group_size: int = 4,
|
||||||
reduction: str = "mean",
|
|
||||||
sync_interval: int = 200,
|
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__(model, device, **kwargs)
|
super().__init__(model, device, **kwargs)
|
||||||
self.ref_model = create_ref_model(model)
|
self.old_model = old_model
|
||||||
|
self.ref_model = ref_model
|
||||||
self.clip_eps = clip_eps
|
self.clip_eps = clip_eps
|
||||||
self.kl_coef = kl_coef
|
self.kl_coef = kl_coef
|
||||||
self.group_size = group_size
|
self.group_size = group_size
|
||||||
self.reduction = reduction
|
|
||||||
self.sync_interval = sync_interval
|
|
||||||
self._step = 0
|
|
||||||
|
|
||||||
def sync_ref_model(self):
|
def sync_old_model(self):
|
||||||
"""Copy current model weights to ref model."""
|
"""Copy current policy weights to old model."""
|
||||||
ref_state = self.model.state_dict()
|
state_dict = self.executor.unwrap_model(self.model)
|
||||||
self.ref_model.load_state_dict(ref_state)
|
if self.executor.use_distributed:
|
||||||
|
state_dict = broadcast_state_dict(state_dict)
|
||||||
|
if state_dict is not None:
|
||||||
|
self.old_model.load_state_dict(state_dict)
|
||||||
|
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
self._step += 1
|
|
||||||
if self._step % self.sync_interval == 0:
|
|
||||||
self.sync_ref_model()
|
|
||||||
|
|
||||||
batch = move_to_device(batch, self.device)
|
batch = move_to_device(batch, self.device)
|
||||||
prompts = batch["prompts"]
|
prompts = batch["prompts"]
|
||||||
responses = batch["responses"]
|
responses = batch["responses"]
|
||||||
@@ -310,33 +408,101 @@ class GRPOStrategy(BaseStrategy):
|
|||||||
responses_flat = responses.view(-1, response_len)
|
responses_flat = responses.view(-1, response_len)
|
||||||
masks_flat = masks.view(-1, response_len)
|
masks_flat = masks.view(-1, response_len)
|
||||||
prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1)
|
prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1)
|
||||||
|
prompt_mask = batch.get("prompt_mask")
|
||||||
|
if prompt_mask is None:
|
||||||
|
prompt_mask = prompts.ne(0)
|
||||||
|
prompt_mask_expanded = (
|
||||||
|
prompt_mask.unsqueeze(1).expand(-1, group_size, -1).flatten(0, 1)
|
||||||
|
)
|
||||||
|
prompt_len = prompt_expanded.size(1)
|
||||||
|
|
||||||
full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1)
|
full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1)
|
||||||
full_masks = torch.cat([torch.ones_like(prompt_expanded), masks_flat], dim=-1)
|
# Prompt tokens are masked out (0) so logprobs are computed only for
|
||||||
|
# response tokens. get_logprobs shifts the mask by one position, so
|
||||||
log_probs_policy = get_logprobs(
|
# the first response token's logprob (predicted from the last prompt
|
||||||
self.model, full_sequences, full_masks, self.reduction
|
# token) is correctly included.
|
||||||
|
full_masks = torch.cat(
|
||||||
|
[torch.zeros_like(prompt_expanded, dtype=torch.bool), masks_flat], dim=-1
|
||||||
)
|
)
|
||||||
log_probs_policy = log_probs_policy.view(batch_size, group_size)
|
|
||||||
|
|
||||||
|
# Build full attention mask: key-padding + causal
|
||||||
|
key_pad = torch.cat([prompt_mask_expanded, masks_flat.bool()], dim=-1)[
|
||||||
|
:, None, None, :
|
||||||
|
]
|
||||||
|
S = key_pad.shape[-1]
|
||||||
|
causal = torch.tril(
|
||||||
|
torch.ones(S, S, dtype=torch.bool, device=full_sequences.device)
|
||||||
|
)[None, None, :, :]
|
||||||
|
attn_mask = key_pad & causal
|
||||||
|
|
||||||
|
# get_logprobs returns [B*G, S-1] (S = prompt_len + response_len).
|
||||||
|
# Response token logprobs occupy the last ``response_len`` positions
|
||||||
|
# (the first response token is predicted from the last prompt token).
|
||||||
|
token_log_probs_policy = get_logprobs(
|
||||||
|
self.model, full_sequences, attn_mask, full_masks, "none"
|
||||||
|
)[:, prompt_len - 1 :]
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
log_probs_ref = get_logprobs(
|
token_log_probs_old = get_logprobs(
|
||||||
self.ref_model, full_sequences, full_masks, self.reduction
|
self.old_model, full_sequences, attn_mask, full_masks, "none"
|
||||||
)
|
)[:, prompt_len - 1 :]
|
||||||
log_probs_ref = log_probs_ref.view(batch_size, group_size)
|
token_log_probs_ref = get_logprobs(
|
||||||
|
self.ref_model, full_sequences, attn_mask, full_masks, "none"
|
||||||
|
)[:, prompt_len - 1 :]
|
||||||
|
|
||||||
eps = torch.finfo(log_probs_policy.dtype).eps
|
# Reshape to [B, G, response_len]
|
||||||
|
token_log_probs_policy = token_log_probs_policy.view(batch_size, group_size, -1)
|
||||||
|
token_log_probs_old = token_log_probs_old.view(batch_size, group_size, -1)
|
||||||
|
token_log_probs_ref = token_log_probs_ref.view(batch_size, group_size, -1)
|
||||||
|
token_masks = masks_flat.view(batch_size, group_size, -1).float()
|
||||||
|
|
||||||
|
# Group-normalized advantages from scalar per-response rewards.
|
||||||
|
eps = 1e-8
|
||||||
mean = rewards.mean(dim=-1, keepdim=True)
|
mean = rewards.mean(dim=-1, keepdim=True)
|
||||||
std = rewards.std(dim=-1, keepdim=True)
|
std = rewards.std(dim=-1, keepdim=True, unbiased=False)
|
||||||
advantages = (rewards - mean) / (std + eps)
|
advantages = (rewards - mean) / (std + eps)
|
||||||
|
# Broadcast scalar advantage to every response token: [B, G, 1]
|
||||||
|
advantages = advantages.unsqueeze(-1)
|
||||||
|
|
||||||
ratio = torch.exp(log_probs_policy - log_probs_ref)
|
# Token-level ratio (π_θ / π_old) and PPO clipping.
|
||||||
|
log_ratio = token_log_probs_policy - token_log_probs_old
|
||||||
|
ratio = torch.exp(log_ratio)
|
||||||
|
|
||||||
surr1 = ratio * advantages
|
surr1 = ratio * advantages
|
||||||
surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * advantages
|
surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * advantages
|
||||||
|
per_token_policy_loss = -torch.min(surr1, surr2)
|
||||||
|
token_count = token_masks.sum().clamp(min=1.0)
|
||||||
|
policy_loss = (per_token_policy_loss * token_masks).sum() / token_count
|
||||||
|
|
||||||
|
# KL penalty to frozen reference model with k1 estimator (non-negative):
|
||||||
|
# k1 = π_ref / π_θ - log(π_ref / π_θ) - 1, where π_ref / π_θ = exp(log_ref - log_policy).
|
||||||
|
log_ref_ratio = token_log_probs_ref - token_log_probs_policy
|
||||||
|
r = torch.exp(log_ref_ratio)
|
||||||
|
kl_per_token = r - torch.log(r + eps) - 1.0
|
||||||
|
kl_penalty = self.kl_coef * (kl_per_token * token_masks).sum() / token_count
|
||||||
|
|
||||||
policy_loss = -torch.min(surr1, surr2).mean()
|
|
||||||
kl_penalty = self.kl_coef * (log_probs_policy - log_probs_ref).square().mean()
|
|
||||||
total_loss = policy_loss + kl_penalty
|
total_loss = policy_loss + kl_penalty
|
||||||
|
|
||||||
return total_loss
|
return total_loss
|
||||||
|
|
||||||
|
def supports_online(self) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]:
|
||||||
|
return {
|
||||||
|
"prompts": result.prompts,
|
||||||
|
"prompt_mask": result.prompt_mask,
|
||||||
|
"responses": result.responses,
|
||||||
|
"masks": result.response_mask,
|
||||||
|
"rewards": result.rewards,
|
||||||
|
}
|
||||||
|
|
||||||
|
def _on_rollout_refresh(self):
|
||||||
|
"""Sync the behaviour policy whenever a fresh rollout arrives."""
|
||||||
|
self.sync_old_model()
|
||||||
|
|
||||||
|
|
||||||
|
# Factory aliases: online variants use the same strategy class; the
|
||||||
|
# ``RolloutRunner`` is injected by ``TrainContextBuilder`` to enable
|
||||||
|
# online mode, so no separate subclass is needed.
|
||||||
|
StrategyFactory.register("online_grpo")(GRPOStrategy)
|
||||||
|
StrategyFactory.register("online_dpo")(DPOStrategy)
|
||||||
|
|||||||
@@ -1,28 +1,32 @@
|
|||||||
import json
|
import json
|
||||||
|
import logging
|
||||||
import os
|
import os
|
||||||
|
import sys
|
||||||
import time
|
import time
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Callable, List, Optional, Protocol, runtime_checkable
|
from typing import IO, Callable, List, Optional, Protocol, runtime_checkable
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch.nn.utils import clip_grad_norm_
|
from torch.utils.checkpoint import checkpoint as torch_checkpoint
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.parallel import only_on_rank
|
from astrai.parallel import only_on_rank
|
||||||
|
from astrai.parallel.setup import get_current_device
|
||||||
from astrai.serialization import Checkpoint
|
from astrai.serialization import Checkpoint
|
||||||
from astrai.trainer.metric_util import (
|
from astrai.trainer.metric_util import (
|
||||||
ctx_get_grad_max,
|
|
||||||
ctx_get_grad_mean,
|
|
||||||
ctx_get_grad_min,
|
|
||||||
ctx_get_grad_nan_num,
|
|
||||||
ctx_get_grad_norm,
|
ctx_get_grad_norm,
|
||||||
ctx_get_grad_std,
|
ctx_get_grad_snr,
|
||||||
ctx_get_loss,
|
ctx_get_loss,
|
||||||
ctx_get_lr,
|
ctx_get_lr,
|
||||||
|
ctx_get_val_loss,
|
||||||
)
|
)
|
||||||
from astrai.trainer.train_context import TrainContext
|
from astrai.trainer.train_context import TrainContext
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@runtime_checkable
|
@runtime_checkable
|
||||||
class TrainCallback(Protocol):
|
class TrainCallback(Protocol):
|
||||||
@@ -42,18 +46,15 @@ class TrainCallback(Protocol):
|
|||||||
def on_epoch_end(self, context: TrainContext):
|
def on_epoch_end(self, context: TrainContext):
|
||||||
"""Called at the end of each epoch."""
|
"""Called at the end of each epoch."""
|
||||||
|
|
||||||
def on_step_begin(self, context: TrainContext):
|
|
||||||
"""Called at the beginning of each step."""
|
|
||||||
|
|
||||||
def on_step_end(self, context: TrainContext):
|
|
||||||
"""Called at the end of each step."""
|
|
||||||
|
|
||||||
def on_batch_begin(self, context: TrainContext):
|
def on_batch_begin(self, context: TrainContext):
|
||||||
"""Called at the beginning of each batch."""
|
"""Called at the beginning of each batch."""
|
||||||
|
|
||||||
def on_batch_end(self, context: TrainContext):
|
def on_batch_end(self, context: TrainContext):
|
||||||
"""Called at the end of each batch."""
|
"""Called at the end of each batch."""
|
||||||
|
|
||||||
|
def on_optimizer_step(self, context: TrainContext):
|
||||||
|
"""Called on every optimizer step (sync step only)."""
|
||||||
|
|
||||||
def on_error(self, context: TrainContext):
|
def on_error(self, context: TrainContext):
|
||||||
"""Called when an error occurs during training."""
|
"""Called when an error occurs during training."""
|
||||||
|
|
||||||
@@ -79,9 +80,45 @@ class GradientClippingCallback(TrainCallback):
|
|||||||
def __init__(self, max_grad_norm: float):
|
def __init__(self, max_grad_norm: float):
|
||||||
self.max_grad_norm = max_grad_norm
|
self.max_grad_norm = max_grad_norm
|
||||||
|
|
||||||
def on_step_end(self, context: TrainContext):
|
def on_optimizer_step(self, context: TrainContext):
|
||||||
_ = context
|
context.grad_norm = context.executor.clip_grad_norm(
|
||||||
clip_grad_norm_(context.model.parameters(), self.max_grad_norm)
|
context.model, self.max_grad_norm
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@CallbackFactory.register("gradient_checkpointing")
|
||||||
|
class GradientCheckpointingCallback(TrainCallback):
|
||||||
|
"""
|
||||||
|
Activation checkpointing callback — trades compute for memory
|
||||||
|
by recomputing specified module activations during the backward pass.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
modules: Module types to apply checkpointing to.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, modules: Optional[List[type]] = None):
|
||||||
|
self.modules = tuple(modules) if modules else ()
|
||||||
|
|
||||||
|
def _enable(self, module: nn.Module):
|
||||||
|
if self.modules and isinstance(module, self.modules):
|
||||||
|
fn = module.forward
|
||||||
|
module._original_forward = fn
|
||||||
|
module.forward = lambda *a, **kw: torch_checkpoint(
|
||||||
|
fn, *a, use_reentrant=False, **kw
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _disable(module: nn.Module):
|
||||||
|
if hasattr(module, "_original_forward"):
|
||||||
|
module.forward = module._original_forward
|
||||||
|
del module._original_forward
|
||||||
|
|
||||||
|
def on_train_begin(self, context: TrainContext):
|
||||||
|
context.model.apply(self._enable)
|
||||||
|
logger.info("Gradient checkpointing enabled")
|
||||||
|
|
||||||
|
def on_train_end(self, context: TrainContext):
|
||||||
|
context.model.apply(self._disable)
|
||||||
|
|
||||||
|
|
||||||
@CallbackFactory.register("checkpoint")
|
@CallbackFactory.register("checkpoint")
|
||||||
@@ -90,54 +127,65 @@ class CheckpointCallback(TrainCallback):
|
|||||||
Checkpoint callback for trainer.
|
Checkpoint callback for trainer.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
extra_keys = ("optimizer", "scheduler")
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
save_dir: str,
|
save_dir: str,
|
||||||
interval: int,
|
interval: int,
|
||||||
weight_only: bool = False,
|
weight_only: bool = False,
|
||||||
state_dict_fn: Optional[Callable[[nn.Module], dict]] = None,
|
|
||||||
save_extra_fn: Optional[Callable[["TrainContext"], dict]] = None,
|
save_extra_fn: Optional[Callable[["TrainContext"], dict]] = None,
|
||||||
):
|
):
|
||||||
self.save_dir = save_dir
|
self.save_dir = save_dir
|
||||||
self.interval = interval
|
self.interval = interval
|
||||||
self.weight_only = weight_only
|
self.weight_only = weight_only
|
||||||
self.state_dict_fn = state_dict_fn
|
self.save_extra_fn = save_extra_fn or CheckpointCallback.save_extra
|
||||||
self.save_extra_fn = save_extra_fn
|
self.last_ckpt_step = None
|
||||||
self.last_ckpt_iter = 0
|
|
||||||
|
def on_train_begin(self, context: TrainContext):
|
||||||
|
self.last_ckpt_step = context.optimizer_step
|
||||||
|
|
||||||
@only_on_rank(0)
|
|
||||||
def _save_checkpoint(self, context: TrainContext):
|
def _save_checkpoint(self, context: TrainContext):
|
||||||
save_path = os.path.join(
|
self.last_ckpt_step = context.optimizer_step
|
||||||
self.save_dir, f"epoch_{context.epoch}_iter_{context.iteration}"
|
|
||||||
)
|
|
||||||
state_dict = (
|
|
||||||
self.state_dict_fn(context.model)
|
|
||||||
if self.state_dict_fn
|
|
||||||
else context.model.state_dict()
|
|
||||||
)
|
|
||||||
|
|
||||||
extra = self.save_extra_fn(context) if self.save_extra_fn else None
|
with context.executor.checkpoint_context(context.model) as state_dict:
|
||||||
context.checkpoint = Checkpoint(
|
if state_dict is not None:
|
||||||
state_dict=state_dict,
|
save_path = os.path.join(
|
||||||
epoch=context.epoch,
|
self.save_dir,
|
||||||
iteration=context.iteration,
|
f"epoch_{context.epoch}_step_{context.optimizer_step}",
|
||||||
extra=extra,
|
)
|
||||||
)
|
extra = self.save_extra_fn(context)
|
||||||
|
meta = context.config.to_dict()
|
||||||
context.checkpoint.save(save_path)
|
context.checkpoint = Checkpoint(
|
||||||
self.last_ckpt_iter = context.iteration
|
state_dict=state_dict,
|
||||||
|
epoch=context.epoch,
|
||||||
|
consumed_samples=context.consumed_samples,
|
||||||
|
config=context.model_config,
|
||||||
|
extra=extra,
|
||||||
|
meta=meta,
|
||||||
|
)
|
||||||
|
context.checkpoint.save(save_path)
|
||||||
|
|
||||||
def on_batch_end(self, context: TrainContext):
|
def on_batch_end(self, context: TrainContext):
|
||||||
if context.iteration - self.last_ckpt_iter >= self.interval:
|
if context.optimizer_step - self.last_ckpt_step >= self.interval:
|
||||||
self._save_checkpoint(context)
|
self._save_checkpoint(context)
|
||||||
|
|
||||||
def on_train_end(self, context: TrainContext):
|
def on_train_end(self, context: TrainContext):
|
||||||
if context.iteration != self.last_ckpt_iter:
|
if context.optimizer_step != self.last_ckpt_step:
|
||||||
self._save_checkpoint(context)
|
self._save_checkpoint(context)
|
||||||
|
|
||||||
def on_error(self, context: TrainContext):
|
def on_error(self, context: TrainContext):
|
||||||
self._save_checkpoint(context)
|
self._save_checkpoint(context)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def save_extra(context: TrainContext) -> dict:
|
||||||
|
extra = {}
|
||||||
|
for name in CheckpointCallback.extra_keys:
|
||||||
|
obj = getattr(context, name, None)
|
||||||
|
if obj:
|
||||||
|
extra[name] = obj.state_dict()
|
||||||
|
return extra
|
||||||
|
|
||||||
|
|
||||||
@CallbackFactory.register("progress_bar")
|
@CallbackFactory.register("progress_bar")
|
||||||
class ProgressBarCallback(TrainCallback):
|
class ProgressBarCallback(TrainCallback):
|
||||||
@@ -145,26 +193,36 @@ class ProgressBarCallback(TrainCallback):
|
|||||||
Progress bar callback for trainer.
|
Progress bar callback for trainer.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, num_epoch: int):
|
def __init__(
|
||||||
|
self, num_epoch: int, log_interval: int = 100, file: Optional[IO[str]] = None
|
||||||
|
):
|
||||||
self.num_epoch = num_epoch
|
self.num_epoch = num_epoch
|
||||||
|
self.log_interval = log_interval
|
||||||
|
self.file = file
|
||||||
self.progress_bar: tqdm = None
|
self.progress_bar: tqdm = None
|
||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
def on_epoch_begin(self, context: TrainContext):
|
def on_epoch_begin(self, context: TrainContext):
|
||||||
|
total_steps = len(context.dataloader) // context.executor.grad_accum_steps
|
||||||
self.progress_bar = tqdm(
|
self.progress_bar = tqdm(
|
||||||
context.dataloader,
|
total=total_steps,
|
||||||
desc=f"Epoch {context.epoch + 1}/{self.num_epoch}",
|
desc=f"Epoch {context.epoch + 1}/{self.num_epoch}",
|
||||||
dynamic_ncols=True,
|
dynamic_ncols=True,
|
||||||
|
file=self.file or sys.stdout,
|
||||||
)
|
)
|
||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
def on_batch_end(self, context: TrainContext):
|
def on_optimizer_step(self, context: TrainContext):
|
||||||
self.progress_bar.set_postfix(
|
postfix = {
|
||||||
{
|
"step": f"{context.optimizer_step:d}",
|
||||||
"loss": f"{context.loss:.4f}",
|
"loss": f"{context.loss:.4f}",
|
||||||
"lr": f"{context.optimizer.param_groups[-1]['lr']:.2e}",
|
"lr": f"{context.optimizer.param_groups[-1]['lr']:.2e}",
|
||||||
}
|
}
|
||||||
)
|
if context.grad_norm is not None:
|
||||||
|
postfix["grad_norm"] = f"{context.grad_norm:.2f}"
|
||||||
|
if context.val_loss is not None:
|
||||||
|
postfix["val_loss"] = f"{context.val_loss:.4f}"
|
||||||
|
self.progress_bar.set_postfix(postfix)
|
||||||
self.progress_bar.update(1)
|
self.progress_bar.update(1)
|
||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
@@ -174,68 +232,116 @@ class ProgressBarCallback(TrainCallback):
|
|||||||
self.progress_bar.close()
|
self.progress_bar.close()
|
||||||
|
|
||||||
|
|
||||||
@CallbackFactory.register("metric_logger")
|
@CallbackFactory.register("metric")
|
||||||
class MetricLoggerCallback(TrainCallback):
|
class MetricCallback(TrainCallback):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
log_dir: str,
|
ckpt_dir: str,
|
||||||
save_interval: int,
|
save_interval: int,
|
||||||
log_interval: int = 10,
|
|
||||||
metrics: List[str] = None,
|
metrics: List[str] = None,
|
||||||
|
val_step: int = 0,
|
||||||
):
|
):
|
||||||
self.last_log_iter = 0
|
self.last_log_flush_step = None
|
||||||
self.save_interval = save_interval
|
self.save_interval = save_interval
|
||||||
self.log_interval = log_interval
|
|
||||||
self.metrics = metrics or ["loss", "lr"]
|
self.metrics = metrics or ["loss", "lr"]
|
||||||
|
self.val_step = val_step
|
||||||
|
self._next_val_step = 0
|
||||||
|
|
||||||
self.log_dir = Path(log_dir) if log_dir else Path.cwd() / "logs"
|
self.ckpt_dir = Path(ckpt_dir) if ckpt_dir else Path.cwd() / "checkpoint"
|
||||||
self.log_dir.mkdir(parents=True, exist_ok=True)
|
|
||||||
|
|
||||||
self.log_cache = []
|
self.log_cache = []
|
||||||
|
|
||||||
self._metric_funcs = {
|
self._metric_funcs = {
|
||||||
"loss": ctx_get_loss,
|
"loss": ctx_get_loss,
|
||||||
"lr": ctx_get_lr,
|
"lr": ctx_get_lr,
|
||||||
|
"val_loss": ctx_get_val_loss,
|
||||||
"grad_norm": ctx_get_grad_norm,
|
"grad_norm": ctx_get_grad_norm,
|
||||||
"grad_std": ctx_get_grad_std,
|
"grad_snr": ctx_get_grad_snr,
|
||||||
"grad_max": ctx_get_grad_max,
|
|
||||||
"grad_min": ctx_get_grad_min,
|
|
||||||
"grad_mean": ctx_get_grad_mean,
|
|
||||||
"grad_nan_num": ctx_get_grad_nan_num,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
def _get_log_data(self, context: TrainContext):
|
def _metrics(self, context: TrainContext, names):
|
||||||
return {
|
return {
|
||||||
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S"),
|
m: self._metric_funcs[m](context)
|
||||||
"epoch": context.epoch,
|
for m in names
|
||||||
"iter": context.iteration,
|
if self._metric_funcs[m](context) is not None
|
||||||
**{m: self._metric_funcs[m](context) for m in self.metrics},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
def _add_log(self, log_data):
|
def _append(self, event_type: str, context: TrainContext, **extra):
|
||||||
self.log_cache.append(log_data)
|
entry = {
|
||||||
|
"type": event_type,
|
||||||
|
"timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"),
|
||||||
|
"epoch": context.epoch,
|
||||||
|
"step": context.optimizer_step,
|
||||||
|
"consumed_samples": context.consumed_samples,
|
||||||
|
**extra,
|
||||||
|
}
|
||||||
|
self.log_cache.append(entry)
|
||||||
|
|
||||||
|
def _run_validation(self, context: TrainContext) -> float:
|
||||||
|
context.model.eval()
|
||||||
|
|
||||||
|
total_loss = 0.0
|
||||||
|
num_batches = 0
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
for batch in context.val_dataloader:
|
||||||
|
loss = context.strategy(batch)
|
||||||
|
total_loss += loss.item()
|
||||||
|
num_batches += 1
|
||||||
|
|
||||||
|
if context.world_size > 1 and dist.is_initialized():
|
||||||
|
stats = torch.tensor(
|
||||||
|
[total_loss, float(num_batches)], device=get_current_device()
|
||||||
|
)
|
||||||
|
dist.all_reduce(stats, op=dist.ReduceOp.SUM)
|
||||||
|
avg_loss = (stats[0] / stats[1]).item()
|
||||||
|
else:
|
||||||
|
avg_loss = total_loss / max(num_batches, 1)
|
||||||
|
|
||||||
|
context.model.train()
|
||||||
|
return avg_loss
|
||||||
|
|
||||||
|
def on_train_begin(self, context: TrainContext):
|
||||||
|
self.last_log_flush_step = context.optimizer_step
|
||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
def _save_log(self, epoch, iter):
|
def _flush(self, epoch, step):
|
||||||
log_file = self.log_dir / f"epoch_{epoch}_iter_{iter}_metric.jsonl"
|
log_file = self.ckpt_dir / f"epoch_{epoch}_step_{step}" / "metric.jsonl"
|
||||||
|
log_file.parent.mkdir(parents=True, exist_ok=True)
|
||||||
with open(log_file, "w") as f:
|
with open(log_file, "w") as f:
|
||||||
for log in self.log_cache:
|
for log in self.log_cache:
|
||||||
f.write(json.dumps(log) + "\n")
|
f.write(json.dumps(log) + "\n")
|
||||||
|
|
||||||
def on_batch_end(self, context):
|
def on_optimizer_step(self, context):
|
||||||
if context.iteration % self.log_interval == 0:
|
context.grad_snr_tracker.update(context.model)
|
||||||
log_data = self._get_log_data(context)
|
|
||||||
self._add_log(log_data)
|
|
||||||
|
|
||||||
if context.iteration - self.last_log_iter >= self.save_interval:
|
if (
|
||||||
self._save_log(context.epoch, context.iteration)
|
context.val_dataloader is not None
|
||||||
self.last_log_iter = context.iteration
|
and self.val_step > 0
|
||||||
|
and context.optimizer_step >= self._next_val_step
|
||||||
|
):
|
||||||
|
context.val_loss = self._run_validation(context)
|
||||||
|
self._next_val_step = context.optimizer_step + self.val_step
|
||||||
|
self._append("validation", context, val_loss=context.val_loss)
|
||||||
|
|
||||||
|
step_metrics = [m for m in self.metrics if m != "val_loss"]
|
||||||
|
self._append("step", context, **self._metrics(context, step_metrics))
|
||||||
|
|
||||||
|
if context.optimizer_step - self.last_log_flush_step >= self.save_interval:
|
||||||
|
self._flush(context.epoch, context.optimizer_step)
|
||||||
|
self.last_log_flush_step = context.optimizer_step
|
||||||
|
|
||||||
|
def on_epoch_end(self, context):
|
||||||
|
self._append("epoch", context)
|
||||||
|
|
||||||
def on_train_end(self, context):
|
def on_train_end(self, context):
|
||||||
if context.iteration != self.last_log_iter:
|
if (
|
||||||
self._save_log(context.epoch, context.iteration)
|
self.last_log_flush_step is None
|
||||||
|
or context.optimizer_step != self.last_log_flush_step
|
||||||
|
):
|
||||||
|
self._flush(context.epoch, context.optimizer_step)
|
||||||
|
self.last_log_flush_step = context.optimizer_step
|
||||||
|
|
||||||
def on_error(self, context):
|
def on_error(self, context):
|
||||||
self._save_log(context.epoch, context.iteration)
|
self._flush(context.epoch, context.optimizer_step)
|
||||||
|
|||||||
+234
-44
@@ -1,101 +1,291 @@
|
|||||||
|
import logging
|
||||||
|
import threading
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Callable, Optional, Self
|
from pathlib import Path
|
||||||
|
from typing import Any, Dict, Optional, Self
|
||||||
|
|
||||||
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch.optim import Optimizer
|
from torch.utils.data import DataLoader, random_split
|
||||||
from torch.optim.lr_scheduler import LRScheduler
|
|
||||||
from torch.utils.data import DataLoader
|
|
||||||
|
|
||||||
from astrai.config.train_config import TrainConfig
|
from astrai.config.train_config import TrainConfig
|
||||||
from astrai.dataset import ResumableDistributedSampler
|
from astrai.dataset import RDSampler
|
||||||
|
from astrai.inference.core.scheduler import InferenceScheduler
|
||||||
|
from astrai.model.components.lora import inject_lora
|
||||||
|
from astrai.parallel.executor import BaseExecutor, ExecutorFactory, create_ref_model
|
||||||
from astrai.parallel.setup import get_current_device, get_rank, get_world_size
|
from astrai.parallel.setup import get_current_device, get_rank, get_world_size
|
||||||
from astrai.serialization import Checkpoint
|
from astrai.protocols import OptimizerProtocol, SchedulerProtocol
|
||||||
|
from astrai.serialization import Checkpoint, load_json
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
from astrai.trainer.metric_util import GradSNRTracker
|
||||||
|
from astrai.trainer.rollout import RolloutGenerator, RolloutRunner
|
||||||
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class TrainContext:
|
class TrainContext:
|
||||||
model: nn.Module = field(default=None)
|
model: nn.Module = field(default=None)
|
||||||
strategy: BaseStrategy = field(default=None)
|
strategy: BaseStrategy = field(default=None)
|
||||||
dataloader: DataLoader = field(default=None)
|
dataloader: DataLoader = field(default=None)
|
||||||
optimizer: Optimizer = field(default=None)
|
optimizer: OptimizerProtocol = field(default=None)
|
||||||
scheduler: LRScheduler = field(default=None)
|
scheduler: SchedulerProtocol = field(default=None)
|
||||||
checkpoint: Checkpoint = field(default=None)
|
checkpoint: Checkpoint = field(default=None)
|
||||||
|
config: TrainConfig = field(default=None)
|
||||||
|
model_config: dict = field(default_factory=dict)
|
||||||
|
executor: BaseExecutor = field(default=None)
|
||||||
epoch: int = field(default=0)
|
epoch: int = field(default=0)
|
||||||
iteration: int = field(default=0)
|
consumed_samples: int = field(default=0)
|
||||||
loss: float = field(default=0.0)
|
loss: float = field(default=0.0)
|
||||||
|
grad_norm: Optional[float] = field(default=None)
|
||||||
|
grad_snr_tracker: GradSNRTracker = field(default_factory=GradSNRTracker)
|
||||||
|
val_dataloader: Optional[DataLoader] = field(default=None)
|
||||||
|
val_loss: Optional[float] = field(default=None)
|
||||||
|
|
||||||
world_size: int = field(default=1)
|
world_size: int = field(default=1)
|
||||||
rank: int = field(default=0)
|
rank: int = field(default=0)
|
||||||
kwargs: dict = field(default_factory=dict)
|
kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
_stop_event: threading.Event = field(default_factory=threading.Event)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def stop_requested(self) -> bool:
|
||||||
|
return self._stop_event.is_set()
|
||||||
|
|
||||||
|
def request_stop(self) -> None:
|
||||||
|
self._stop_event.set()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def optimizer_step(self) -> int:
|
||||||
|
return self.consumed_samples // (
|
||||||
|
self.config.batch_per_device
|
||||||
|
* self.world_size
|
||||||
|
* self.config.grad_accum_steps
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TrainContextBuilder:
|
class TrainContextBuilder:
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: TrainConfig,
|
config: TrainConfig,
|
||||||
load_extra_fn: Optional[Callable[[dict, "TrainContext"], None]] = None,
|
|
||||||
):
|
):
|
||||||
self.config = config
|
self.config = config
|
||||||
self._checkpoint: Optional[Checkpoint] = None
|
self._param_path: Optional[str] = None
|
||||||
self._load_extra_fn = load_extra_fn
|
self._resume: bool = False
|
||||||
|
|
||||||
def with_checkpoint(self, checkpoint: Optional[Checkpoint]) -> Self:
|
def with_param_path(self, param_path: Optional[str], resume: bool = False) -> Self:
|
||||||
self._checkpoint = checkpoint
|
self._param_path = param_path
|
||||||
|
self._resume = resume
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def build(self) -> TrainContext:
|
def build(self) -> TrainContext:
|
||||||
context = TrainContext(
|
cfg = self.config
|
||||||
model=self.config.model,
|
device = get_current_device()
|
||||||
world_size=get_world_size(),
|
|
||||||
rank=get_rank(),
|
executor = ExecutorFactory.create(
|
||||||
|
cfg.parallel_mode,
|
||||||
|
grad_accum_steps=cfg.grad_accum_steps,
|
||||||
|
**cfg.executor_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
device = get_current_device()
|
model_config = {}
|
||||||
context.model = context.model.to(device=device)
|
if self._param_path:
|
||||||
|
config_path = Path(self._param_path) / "config.json"
|
||||||
|
if config_path.exists():
|
||||||
|
model_config = load_json(config_path)
|
||||||
|
|
||||||
if self.config.nprocs > 1 and self.config.parallel_wrapper:
|
preloaded_state_dict = None
|
||||||
context.model = self.config.parallel_wrapper(context.model)
|
preloaded_epoch = cfg.start_epoch
|
||||||
|
preloaded_consumed = cfg.start_samples * get_world_size()
|
||||||
|
preloaded_checkpoint = None
|
||||||
|
if self._param_path:
|
||||||
|
checkpoint = Checkpoint.load_any(self._param_path)
|
||||||
|
if checkpoint is not None:
|
||||||
|
preloaded_state_dict = checkpoint.state_dict
|
||||||
|
if checkpoint.config:
|
||||||
|
model_config = checkpoint.config
|
||||||
|
if self._resume:
|
||||||
|
preloaded_epoch = checkpoint.epoch
|
||||||
|
per_step = (
|
||||||
|
cfg.batch_per_device * get_world_size() * cfg.grad_accum_steps
|
||||||
|
)
|
||||||
|
preloaded_consumed = (
|
||||||
|
checkpoint.consumed_samples // per_step
|
||||||
|
) * per_step
|
||||||
|
preloaded_checkpoint = checkpoint
|
||||||
|
|
||||||
if self._checkpoint is not None:
|
if not model_config and hasattr(cfg.model_fn(), "config"):
|
||||||
context.epoch = max(self._checkpoint.epoch, self.config.start_epoch)
|
model_config = cfg.model_fn().config.to_dict()
|
||||||
context.iteration = max(self._checkpoint.iteration, self.config.start_batch)
|
|
||||||
context.model.load_state_dict(self._checkpoint.state_dict)
|
def _before_wrap(m):
|
||||||
context.checkpoint = self._checkpoint
|
m = m.to(device=device)
|
||||||
else:
|
if cfg.lora is not None:
|
||||||
context.checkpoint = Checkpoint(
|
inject_lora(
|
||||||
state_dict=context.model.state_dict(),
|
m,
|
||||||
|
r=cfg.lora.r,
|
||||||
|
alpha=cfg.lora.alpha,
|
||||||
|
target_modules=set(cfg.lora.target_modules),
|
||||||
|
)
|
||||||
|
if preloaded_state_dict is not None:
|
||||||
|
m.load_state_dict(preloaded_state_dict, strict=False)
|
||||||
|
return m
|
||||||
|
|
||||||
|
def _after_wrap(m):
|
||||||
|
if cfg.compile_mode is not None:
|
||||||
|
logger.info("torch.compile enabled (mode=%s)", cfg.compile_mode)
|
||||||
|
m = torch.compile(m, mode=cfg.compile_mode)
|
||||||
|
return m
|
||||||
|
|
||||||
|
context = TrainContext(
|
||||||
|
world_size=get_world_size(),
|
||||||
|
rank=get_rank(),
|
||||||
|
config=cfg,
|
||||||
|
model_config=model_config,
|
||||||
|
executor=executor,
|
||||||
|
epoch=preloaded_epoch,
|
||||||
|
consumed_samples=preloaded_consumed,
|
||||||
|
checkpoint=preloaded_checkpoint,
|
||||||
|
)
|
||||||
|
|
||||||
|
context.model, context.optimizer, context.scheduler = executor.prepare(
|
||||||
|
cfg.model_fn,
|
||||||
|
cfg.optimizer_fn,
|
||||||
|
cfg.scheduler_fn,
|
||||||
|
before_wrap=_before_wrap,
|
||||||
|
after_wrap=_after_wrap,
|
||||||
|
)
|
||||||
|
|
||||||
|
train_dataset = cfg.dataset
|
||||||
|
val_dataset = cfg.val_dataset
|
||||||
|
|
||||||
|
if val_dataset is None and cfg.val_split is not None:
|
||||||
|
n_total = len(cfg.dataset)
|
||||||
|
n_val = max(1, int(n_total * cfg.val_split))
|
||||||
|
n_train = n_total - n_val
|
||||||
|
generator = torch.Generator().manual_seed(cfg.random_seed)
|
||||||
|
train_dataset, val_dataset = random_split(
|
||||||
|
cfg.dataset, [n_train, n_val], generator=generator
|
||||||
)
|
)
|
||||||
|
|
||||||
context.optimizer = self.config.optimizer_fn(context.model)
|
sampler_offset = context.consumed_samples // context.world_size
|
||||||
context.scheduler = self.config.scheduler_fn(context.optimizer)
|
|
||||||
|
|
||||||
if self._checkpoint and self._checkpoint.extra and self._load_extra_fn:
|
if self._resume and sampler_offset > 0:
|
||||||
self._load_extra_fn(self._checkpoint.extra, context)
|
offset = context.world_size - 1
|
||||||
|
num_samples_per_replica = (
|
||||||
|
len(train_dataset) + offset
|
||||||
|
) // context.world_size
|
||||||
|
if num_samples_per_replica > 0:
|
||||||
|
context.epoch = sampler_offset // num_samples_per_replica
|
||||||
|
|
||||||
cfg = self.config
|
sampler = RDSampler(
|
||||||
sampler_offset = context.iteration * cfg.batch_size
|
data_source=train_dataset,
|
||||||
sampler = ResumableDistributedSampler(
|
|
||||||
data_source=cfg.dataset,
|
|
||||||
start_epoch=context.epoch,
|
start_epoch=context.epoch,
|
||||||
start_iter=sampler_offset,
|
start_iter=sampler_offset,
|
||||||
seed=cfg.random_seed,
|
seed=cfg.random_seed,
|
||||||
)
|
)
|
||||||
context.dataloader = DataLoader(
|
context.dataloader = DataLoader(
|
||||||
cfg.dataset,
|
train_dataset,
|
||||||
batch_size=cfg.batch_size,
|
batch_size=cfg.batch_per_device,
|
||||||
sampler=sampler,
|
sampler=sampler,
|
||||||
num_workers=cfg.num_workers,
|
num_workers=cfg.num_workers,
|
||||||
pin_memory=cfg.pin_memory,
|
pin_memory=cfg.pin_memory,
|
||||||
prefetch_factor=cfg.prefetch_factor,
|
prefetch_factor=cfg.prefetch_factor,
|
||||||
|
collate_fn=cfg.collate_fn,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if val_dataset is not None:
|
||||||
|
val_sampler = RDSampler(
|
||||||
|
data_source=val_dataset,
|
||||||
|
start_epoch=0,
|
||||||
|
start_iter=0,
|
||||||
|
seed=cfg.random_seed,
|
||||||
|
shuffle=False,
|
||||||
|
)
|
||||||
|
context.val_dataloader = DataLoader(
|
||||||
|
val_dataset,
|
||||||
|
batch_size=cfg.batch_per_device,
|
||||||
|
sampler=val_sampler,
|
||||||
|
num_workers=cfg.num_workers,
|
||||||
|
pin_memory=cfg.pin_memory,
|
||||||
|
prefetch_factor=cfg.prefetch_factor,
|
||||||
|
collate_fn=cfg.collate_fn,
|
||||||
|
)
|
||||||
|
|
||||||
|
if context.checkpoint and context.checkpoint.extra:
|
||||||
|
extra = context.checkpoint.extra
|
||||||
|
for name in ("optimizer", "scheduler"):
|
||||||
|
if name in extra:
|
||||||
|
obj = getattr(context, name, None)
|
||||||
|
if obj is not None:
|
||||||
|
obj.load_state_dict(extra[name])
|
||||||
|
|
||||||
|
strategy_kwargs = dict(cfg.extra_kwargs)
|
||||||
|
|
||||||
|
needs_ref = cfg.strategy in (
|
||||||
|
"dpo",
|
||||||
|
"grpo",
|
||||||
|
"online_grpo",
|
||||||
|
"online_dpo",
|
||||||
|
)
|
||||||
|
needs_old = cfg.strategy in ("grpo", "online_grpo")
|
||||||
|
|
||||||
|
if needs_ref:
|
||||||
|
strategy_kwargs["ref_model"] = create_ref_model(
|
||||||
|
cfg.model_fn, executor=executor, model=context.model, device=device
|
||||||
|
)
|
||||||
|
|
||||||
|
if needs_old:
|
||||||
|
strategy_kwargs["old_model"] = create_ref_model(
|
||||||
|
cfg.model_fn, executor=executor, model=context.model, device=device
|
||||||
|
)
|
||||||
|
|
||||||
context.strategy = StrategyFactory.create(
|
context.strategy = StrategyFactory.create(
|
||||||
|
cfg.strategy,
|
||||||
model=context.model,
|
model=context.model,
|
||||||
train_type=self.config.strategy,
|
|
||||||
device=device,
|
device=device,
|
||||||
**self.config.extra_kwargs,
|
executor=executor,
|
||||||
|
**strategy_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Enable online rollout when the train_type is an ``online_*`` variant.
|
||||||
|
is_online = cfg.strategy.startswith("online_")
|
||||||
|
if is_online:
|
||||||
|
if not context.strategy.supports_online():
|
||||||
|
raise ValueError(
|
||||||
|
f"Strategy '{cfg.strategy}' does not support online rollout"
|
||||||
|
)
|
||||||
|
if cfg.reward_model_fn is None:
|
||||||
|
raise ValueError("reward_model_fn is required for online RL strategies")
|
||||||
|
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(self._param_path)
|
||||||
|
reward_model = cfg.reward_model_fn()
|
||||||
|
|
||||||
|
group_size = strategy_kwargs.get("group_size", 1)
|
||||||
|
rollout_batch_size = group_size * max(1, cfg.batch_per_device)
|
||||||
|
max_seq_len = getattr(context.model.config, "max_position_embeddings", None)
|
||||||
|
|
||||||
|
scheduler = InferenceScheduler(
|
||||||
|
model=context.model,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
max_batch_size=rollout_batch_size,
|
||||||
|
max_seq_len=max_seq_len,
|
||||||
|
)
|
||||||
|
|
||||||
|
generator = RolloutGenerator(
|
||||||
|
scheduler=scheduler,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
max_tokens=cfg.rollout_max_tokens,
|
||||||
|
group_size=group_size,
|
||||||
|
temperature=cfg.rollout_temperature,
|
||||||
|
top_k=cfg.rollout_top_k,
|
||||||
|
top_p=cfg.rollout_top_p,
|
||||||
|
)
|
||||||
|
runner = RolloutRunner(
|
||||||
|
generator=generator,
|
||||||
|
reward_model=reward_model,
|
||||||
|
rollout_interval=cfg.rollout_interval,
|
||||||
|
)
|
||||||
|
context.strategy.set_rollout_runner(runner)
|
||||||
|
|
||||||
return context
|
return context
|
||||||
|
|||||||
+74
-40
@@ -1,10 +1,14 @@
|
|||||||
import logging
|
import logging
|
||||||
from itertools import batched
|
|
||||||
from typing import List, Optional
|
from typing import List, Optional
|
||||||
|
|
||||||
|
import torch.distributed as dist
|
||||||
|
|
||||||
from astrai.config import TrainConfig
|
from astrai.config import TrainConfig
|
||||||
from astrai.parallel.setup import spawn_parallel_fn
|
from astrai.parallel.setup import spawn_parallel_fn
|
||||||
from astrai.serialization import Checkpoint
|
from astrai.signal_handler import (
|
||||||
|
register_signal_handlers,
|
||||||
|
unregister_signal_handlers,
|
||||||
|
)
|
||||||
from astrai.trainer.train_callback import (
|
from astrai.trainer.train_callback import (
|
||||||
CallbackFactory,
|
CallbackFactory,
|
||||||
TrainCallback,
|
TrainCallback,
|
||||||
@@ -26,17 +30,27 @@ class Trainer:
|
|||||||
|
|
||||||
def _get_default_callbacks(self) -> List[TrainCallback]:
|
def _get_default_callbacks(self) -> List[TrainCallback]:
|
||||||
cfg = self.train_config
|
cfg = self.train_config
|
||||||
return [
|
callbacks = [
|
||||||
|
CallbackFactory.create(
|
||||||
|
"gradient_checkpointing",
|
||||||
|
modules=cfg.gradient_checkpointing_modules,
|
||||||
|
),
|
||||||
|
CallbackFactory.create(
|
||||||
|
"checkpoint",
|
||||||
|
cfg.ckpt_dir,
|
||||||
|
cfg.ckpt_interval,
|
||||||
|
),
|
||||||
|
CallbackFactory.create(
|
||||||
|
"metric",
|
||||||
|
ckpt_dir=cfg.ckpt_dir,
|
||||||
|
save_interval=cfg.ckpt_interval,
|
||||||
|
metrics=cfg.metrics,
|
||||||
|
val_step=cfg.val_step,
|
||||||
|
),
|
||||||
CallbackFactory.create("progress_bar", cfg.n_epoch),
|
CallbackFactory.create("progress_bar", cfg.n_epoch),
|
||||||
CallbackFactory.create("checkpoint", cfg.ckpt_dir, cfg.ckpt_interval),
|
|
||||||
CallbackFactory.create("metric_logger", cfg.ckpt_dir, cfg.ckpt_interval),
|
|
||||||
CallbackFactory.create("gradient_clipping", cfg.max_grad_norm),
|
CallbackFactory.create("gradient_clipping", cfg.max_grad_norm),
|
||||||
]
|
]
|
||||||
|
return callbacks
|
||||||
def _build_context(self, checkpoint: Optional[Checkpoint]) -> TrainContext:
|
|
||||||
return (
|
|
||||||
TrainContextBuilder(self.train_config).with_checkpoint(checkpoint).build()
|
|
||||||
)
|
|
||||||
|
|
||||||
def _call_callbacks(self, method_name: str, context: TrainContext):
|
def _call_callbacks(self, method_name: str, context: TrainContext):
|
||||||
for callback in self.callbacks:
|
for callback in self.callbacks:
|
||||||
@@ -44,56 +58,76 @@ class Trainer:
|
|||||||
if method:
|
if method:
|
||||||
method(context)
|
method(context)
|
||||||
|
|
||||||
def train(self, checkpoint: Optional[Checkpoint] = None):
|
def _trainer_loop(self, param_path: Optional[str] = None, resume: bool = False):
|
||||||
config = self.train_config
|
context = (
|
||||||
spawn_parallel_fn(
|
TrainContextBuilder(self.train_config)
|
||||||
self._train_impl,
|
.with_param_path(param_path, resume=resume)
|
||||||
backend=config.backend,
|
.build()
|
||||||
world_size=config.nprocs,
|
|
||||||
master_addr=config.master_addr,
|
|
||||||
master_port=config.master_port,
|
|
||||||
device_type=config.device_type,
|
|
||||||
checkpoint=checkpoint,
|
|
||||||
)
|
)
|
||||||
|
register_signal_handlers(context)
|
||||||
def _train_impl(self, checkpoint: Optional[Checkpoint] = None) -> Checkpoint:
|
executor = context.executor
|
||||||
context = self._build_context(checkpoint)
|
|
||||||
self._call_callbacks("on_train_begin", context)
|
self._call_callbacks("on_train_begin", context)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
context.model.train()
|
context.model.train()
|
||||||
accumulation_steps = max(self.train_config.accumulation_steps, 1)
|
|
||||||
|
|
||||||
for epoch in range(context.epoch, self.train_config.n_epoch):
|
for epoch in range(context.epoch, context.config.n_epoch):
|
||||||
|
if context.stop_requested:
|
||||||
|
break
|
||||||
context.epoch = epoch
|
context.epoch = epoch
|
||||||
self._call_callbacks("on_epoch_begin", context)
|
self._call_callbacks("on_epoch_begin", context)
|
||||||
|
|
||||||
for steps in batched(context.dataloader, accumulation_steps):
|
for batch in context.dataloader:
|
||||||
self._call_callbacks("on_step_begin", context)
|
if context.stop_requested:
|
||||||
|
break
|
||||||
step_batch_nums = len(steps)
|
with executor.accumulate(context.model):
|
||||||
for batch in steps:
|
|
||||||
self._call_callbacks("on_batch_begin", context)
|
self._call_callbacks("on_batch_begin", context)
|
||||||
loss = context.strategy(batch)
|
loss = context.strategy(batch)
|
||||||
context.loss = loss.item()
|
context.loss = loss.item()
|
||||||
context.iteration += 1
|
stand_loss = loss / executor.grad_accum_steps
|
||||||
|
executor.backward(stand_loss)
|
||||||
stand_loss = loss / step_batch_nums
|
context.consumed_samples += (
|
||||||
stand_loss.backward()
|
context.config.batch_per_device * context.world_size
|
||||||
|
)
|
||||||
self._call_callbacks("on_batch_end", context)
|
self._call_callbacks("on_batch_end", context)
|
||||||
|
|
||||||
self._call_callbacks("on_step_end", context)
|
if executor.sync_gradients:
|
||||||
context.optimizer.step()
|
self._call_callbacks("on_optimizer_step", context)
|
||||||
context.optimizer.zero_grad()
|
context.optimizer.step()
|
||||||
|
context.strategy.on_optimizer_step()
|
||||||
|
context.optimizer.zero_grad()
|
||||||
|
|
||||||
if context.scheduler:
|
if context.scheduler:
|
||||||
context.scheduler.step()
|
context.scheduler.step()
|
||||||
|
|
||||||
self._call_callbacks("on_epoch_end", context)
|
self._call_callbacks("on_epoch_end", context)
|
||||||
|
|
||||||
|
if context.stop_requested:
|
||||||
|
logger.warning(
|
||||||
|
"Training interrupted by signal, saving emergency checkpoint..."
|
||||||
|
)
|
||||||
|
self._call_callbacks("on_error", context)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Training failed: {str(e)}", exc_info=True)
|
logger.error("Training failed: %s", str(e), exc_info=True)
|
||||||
self._call_callbacks("on_error", context)
|
self._call_callbacks("on_error", context)
|
||||||
raise
|
raise
|
||||||
finally:
|
finally:
|
||||||
self._call_callbacks("on_train_end", context)
|
self._call_callbacks("on_train_end", context)
|
||||||
|
if executor.use_distributed and dist.is_initialized():
|
||||||
|
dist.barrier()
|
||||||
|
unregister_signal_handlers()
|
||||||
|
|
||||||
|
def train(self, param_path: Optional[str] = None, resume: bool = False):
|
||||||
|
cfg = self.train_config
|
||||||
|
spawn_parallel_fn(
|
||||||
|
self._trainer_loop,
|
||||||
|
backend=cfg.backend,
|
||||||
|
world_size=cfg.nprocs,
|
||||||
|
master_addr=cfg.master_addr,
|
||||||
|
master_port=cfg.master_port,
|
||||||
|
device_type=cfg.device_type,
|
||||||
|
start_method=cfg.start_method,
|
||||||
|
param_path=param_path,
|
||||||
|
resume=resume,
|
||||||
|
)
|
||||||
|
|||||||
@@ -0,0 +1,2 @@
|
|||||||
|
# Source directory for CUDA kernels — build-time only.
|
||||||
|
# Compiled .so files live in astrAI/_ext/.
|
||||||
@@ -0,0 +1,76 @@
|
|||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
def cuda_toolkit_version() -> tuple[int, int] | None:
|
||||||
|
"""Return ``(major, minor)`` of the nvcc on PATH, or ``None``.
|
||||||
|
|
||||||
|
Used by ``setup.py`` to detect nvcc/torch CUDA version mismatches
|
||||||
|
(e.g. nvcc 13.0 with a cu128 torch wheel) which cause cryptic ABI errors.
|
||||||
|
"""
|
||||||
|
import shutil
|
||||||
|
import subprocess
|
||||||
|
|
||||||
|
nvcc = shutil.which("nvcc")
|
||||||
|
if nvcc is None:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
out = subprocess.check_output(
|
||||||
|
[nvcc, "--version"], stderr=subprocess.STDOUT, text=True
|
||||||
|
)
|
||||||
|
for line in out.splitlines():
|
||||||
|
if "release" in line:
|
||||||
|
ver = line.split("release")[1].split(",")[0].strip()
|
||||||
|
major, minor = ver.split(".")
|
||||||
|
return (int(major), int(minor))
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _arch_flags() -> list[str]:
|
||||||
|
import torch
|
||||||
|
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
cap = torch.cuda.get_device_capability()
|
||||||
|
else:
|
||||||
|
cap = (8, 0)
|
||||||
|
ver = f"{cap[0]}{cap[1]}"
|
||||||
|
flags = [f"-gencode=arch=compute_{ver},code=sm_{ver}"]
|
||||||
|
# tensor-core mma path (mma.sync.m16n8k16.bf16) requires sm_80+; decide the
|
||||||
|
# kernel dispatch at build time via this define rather than at runtime.
|
||||||
|
if cap[0] < 8:
|
||||||
|
flags.append("-DASTRAI_NO_MMA")
|
||||||
|
return flags
|
||||||
|
|
||||||
|
|
||||||
|
_kernels_dir = Path("csrc/kernels")
|
||||||
|
REGISTRY: dict[str, dict] = {}
|
||||||
|
|
||||||
|
CXX_FLAGS = ["-O3", "-funroll-loops"]
|
||||||
|
NVCC_FLAGS = [
|
||||||
|
"-O3",
|
||||||
|
"--expt-relaxed-constexpr",
|
||||||
|
"--use_fast_math",
|
||||||
|
"--ptxas-options=-O3,-v",
|
||||||
|
"--extra-device-vectorization",
|
||||||
|
"--threads=16",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def register(name: str, sources: list[str] | None = None, **kwargs):
|
||||||
|
if sources is None:
|
||||||
|
sources = [str(_kernels_dir / f"{name}.cu")]
|
||||||
|
REGISTRY[name] = {
|
||||||
|
"sources": sources,
|
||||||
|
"cxx_flags": [*CXX_FLAGS],
|
||||||
|
"nvcc_flags": [*NVCC_FLAGS, *_arch_flags()],
|
||||||
|
"extra_link_args": kwargs.pop("extra_link_args", []),
|
||||||
|
**kwargs,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
register("attn_decode")
|
||||||
|
register("attn_prefill")
|
||||||
|
register("attn_paged_decode")
|
||||||
|
register("attn_paged_prefill")
|
||||||
|
register("rotary_emb")
|
||||||
@@ -0,0 +1,99 @@
|
|||||||
|
#pragma once
|
||||||
|
|
||||||
|
// Tensor layout for Q/K/V tensors passed to attention kernels.
|
||||||
|
// Internally, kernels always operate on BHLD [batch, n_heads, seq_len, head_dim].
|
||||||
|
// When the caller passes BLHD, dims 1 and 2 are transposed at entry.
|
||||||
|
enum TensorLayout : int {
|
||||||
|
BHLD = 0, // [batch, n_heads, seq_len, head_dim]
|
||||||
|
BLHD = 1, // [batch, seq_len, n_heads, head_dim]
|
||||||
|
};
|
||||||
|
|
||||||
|
|
||||||
|
template<typename T, typename AT = float>
|
||||||
|
struct AttentionParams {
|
||||||
|
int batch;
|
||||||
|
int q_head;
|
||||||
|
int kv_head;
|
||||||
|
int q_len;
|
||||||
|
int kv_len;
|
||||||
|
int head_dim;
|
||||||
|
int use_mask;
|
||||||
|
int causal_offset; // -1 = non-causal; >=0 = absolute position of first Q token
|
||||||
|
int num_splits;
|
||||||
|
float scale;
|
||||||
|
|
||||||
|
// Q strides (element offsets for each dim — layout-agnostic)
|
||||||
|
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
|
||||||
|
// KV strides (K and V share the same layout — only base pointers differ)
|
||||||
|
int kv_stride_b, kv_stride_h, kv_stride_l, kv_stride_d;
|
||||||
|
|
||||||
|
// Mask: 2D [batch, kv_len], 3D [batch, q_len, kv_len],
|
||||||
|
// or 4D [batch, n_heads, q_len, kv_len] (head dim broadcasts when stride=0)
|
||||||
|
int mask_b_stride; // batch stride
|
||||||
|
int mask_h_stride; // head stride (0 = broadcast across heads)
|
||||||
|
int mask_q_stride; // q stride (0 = all q rows share)
|
||||||
|
|
||||||
|
const T* __restrict__ q;
|
||||||
|
const T* __restrict__ k;
|
||||||
|
const T* __restrict__ v;
|
||||||
|
const bool* __restrict__ mask;
|
||||||
|
|
||||||
|
T* __restrict__ o;
|
||||||
|
AT* __restrict__ o_part;
|
||||||
|
AT* __restrict__ ml_part;
|
||||||
|
};
|
||||||
|
|
||||||
|
// ---- PagedAttentionParams ----
|
||||||
|
// SGLang-style indirect params over a shared KV pool.
|
||||||
|
// k_cache/v_cache: [size, kv_head, head_dim] (bare buffers, no gather).
|
||||||
|
// req_to_token: [num_reqs, max_context_len] token -> slot.
|
||||||
|
// req_pool_indices:[batch] rows of the current batch into req_to_token.
|
||||||
|
// kv_indptr: [batch+1] prefix sum of per-request seq_lens (device).
|
||||||
|
// qo_indptr: [batch+1] prefix sum of per-request q_len (prefill) or
|
||||||
|
// nullptr for decode (q_len == 1 everywhere).
|
||||||
|
template<typename T, typename AT = float>
|
||||||
|
struct PagedAttentionParams {
|
||||||
|
int batch;
|
||||||
|
int q_head;
|
||||||
|
int kv_head;
|
||||||
|
int head_dim;
|
||||||
|
int num_splits;
|
||||||
|
int use_mask;
|
||||||
|
int causal_offset; // -1 = non-causal; >=0 = causal (per-request offset
|
||||||
|
// computed inside kernel from kv_indptr/qo_indptr)
|
||||||
|
float scale;
|
||||||
|
|
||||||
|
// Q: [total_q, q_head, head_dim] (3D flattened — no batch dim).
|
||||||
|
// For decode total_q == batch (q_len=1 per request).
|
||||||
|
// For prefill total_q == qo_indptr[batch].
|
||||||
|
int q_stride_l, q_stride_h, q_stride_d;
|
||||||
|
|
||||||
|
// Q: [total_q, q_head, head_dim]
|
||||||
|
const T* __restrict__ q;
|
||||||
|
|
||||||
|
// Flat KV pool: [size, kv_head, head_dim]
|
||||||
|
const T* __restrict__ k_cache;
|
||||||
|
const T* __restrict__ v_cache;
|
||||||
|
|
||||||
|
// Indexing
|
||||||
|
const int64_t* __restrict__ req_to_token; // [num_reqs, max_context_len]
|
||||||
|
const int64_t* __restrict__ req_pool_indices; // [batch]
|
||||||
|
const int* __restrict__ kv_indptr; // [batch+1]
|
||||||
|
const int* __restrict__ qo_indptr; // [batch+1] or nullptr (decode)
|
||||||
|
int max_context_len; // req_to_token stride (dim 1)
|
||||||
|
int max_seq_len; // max per-request seq_len (host-side, for split computation)
|
||||||
|
int total_q; // total Q tokens across all requests (host-side, for grid)
|
||||||
|
int max_q_len; // max per-request q_len (host-side, for prefill grid)
|
||||||
|
|
||||||
|
// Mask: [batch, max_seq_len] (decode) or [batch, 1, q_len, kv_len]
|
||||||
|
// (prefill, optional). mask_h_stride/mask_q_stride are 0 when those
|
||||||
|
// dims are size 1 (broadcast).
|
||||||
|
int mask_b_stride;
|
||||||
|
int mask_h_stride;
|
||||||
|
int mask_q_stride;
|
||||||
|
const bool* __restrict__ mask;
|
||||||
|
|
||||||
|
T* __restrict__ o;
|
||||||
|
AT* __restrict__ o_part;
|
||||||
|
AT* __restrict__ ml_part;
|
||||||
|
};
|
||||||
@@ -0,0 +1,38 @@
|
|||||||
|
#include "attn_dispatchers.cuh"
|
||||||
|
#include "attn_entry_utils.cuh"
|
||||||
|
|
||||||
|
torch::Tensor attn_decode(
|
||||||
|
torch::Tensor q,
|
||||||
|
torch::Tensor k,
|
||||||
|
torch::Tensor v,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
int64_t causal_offset,
|
||||||
|
double scale,
|
||||||
|
int64_t layout
|
||||||
|
) {
|
||||||
|
AttentionParams<bf16> p;
|
||||||
|
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
|
||||||
|
TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1");
|
||||||
|
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
||||||
|
|
||||||
|
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
||||||
|
auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O;
|
||||||
|
p.o = (bf16*)O_view.data_ptr();
|
||||||
|
|
||||||
|
alloc_split_partials(p);
|
||||||
|
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p);
|
||||||
|
C10_CUDA_CHECK(cudaGetLastError());
|
||||||
|
return O;
|
||||||
|
}
|
||||||
|
|
||||||
|
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||||
|
m.def("attn_decode", &attn_decode,
|
||||||
|
py::arg("q"),
|
||||||
|
py::arg("k"),
|
||||||
|
py::arg("v"),
|
||||||
|
py::arg("mask") = py::none(),
|
||||||
|
py::arg("causal_offset") = -1,
|
||||||
|
py::arg("scale") = 0.0,
|
||||||
|
py::arg("layout") = (int64_t)BHLD,
|
||||||
|
"GQA decode (tensor-core head-packing on sm_80+, scalar fallback)");
|
||||||
|
}
|
||||||
@@ -0,0 +1,129 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include <float.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
#include "attn_warp_utils.cuh"
|
||||||
|
constexpr int DC_CHUNK = 64;
|
||||||
|
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
__global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
||||||
|
int batch = blockIdx.x / p.kv_head;
|
||||||
|
int kv_head = blockIdx.x % p.kv_head;
|
||||||
|
int split = blockIdx.z;
|
||||||
|
int group_size = blockDim.y;
|
||||||
|
int q_head = kv_head * group_size + threadIdx.y;
|
||||||
|
int lane = threadIdx.x;
|
||||||
|
int hd_per_thread = p.head_dim / 32;
|
||||||
|
|
||||||
|
// Q: [batch, q_head, q_len=1, head_dim] — stride-based
|
||||||
|
float q_reg[8];
|
||||||
|
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
||||||
|
+ lane * hd_per_thread * p.q_stride_d;
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
|
||||||
|
|
||||||
|
// KV: [batch, kv_head, kv_len, head_dim] — stride-based base
|
||||||
|
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||||
|
int mask_base = batch * p.mask_b_stride + q_head * p.mask_h_stride;
|
||||||
|
|
||||||
|
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
|
||||||
|
|
||||||
|
extern __shared__ __align__(16) bf16 k_smem[];
|
||||||
|
|
||||||
|
// Split-KV: each split processes a contiguous subset of chunks
|
||||||
|
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||||
|
int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits;
|
||||||
|
int ch_begin = split * chunks_per_split;
|
||||||
|
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
|
||||||
|
|
||||||
|
for (int ci = ch_begin; ci < ch_end; ci++) {
|
||||||
|
int chunk_start = ci * DC_CHUNK;
|
||||||
|
int this_chunk = min(DC_CHUNK, p.kv_len - chunk_start);
|
||||||
|
|
||||||
|
// Load K into shared memory (gather from strided global)
|
||||||
|
int total = this_chunk * p.head_dim;
|
||||||
|
for (int i = threadIdx.y * 32 + lane; i < total;
|
||||||
|
i += blockDim.x * blockDim.y) {
|
||||||
|
int s = i / p.head_dim;
|
||||||
|
int d_dim = i % p.head_dim;
|
||||||
|
int kv_idx = chunk_start + s;
|
||||||
|
int g_off = kv_base + kv_idx * p.kv_stride_l + d_dim * p.kv_stride_d;
|
||||||
|
k_smem[i] = p.k[g_off];
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
|
||||||
|
for (int s = 0; s < this_chunk; s++) {
|
||||||
|
float partial = 0.0f;
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
partial += q_reg[i] * __bfloat162float(
|
||||||
|
k_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
||||||
|
partial = warp_reduce_sum(partial) * p.scale;
|
||||||
|
|
||||||
|
int kv_idx = chunk_start + s;
|
||||||
|
if constexpr (HasMask) {
|
||||||
|
if (!p.mask[mask_base + kv_idx])
|
||||||
|
partial = -FLT_MAX;
|
||||||
|
}
|
||||||
|
if constexpr (IsCausal) {
|
||||||
|
if (kv_idx > p.causal_offset)
|
||||||
|
partial = -FLT_MAX;
|
||||||
|
}
|
||||||
|
|
||||||
|
float new_m = fmaxf(m, partial);
|
||||||
|
float alpha = __expf(m - new_m);
|
||||||
|
float beta = __expf(partial - new_m);
|
||||||
|
d = d * alpha + beta;
|
||||||
|
|
||||||
|
int v_off = kv_base + kv_idx * p.kv_stride_l
|
||||||
|
+ lane * hd_per_thread * p.kv_stride_d;
|
||||||
|
for (int i = 0; i < hd_per_thread; i++)
|
||||||
|
acc_reg[i] = fmaf(acc_reg[i], alpha,
|
||||||
|
__bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta);
|
||||||
|
m = new_m;
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- write UN-normalised partials for this split ----
|
||||||
|
size_t bh = (size_t)batch * p.q_head + q_head;
|
||||||
|
size_t slot = bh * MAX_SPLITS + split;
|
||||||
|
int d0 = lane * hd_per_thread;
|
||||||
|
for (int i = 0; i < hd_per_thread; i++) {
|
||||||
|
int dd = d0 + i;
|
||||||
|
p.o_part[slot * p.head_dim + dd] = acc_reg[i];
|
||||||
|
}
|
||||||
|
if (lane == 0) {
|
||||||
|
p.ml_part[slot * 2] = m;
|
||||||
|
p.ml_part[slot * 2 + 1] = d;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
__global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
|
||||||
|
int bh = blockIdx.x;
|
||||||
|
int d = threadIdx.x;
|
||||||
|
if (d >= p.head_dim) return;
|
||||||
|
|
||||||
|
int batch = bh / p.q_head;
|
||||||
|
int q_head = bh % p.q_head;
|
||||||
|
|
||||||
|
size_t split_base = (size_t)bh * MAX_SPLITS;
|
||||||
|
const float* mlp = p.ml_part + split_base * 2;
|
||||||
|
const float* op = p.o_part + split_base * p.head_dim;
|
||||||
|
|
||||||
|
float m = -FLT_MAX, l = 0.0f, acc = 0.0f;
|
||||||
|
for (int s = 0; s < p.num_splits; s++) {
|
||||||
|
float mi = mlp[s * 2];
|
||||||
|
if (mi <= -FLT_MAX) continue;
|
||||||
|
float li = mlp[s * 2 + 1];
|
||||||
|
float nm = fmaxf(m, mi);
|
||||||
|
float corr = __expf(m - nm);
|
||||||
|
float e = __expf(mi - nm);
|
||||||
|
acc = fmaf(acc, corr, op[s * p.head_dim + d] * e);
|
||||||
|
l = fmaf(l, corr, li * e);
|
||||||
|
m = nm;
|
||||||
|
}
|
||||||
|
|
||||||
|
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
||||||
|
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d;
|
||||||
|
p.o[o_off] = __float2bfloat16(acc * inv);
|
||||||
|
}
|
||||||
@@ -0,0 +1,169 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <cfloat>
|
||||||
|
#include <cuda_bf16.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
#include "attn_mma_utils.cuh"
|
||||||
|
#include "attn_warp_utils.cuh"
|
||||||
|
|
||||||
|
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing.
|
||||||
|
// Decode has q_len == 1, so we pack G = q_head/kv_head query heads into the
|
||||||
|
// M=16 rows of mma.sync.m16n8k16, turning G independent GEMVs into a single
|
||||||
|
// GEMM that reuses each loaded K/V tile across all G heads.
|
||||||
|
//
|
||||||
|
// IsCausal and HasMask are compile-time bools — no runtime branch in the
|
||||||
|
// inner compute loop.
|
||||||
|
//
|
||||||
|
// Traits = KernelTraits<HEAD_DIM, BC=32, WARPS=1, STAGES=<2 or 1>>.
|
||||||
|
template <typename Traits, bool IsCausal, bool HasMask>
|
||||||
|
__global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
||||||
|
const int lane = threadIdx.x;
|
||||||
|
const int gid = lane >> 2;
|
||||||
|
const int tid4 = lane & 3;
|
||||||
|
|
||||||
|
const int pass = blockIdx.x / p.kv_head;
|
||||||
|
const int kv_head = blockIdx.x % p.kv_head;
|
||||||
|
const int batch = blockIdx.y;
|
||||||
|
const int split = blockIdx.z;
|
||||||
|
|
||||||
|
constexpr int MAX_G = 16;
|
||||||
|
const int G_total = p.q_head / p.kv_head;
|
||||||
|
const int g_begin = pass * MAX_G;
|
||||||
|
const int G = min(MAX_G, G_total - g_begin);
|
||||||
|
const int q_head0 = kv_head * G_total + g_begin;
|
||||||
|
|
||||||
|
// Double-buffered shared memory for K/V (no sQ needed)
|
||||||
|
__shared__ __align__(16) bf16 sK[Traits::STAGES * Traits::BC * Traits::LD];
|
||||||
|
__shared__ __align__(16) bf16 sV[Traits::STAGES * Traits::BC * Traits::LD];
|
||||||
|
|
||||||
|
// Load Q directly from global into mma A-operand registers.
|
||||||
|
// stride_row = p.q_stride_h for decode (q_len=1).
|
||||||
|
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
|
||||||
|
const int qra = gid;
|
||||||
|
const int qrb = gid + 8;
|
||||||
|
const bool va = qra < G, vb = qrb < G;
|
||||||
|
unsigned Qa[Traits::KD][4];
|
||||||
|
load_q_mma_frags<Traits::KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
|
||||||
|
qra, qrb, va, vb, tid4, Qa);
|
||||||
|
|
||||||
|
float Oacc[Traits::DN8][4];
|
||||||
|
#pragma unroll
|
||||||
|
for (int j = 0; j < Traits::DN8; j++)
|
||||||
|
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||||
|
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||||
|
|
||||||
|
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||||
|
const int tiles_total = (p.kv_len + Traits::BC - 1) / Traits::BC;
|
||||||
|
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
|
||||||
|
const int ti_begin = split * tiles_per_split;
|
||||||
|
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
|
||||||
|
|
||||||
|
// ---- Load tile lambda: predicated cp.async ----
|
||||||
|
auto load_tile = [&](int ti, int buf) {
|
||||||
|
int kv0 = ti * Traits::BC;
|
||||||
|
bf16* dK = sK + buf * Traits::BC * Traits::LD;
|
||||||
|
bf16* dV = sV + buf * Traits::BC * Traits::LD;
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = lane * Traits::VEC; i < Traits::TOTAL;
|
||||||
|
i += Traits::NUM_THREADS * Traits::VEC) {
|
||||||
|
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
||||||
|
int kc = kv0 + r;
|
||||||
|
bool valid = kc < p.kv_len;
|
||||||
|
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
||||||
|
int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
|
||||||
|
cp_async_16_pred(&dK[off], &p.k[g_off], valid);
|
||||||
|
cp_async_16_pred(&dV[off], &p.v[g_off], valid);
|
||||||
|
}
|
||||||
|
cp_async_commit();
|
||||||
|
};
|
||||||
|
|
||||||
|
// ---- Multi-stage cp.async pipeline ----
|
||||||
|
// Prologue loads STAGES tiles; each loop iteration waits only for the
|
||||||
|
// oldest outstanding group (wait_group<STAGES-1>) so the STAGES-1 newer
|
||||||
|
// tile loads stay in flight and overlap with the current tile's compute.
|
||||||
|
constexpr int STAGES = Traits::STAGES;
|
||||||
|
const int ntiles = ti_end - ti_begin;
|
||||||
|
|
||||||
|
auto process_tile = [&](int it, int buf) {
|
||||||
|
const bf16* bK = sK + buf * Traits::BC * Traits::LD;
|
||||||
|
const bf16* bV = sV + buf * Traits::BC * Traits::LD;
|
||||||
|
int kv0 = (ti_begin + it) * Traits::BC;
|
||||||
|
|
||||||
|
float Sacc[Traits::NC8][4];
|
||||||
|
mma_compute_scores<Traits>(Qa, bK, lane, Sacc);
|
||||||
|
|
||||||
|
#pragma unroll
|
||||||
|
for (int n8 = 0; n8 < Traits::NC8; n8++)
|
||||||
|
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||||
|
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||||
|
|
||||||
|
// Decode: q_len=1, so qrow0=qrow1=0
|
||||||
|
int maxc = IsCausal ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
||||||
|
mma_softmax_tile<Traits, HasMask>(kv0, maxc, maxc,
|
||||||
|
0, 0,
|
||||||
|
p.mask_b_stride, 0, 0,
|
||||||
|
batch, 0,
|
||||||
|
p.mask,
|
||||||
|
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||||
|
|
||||||
|
mma_pv_accumulate<Traits>(Sacc, bV, lane, Oacc);
|
||||||
|
};
|
||||||
|
|
||||||
|
if (ntiles >= STAGES) {
|
||||||
|
#pragma unroll
|
||||||
|
for (int i = 0; i < STAGES; i++)
|
||||||
|
load_tile(ti_begin + i, i);
|
||||||
|
|
||||||
|
for (int it = 0; it < ntiles; it++) {
|
||||||
|
cp_async_wait_group<STAGES - 1>();
|
||||||
|
__syncwarp();
|
||||||
|
process_tile(it, it & (STAGES - 1));
|
||||||
|
__syncwarp();
|
||||||
|
if (it + STAGES < ntiles)
|
||||||
|
load_tile(ti_begin + it + STAGES, (it + STAGES) & (STAGES - 1));
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// Fewer tiles than stages: load all, wait for all, process.
|
||||||
|
for (int i = 0; i < ntiles; i++)
|
||||||
|
load_tile(ti_begin + i, i);
|
||||||
|
cp_async_wait_group<0>();
|
||||||
|
__syncwarp();
|
||||||
|
for (int it = 0; it < ntiles; it++)
|
||||||
|
process_tile(it, it);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- write UN-normalised partials for this split ----
|
||||||
|
auto split_slot = [&](int h) -> size_t {
|
||||||
|
size_t bh = (size_t)batch * p.q_head + h;
|
||||||
|
return bh * MAX_SPLITS + split;
|
||||||
|
};
|
||||||
|
#pragma unroll
|
||||||
|
for (int dn8 = 0; dn8 < Traits::DN8; dn8++) {
|
||||||
|
int d = dn8 * 8 + 2 * tid4;
|
||||||
|
int r0 = gid, r1 = gid + 8;
|
||||||
|
if (r0 < G) {
|
||||||
|
int h = q_head0 + r0;
|
||||||
|
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
|
||||||
|
op[d] = Oacc[dn8][0];
|
||||||
|
op[d + 1] = Oacc[dn8][1];
|
||||||
|
}
|
||||||
|
if (r1 < G) {
|
||||||
|
int h = q_head0 + r1;
|
||||||
|
float* op = p.o_part + split_slot(h) * Traits::HEAD_DIM;
|
||||||
|
op[d] = Oacc[dn8][2];
|
||||||
|
op[d + 1] = Oacc[dn8][3];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (tid4 == 0) {
|
||||||
|
int r0 = gid, r1 = gid + 8;
|
||||||
|
if (r0 < G) {
|
||||||
|
int h = q_head0 + r0;
|
||||||
|
float* mp = p.ml_part + split_slot(h) * 2;
|
||||||
|
mp[0] = m0; mp[1] = l0;
|
||||||
|
}
|
||||||
|
if (r1 < G) {
|
||||||
|
int h = q_head0 + r1;
|
||||||
|
float* mp = p.ml_part + split_slot(h) * 2;
|
||||||
|
mp[0] = m1; mp[1] = l1;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,232 @@
|
|||||||
|
#pragma once
|
||||||
|
// Shared attention dispatchers — used by both production .cu and test .cu.
|
||||||
|
// No torch dependency; pure CUDA.
|
||||||
|
|
||||||
|
#include <cuda_runtime.h>
|
||||||
|
#include <algorithm>
|
||||||
|
#include "attn_warp_utils.cuh"
|
||||||
|
#include "attn_prefill_split_q.cuh"
|
||||||
|
#include "attn_decode_split_kv.cuh"
|
||||||
|
#include "attn_paged_decode_split_kv.cuh"
|
||||||
|
#include "attn_paged_prefill_split_q.cuh"
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
#include "attn_prefill_split_q_mma.cuh"
|
||||||
|
#include "attn_decode_split_kv_mma.cuh"
|
||||||
|
#include "attn_paged_decode_split_kv_mma.cuh"
|
||||||
|
#include "attn_paged_prefill_split_q_mma.cuh"
|
||||||
|
#endif
|
||||||
|
|
||||||
|
// Cached SM count — cudaDeviceGetAttribute is a host-side call that was
|
||||||
|
// invoked on every decode/paged-decode launch. Cache per-device so multi-GPU
|
||||||
|
// setups with heterogeneous GPUs still get the right count, while the common
|
||||||
|
// single-GPU path hits the cache after the first call.
|
||||||
|
inline int get_sm_count() {
|
||||||
|
int dev = 0;
|
||||||
|
cudaGetDevice(&dev);
|
||||||
|
static int cached_dev = -1;
|
||||||
|
static int cached_count = 0;
|
||||||
|
if (dev != cached_dev) {
|
||||||
|
cudaDeviceGetAttribute(&cached_count, cudaDevAttrMultiProcessorCount, dev);
|
||||||
|
cached_dev = dev;
|
||||||
|
}
|
||||||
|
return cached_count;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Split-KV: compute number of splits to fill all SMs for small-batch decode.
|
||||||
|
// Caps splits so each split processes at least `min_tiles_per_split` tiles,
|
||||||
|
// avoiding excessive loop/prologue overhead when tiles are small.
|
||||||
|
inline int compute_num_splits(int base_blocks, int tiles_total,
|
||||||
|
int min_tiles_per_split = 1) {
|
||||||
|
int sm_count = get_sm_count();
|
||||||
|
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
|
||||||
|
int max_by_work = tiles_total / min_tiles_per_split;
|
||||||
|
return std::max(1, std::min(n, std::min(max_by_work, MAX_SPLITS)));
|
||||||
|
}
|
||||||
|
|
||||||
|
// Dispatch IsCausal × HasMask — eliminates the duplicated 4-way if/else
|
||||||
|
// ladder that appeared in each dispatch_* function. FN must be a function
|
||||||
|
// template <int HEAD_DIM, bool IsCausal, bool HasMask>; HEAD_DIM is forwarded
|
||||||
|
// as the first template argument so callers only spell it once.
|
||||||
|
//
|
||||||
|
// Usage: DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_mma, HEAD_DIM, p, group_size);
|
||||||
|
#define DISPATCH_CAUSAL_MASK(is_causal, has_mask, FN, HEAD_DIM, ...) \
|
||||||
|
do { \
|
||||||
|
if (is_causal) { \
|
||||||
|
if (has_mask) FN<HEAD_DIM, true, true>(__VA_ARGS__); \
|
||||||
|
else FN<HEAD_DIM, true, false>(__VA_ARGS__); \
|
||||||
|
} else { \
|
||||||
|
if (has_mask) FN<HEAD_DIM, false, true>(__VA_ARGS__); \
|
||||||
|
else FN<HEAD_DIM, false, false>(__VA_ARGS__); \
|
||||||
|
} \
|
||||||
|
} while (0)
|
||||||
|
|
||||||
|
// ======================================================================
|
||||||
|
// Prefill
|
||||||
|
// ======================================================================
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
static inline void launch_prefill_mma(AttentionParams<bf16>& p) {
|
||||||
|
constexpr int WARPS = 4;
|
||||||
|
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
||||||
|
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
|
||||||
|
dim3 grid((p.q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS), p.q_head, p.batch);
|
||||||
|
dim3 block(Traits::NUM_THREADS);
|
||||||
|
attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block>>>(p);
|
||||||
|
}
|
||||||
|
#endif
|
||||||
|
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
static inline void launch_prefill_scalar(AttentionParams<bf16>& p) {
|
||||||
|
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
||||||
|
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
|
||||||
|
dim3 block(G, ROWS);
|
||||||
|
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask><<<grid, block>>>(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static inline void dispatch_prefill(AttentionParams<bf16>& p) {
|
||||||
|
bool is_causal = (p.causal_offset >= 0);
|
||||||
|
bool has_mask = (p.use_mask && p.mask);
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_mma, HEAD_DIM, p);
|
||||||
|
#else
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_prefill_scalar, HEAD_DIM, p);
|
||||||
|
#endif
|
||||||
|
}
|
||||||
|
|
||||||
|
// ======================================================================
|
||||||
|
// Decode
|
||||||
|
// ======================================================================
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
// BC=16: halves smem (16KB vs 32KB) → doubles occupancy (6 vs 3 blocks/SM).
|
||||||
|
// For D=256, BC=16 also reduces register pressure (fewer Sacc/PV frags),
|
||||||
|
// enabling STAGES=2 (double-buffer) within the 32KB smem budget — eliminates
|
||||||
|
// the 176-byte spill that STAGES=1+BC=32 suffered.
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
static inline void launch_decode_mma(AttentionParams<bf16>& p, int group_size) {
|
||||||
|
int G = p.q_head / p.kv_head;
|
||||||
|
constexpr int MAX_G = 16;
|
||||||
|
int num_passes = (G + MAX_G - 1) / MAX_G;
|
||||||
|
constexpr int BC = 16;
|
||||||
|
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||||
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total, 2);
|
||||||
|
constexpr int STAGES = 2;
|
||||||
|
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
||||||
|
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
||||||
|
attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask><<<grid, 32>>>(p);
|
||||||
|
}
|
||||||
|
#endif
|
||||||
|
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
static inline void launch_decode_scalar(AttentionParams<bf16>& p, int group_size) {
|
||||||
|
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||||
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||||
|
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
||||||
|
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
|
||||||
|
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
||||||
|
dim3 block(32, g);
|
||||||
|
attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static inline void dispatch_decode(AttentionParams<bf16>& p) {
|
||||||
|
bool is_causal = (p.causal_offset >= 0);
|
||||||
|
bool has_mask = (p.use_mask && p.mask);
|
||||||
|
int group_size = p.q_head / p.kv_head;
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_mma, HEAD_DIM, p, group_size);
|
||||||
|
#else
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_decode_scalar, HEAD_DIM, p, group_size);
|
||||||
|
#endif
|
||||||
|
|
||||||
|
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ======================================================================
|
||||||
|
// Paged Decode (SGLang-style: flat pool + req_to_token + kv_indptr)
|
||||||
|
// ======================================================================
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
static inline void launch_paged_decode_mma(PagedAttentionParams<bf16>& p, int) {
|
||||||
|
int G = p.q_head / p.kv_head;
|
||||||
|
constexpr int MAX_G = 16;
|
||||||
|
constexpr int BC = 16;
|
||||||
|
int num_passes = (G + MAX_G - 1) / MAX_G;
|
||||||
|
int tiles_total = (p.max_seq_len + BC - 1) / BC;
|
||||||
|
p.num_splits = compute_num_splits(p.batch * p.kv_head * num_passes, tiles_total, 2);
|
||||||
|
constexpr int STAGES = 2;
|
||||||
|
using Traits = KernelTraits<HEAD_DIM, BC, 1, STAGES>;
|
||||||
|
dim3 grid(p.kv_head * num_passes, p.batch, p.num_splits);
|
||||||
|
paged_attn_decode_split_kv_mma_kernel<Traits, IsCausal, HasMask> <<<grid, 32>>>(p);
|
||||||
|
}
|
||||||
|
#endif
|
||||||
|
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
static inline void launch_paged_decode_scalar(PagedAttentionParams<bf16>& p, int group_size) {
|
||||||
|
int chunks_total = (p.max_seq_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
||||||
|
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||||
|
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
|
||||||
|
int g = min(group_size, 32);
|
||||||
|
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
|
||||||
|
dim3 block(32, g);
|
||||||
|
paged_attn_decode_split_kv_kernel<HEAD_DIM, IsCausal, HasMask><<<grid, block, smem>>>(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static inline void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
|
||||||
|
bool is_causal = (p.causal_offset >= 0);
|
||||||
|
bool has_mask = (p.use_mask && p.mask);
|
||||||
|
int group_size = p.q_head / p.kv_head;
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_mma, HEAD_DIM, p, 0);
|
||||||
|
#else
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_decode_scalar, HEAD_DIM, p, group_size);
|
||||||
|
#endif
|
||||||
|
|
||||||
|
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ======================================================================
|
||||||
|
// Paged Prefill (SGLang-style: flat pool + ragged batch)
|
||||||
|
// ======================================================================
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
static inline void launch_paged_prefill_mma(PagedAttentionParams<bf16>& p) {
|
||||||
|
constexpr int WARPS = 4;
|
||||||
|
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
||||||
|
using Traits = KernelTraits<HEAD_DIM, BC, WARPS, 2>;
|
||||||
|
int max_q_tiles = (p.max_q_len + Traits::BR * WARPS - 1) / (Traits::BR * WARPS);
|
||||||
|
dim3 grid(max_q_tiles, p.q_head, p.batch);
|
||||||
|
dim3 block(Traits::NUM_THREADS);
|
||||||
|
paged_attn_prefill_split_q_mma_kernel<Traits, IsCausal, HasMask><<<grid, block>>>(p);
|
||||||
|
}
|
||||||
|
#endif
|
||||||
|
|
||||||
|
template <int HEAD_DIM, bool IsCausal, bool HasMask>
|
||||||
|
static inline void launch_paged_prefill_scalar(PagedAttentionParams<bf16>& p) {
|
||||||
|
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
||||||
|
int max_q_tiles = (p.max_q_len + ROWS - 1) / ROWS;
|
||||||
|
dim3 grid(max_q_tiles, p.q_head, p.batch);
|
||||||
|
dim3 block(G, ROWS);
|
||||||
|
paged_attn_prefill_split_q_kernel<HEAD_DIM, G, ROWS, P_BC, IsCausal, HasMask>
|
||||||
|
<<<grid, block>>>(p);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int HEAD_DIM>
|
||||||
|
static inline void dispatch_paged_prefill(PagedAttentionParams<bf16>& p) {
|
||||||
|
bool is_causal = (p.causal_offset >= 0);
|
||||||
|
bool has_mask = (p.use_mask && p.mask);
|
||||||
|
|
||||||
|
#ifndef ASTRAI_NO_MMA
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_prefill_mma, HEAD_DIM, p);
|
||||||
|
#else
|
||||||
|
DISPATCH_CAUSAL_MASK(is_causal, has_mask, launch_paged_prefill_scalar, HEAD_DIM, p);
|
||||||
|
#endif
|
||||||
|
}
|
||||||
@@ -0,0 +1,311 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <float.h>
|
||||||
|
#include <torch/extension.h>
|
||||||
|
#include <c10/cuda/CUDAGuard.h>
|
||||||
|
#include "attn_common.h"
|
||||||
|
#include "attn_warp_utils.cuh"
|
||||||
|
|
||||||
|
using bf16 = __nv_bfloat16;
|
||||||
|
|
||||||
|
// Dispatch head_dim: shared macro — avoids C++20 lambda template syntax.
|
||||||
|
// Usage: DISPATCH_HEAD_DIM(hd, fn, arg)
|
||||||
|
// Expands to: fn<32>(arg); fn<64>(arg); etc.
|
||||||
|
#define DISPATCH_HEAD_DIM(hd, fn, arg) \
|
||||||
|
switch (hd) { \
|
||||||
|
case 32: fn<32>(arg); break; \
|
||||||
|
case 64: fn<64>(arg); break; \
|
||||||
|
case 128: fn<128>(arg); break; \
|
||||||
|
case 256: fn<256>(arg); break; \
|
||||||
|
default: \
|
||||||
|
TORCH_CHECK(false, "unsupported head_dim ", hd, \
|
||||||
|
" (supported: 32, 64, 128, 256)"); \
|
||||||
|
}
|
||||||
|
|
||||||
|
// The split kernel unconditionally writes every (batch, q_head, split) slot it
|
||||||
|
// owns — including empty split ranges, which store m = -FLT_MAX so the combine
|
||||||
|
// skips them. Allocators are therefore left uninitialized (torch::empty); the
|
||||||
|
// per-call memset (torch::zeros / torch::full) was pure overhead.
|
||||||
|
template<typename P>
|
||||||
|
inline void alloc_split_partials(P& p) {
|
||||||
|
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
|
||||||
|
auto o_part = torch::empty(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, p.head_dim}, fopt);
|
||||||
|
auto ml_part = torch::empty(at::IntArrayRef{p.batch, p.q_head, MAX_SPLITS, 2}, fopt);
|
||||||
|
p.o_part = (float*)o_part.data_ptr();
|
||||||
|
p.ml_part = (float*)ml_part.data_ptr();
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- Shared Q-dims + strides extraction ----
|
||||||
|
template <typename P>
|
||||||
|
inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
|
||||||
|
if (layout == BLHD) q = q.transpose(1, 2);
|
||||||
|
p.batch = (int)q.size(0);
|
||||||
|
p.q_head = (int)q.size(1);
|
||||||
|
p.q_len = (int)q.size(2);
|
||||||
|
p.head_dim = (int)q.size(3);
|
||||||
|
p.q_stride_b = (int)q.stride(0);
|
||||||
|
p.q_stride_h = (int)q.stride(1);
|
||||||
|
p.q_stride_l = (int)q.stride(2);
|
||||||
|
p.q_stride_d = (int)q.stride(3);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- Shared mask packing ----
|
||||||
|
// Accepts 2D [batch, kv_len], 3D [batch, q_len, kv_len],
|
||||||
|
// or 4D [batch, n_heads, q_len, kv_len].
|
||||||
|
// Head/q dimensions with size 1 broadcast (stride set to 0).
|
||||||
|
template <typename P>
|
||||||
|
inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
|
||||||
|
if (p.use_mask) {
|
||||||
|
auto m = mask.value();
|
||||||
|
TORCH_CHECK(m.is_cuda(), "mask must be on CUDA");
|
||||||
|
TORCH_CHECK(m.dtype() == torch::kBool, "mask must be bool");
|
||||||
|
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
|
||||||
|
TORCH_CHECK(m.size(m.dim() - 1) == p.kv_len, "mask kv_len mismatch");
|
||||||
|
if (m.dim() == 2) {
|
||||||
|
p.mask_b_stride = (int)m.stride(0);
|
||||||
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_q_stride = 0;
|
||||||
|
} else if (m.dim() == 3) {
|
||||||
|
TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_len, "mask q_len mismatch");
|
||||||
|
p.mask_b_stride = (int)m.stride(0);
|
||||||
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_q_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
|
||||||
|
} else if (m.dim() == 4) {
|
||||||
|
TORCH_CHECK(m.size(2) == 1 || m.size(2) == p.q_len, "mask q_len mismatch");
|
||||||
|
p.mask_b_stride = (int)m.stride(0);
|
||||||
|
p.mask_h_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
|
||||||
|
p.mask_q_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2);
|
||||||
|
} else {
|
||||||
|
TORCH_CHECK(false, "mask must be 2D, 3D, or 4D");
|
||||||
|
}
|
||||||
|
p.mask = m.data_ptr<bool>();
|
||||||
|
} else {
|
||||||
|
p.mask = nullptr;
|
||||||
|
p.mask_b_stride = 0;
|
||||||
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_q_stride = 0;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- attn_pack_params (contiguous KV) ----
|
||||||
|
template<typename T>
|
||||||
|
inline void attn_pack_params(
|
||||||
|
torch::Tensor q,
|
||||||
|
torch::Tensor k,
|
||||||
|
torch::Tensor v,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
int64_t causal_offset,
|
||||||
|
double scale,
|
||||||
|
int64_t layout,
|
||||||
|
AttentionParams<T>& p
|
||||||
|
) {
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
|
||||||
|
TORCH_CHECK(q.is_cuda() && k.is_cuda() && v.is_cuda());
|
||||||
|
TORCH_CHECK(q.dtype() == torch::kBFloat16);
|
||||||
|
TORCH_CHECK(k.dtype() == torch::kBFloat16);
|
||||||
|
TORCH_CHECK(v.dtype() == torch::kBFloat16);
|
||||||
|
TORCH_CHECK(k.sizes() == v.sizes(), "K and V must have identical shapes");
|
||||||
|
TORCH_CHECK(q.dim() == 4 && k.dim() == 4, "Q/K/V must be 4D");
|
||||||
|
|
||||||
|
extract_q_dims_and_strides(q, layout, p);
|
||||||
|
|
||||||
|
if (layout == BLHD) k = k.transpose(1, 2), v = v.transpose(1, 2);
|
||||||
|
|
||||||
|
p.kv_head = (int)k.size(1);
|
||||||
|
p.kv_len = (int)k.size(2);
|
||||||
|
TORCH_CHECK(k.size(3) == p.head_dim, "K/V head_dim must match Q");
|
||||||
|
|
||||||
|
p.kv_stride_b = (int)k.stride(0);
|
||||||
|
p.kv_stride_h = (int)k.stride(1);
|
||||||
|
p.kv_stride_l = (int)k.stride(2);
|
||||||
|
p.kv_stride_d = (int)k.stride(3);
|
||||||
|
|
||||||
|
p.causal_offset = (int)causal_offset;
|
||||||
|
p.use_mask = mask.has_value() ? 1 : 0;
|
||||||
|
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
|
||||||
|
|
||||||
|
p.q = (const T*)q.data_ptr();
|
||||||
|
p.k = (const T*)k.data_ptr();
|
||||||
|
p.v = (const T*)v.data_ptr();
|
||||||
|
p.o = nullptr;
|
||||||
|
p.o_part = nullptr;
|
||||||
|
p.ml_part = nullptr;
|
||||||
|
|
||||||
|
pack_mask(mask, p);
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- attn_pack_paged_decode_params ----
|
||||||
|
// SGLang-style: flat KV pool + req_to_token indexing + variable
|
||||||
|
// seq_lens via kv_indptr. Q is [batch, q_head, head_dim] (q_len=1 per req).
|
||||||
|
template<typename T>
|
||||||
|
inline void attn_pack_paged_decode_params(
|
||||||
|
torch::Tensor q,
|
||||||
|
torch::Tensor k_cache,
|
||||||
|
torch::Tensor v_cache,
|
||||||
|
torch::Tensor req_to_token,
|
||||||
|
torch::Tensor req_pool_indices,
|
||||||
|
torch::Tensor kv_indptr,
|
||||||
|
int64_t max_seq_len,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
int64_t causal_offset,
|
||||||
|
double scale,
|
||||||
|
PagedAttentionParams<T>& p
|
||||||
|
) {
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
|
||||||
|
TORCH_CHECK(q.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
|
||||||
|
TORCH_CHECK(req_to_token.is_cuda() && req_pool_indices.is_cuda() && kv_indptr.is_cuda());
|
||||||
|
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
|
||||||
|
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
|
||||||
|
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
|
||||||
|
TORCH_CHECK(req_to_token.dtype() == torch::kLong, "req_to_token must be int64");
|
||||||
|
TORCH_CHECK(req_pool_indices.dtype() == torch::kLong, "req_pool_indices must be int64");
|
||||||
|
TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32");
|
||||||
|
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must match");
|
||||||
|
TORCH_CHECK(k_cache.dim() == 3, "k_cache must be 3D [size, kv_head, head_dim]");
|
||||||
|
TORCH_CHECK(q.dim() == 3, "q must be 3D [batch, q_head, head_dim]");
|
||||||
|
|
||||||
|
p.batch = (int)q.size(0);
|
||||||
|
p.q_head = (int)q.size(1);
|
||||||
|
p.head_dim = (int)q.size(2);
|
||||||
|
p.kv_head = (int)k_cache.size(1);
|
||||||
|
TORCH_CHECK(k_cache.size(2) == p.head_dim, "k_cache head_dim mismatch");
|
||||||
|
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
||||||
|
TORCH_CHECK(p.q_head % p.kv_head == 0, "q_head must be divisible by kv_head");
|
||||||
|
|
||||||
|
p.q_stride_l = (int)q.stride(0);
|
||||||
|
p.q_stride_h = (int)q.stride(1);
|
||||||
|
p.q_stride_d = (int)q.stride(2);
|
||||||
|
|
||||||
|
p.k_cache = (const T*)k_cache.data_ptr();
|
||||||
|
p.v_cache = (const T*)v_cache.data_ptr();
|
||||||
|
p.q = (const T*)q.data_ptr();
|
||||||
|
p.req_to_token = req_to_token.data_ptr<int64_t>();
|
||||||
|
p.req_pool_indices = req_pool_indices.data_ptr<int64_t>();
|
||||||
|
p.kv_indptr = kv_indptr.data_ptr<int>();
|
||||||
|
p.qo_indptr = nullptr;
|
||||||
|
p.max_context_len = (int)req_to_token.size(1);
|
||||||
|
p.max_seq_len = (int)max_seq_len;
|
||||||
|
p.total_q = p.batch; // decode: 1 Q token per request
|
||||||
|
p.max_q_len = 1;
|
||||||
|
|
||||||
|
p.causal_offset = (int)causal_offset;
|
||||||
|
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
|
||||||
|
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
|
||||||
|
|
||||||
|
if (p.use_mask) {
|
||||||
|
auto m = mask.value();
|
||||||
|
TORCH_CHECK(m.is_cuda() && m.dtype() == torch::kBool, "mask must be bool CUDA");
|
||||||
|
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
|
||||||
|
p.mask_b_stride = (int)m.stride(0);
|
||||||
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_q_stride = 0;
|
||||||
|
p.mask = m.data_ptr<bool>();
|
||||||
|
} else {
|
||||||
|
p.mask = nullptr;
|
||||||
|
p.mask_b_stride = 0;
|
||||||
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_q_stride = 0;
|
||||||
|
}
|
||||||
|
|
||||||
|
p.o = nullptr;
|
||||||
|
p.o_part = nullptr;
|
||||||
|
p.ml_part = nullptr;
|
||||||
|
}
|
||||||
|
|
||||||
|
// ---- attn_pack_paged_prefill_params ----
|
||||||
|
// SGLang-style: flat KV pool + req_to_token + ragged batch via qo_indptr.
|
||||||
|
// Q is [total_q, q_head, head_dim] (flattened across all requests).
|
||||||
|
template<typename T>
|
||||||
|
inline void attn_pack_paged_prefill_params(
|
||||||
|
torch::Tensor q,
|
||||||
|
torch::Tensor k_cache,
|
||||||
|
torch::Tensor v_cache,
|
||||||
|
torch::Tensor req_to_token,
|
||||||
|
torch::Tensor req_pool_indices,
|
||||||
|
torch::Tensor kv_indptr,
|
||||||
|
torch::Tensor qo_indptr,
|
||||||
|
c10::optional<torch::Tensor> mask,
|
||||||
|
int64_t max_q_len,
|
||||||
|
int64_t causal_offset,
|
||||||
|
double scale,
|
||||||
|
PagedAttentionParams<T>& p
|
||||||
|
) {
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
|
|
||||||
|
TORCH_CHECK(q.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
|
||||||
|
TORCH_CHECK(req_to_token.is_cuda() && req_pool_indices.is_cuda());
|
||||||
|
TORCH_CHECK(kv_indptr.is_cuda() && qo_indptr.is_cuda());
|
||||||
|
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
|
||||||
|
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
|
||||||
|
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
|
||||||
|
TORCH_CHECK(req_to_token.dtype() == torch::kLong, "req_to_token must be int64");
|
||||||
|
TORCH_CHECK(req_pool_indices.dtype() == torch::kLong, "req_pool_indices must be int64");
|
||||||
|
TORCH_CHECK(kv_indptr.dtype() == torch::kInt32, "kv_indptr must be int32");
|
||||||
|
TORCH_CHECK(qo_indptr.dtype() == torch::kInt32, "qo_indptr must be int32");
|
||||||
|
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must match");
|
||||||
|
TORCH_CHECK(k_cache.dim() == 3, "k_cache must be 3D [size, kv_head, head_dim]");
|
||||||
|
TORCH_CHECK(q.dim() == 3, "q must be 3D [total_q, q_head, head_dim]");
|
||||||
|
|
||||||
|
p.q_head = (int)q.size(1);
|
||||||
|
p.head_dim = (int)q.size(2);
|
||||||
|
p.kv_head = (int)k_cache.size(1);
|
||||||
|
p.batch = (int)req_pool_indices.size(0);
|
||||||
|
TORCH_CHECK(k_cache.size(2) == p.head_dim, "k_cache head_dim mismatch");
|
||||||
|
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
|
||||||
|
TORCH_CHECK(p.q_head % p.kv_head == 0, "q_head must be divisible by kv_head");
|
||||||
|
TORCH_CHECK(kv_indptr.size(0) == p.batch + 1, "kv_indptr must be [batch+1]");
|
||||||
|
TORCH_CHECK(qo_indptr.size(0) == p.batch + 1, "qo_indptr must be [batch+1]");
|
||||||
|
|
||||||
|
p.q_stride_l = (int)q.stride(0);
|
||||||
|
p.q_stride_h = (int)q.stride(1);
|
||||||
|
p.q_stride_d = (int)q.stride(2);
|
||||||
|
|
||||||
|
p.k_cache = (const T*)k_cache.data_ptr();
|
||||||
|
p.v_cache = (const T*)v_cache.data_ptr();
|
||||||
|
p.q = (const T*)q.data_ptr();
|
||||||
|
p.req_to_token = req_to_token.data_ptr<int64_t>();
|
||||||
|
p.req_pool_indices = req_pool_indices.data_ptr<int64_t>();
|
||||||
|
p.kv_indptr = kv_indptr.data_ptr<int>();
|
||||||
|
p.qo_indptr = qo_indptr.data_ptr<int>();
|
||||||
|
p.max_context_len = (int)req_to_token.size(1);
|
||||||
|
p.total_q = (int)q.size(0); // prefill: flattened Q across all requests
|
||||||
|
p.max_q_len = (int)max_q_len;
|
||||||
|
// max_seq_len is unused by the prefill path (decode uses it for split
|
||||||
|
// computation); fill with max_q_len only to keep the POD struct defined.
|
||||||
|
p.max_seq_len = p.max_q_len;
|
||||||
|
|
||||||
|
p.causal_offset = (int)causal_offset;
|
||||||
|
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
|
||||||
|
if (p.use_mask) {
|
||||||
|
auto m = mask.value();
|
||||||
|
TORCH_CHECK(m.is_cuda() && m.dtype() == torch::kBool, "mask must be bool CUDA");
|
||||||
|
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
|
||||||
|
if (m.dim() == 2) {
|
||||||
|
TORCH_CHECK(m.size(1) <= p.max_context_len, "mask kv_len mismatch");
|
||||||
|
p.mask_b_stride = (int)m.stride(0);
|
||||||
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_q_stride = 0;
|
||||||
|
} else if (m.dim() == 4) {
|
||||||
|
TORCH_CHECK(m.size(1) == 1 || m.size(1) == p.q_head, "mask head mismatch");
|
||||||
|
TORCH_CHECK(m.size(2) == 1 || m.size(2) == p.max_q_len, "mask q_len mismatch");
|
||||||
|
TORCH_CHECK(m.size(3) <= p.max_context_len, "mask kv_len mismatch");
|
||||||
|
p.mask_b_stride = (int)m.stride(0);
|
||||||
|
p.mask_h_stride = (m.size(1) == 1) ? 0 : (int)m.stride(1);
|
||||||
|
p.mask_q_stride = (m.size(2) == 1) ? 0 : (int)m.stride(2);
|
||||||
|
} else {
|
||||||
|
TORCH_CHECK(false, "mask must be 2D or 4D");
|
||||||
|
}
|
||||||
|
p.mask = m.data_ptr<bool>();
|
||||||
|
} else {
|
||||||
|
p.mask = nullptr;
|
||||||
|
p.mask_b_stride = 0;
|
||||||
|
p.mask_h_stride = 0;
|
||||||
|
p.mask_q_stride = 0;
|
||||||
|
}
|
||||||
|
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
|
||||||
|
|
||||||
|
p.o = nullptr;
|
||||||
|
p.o_part = nullptr;
|
||||||
|
p.ml_part = nullptr;
|
||||||
|
}
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user