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

|
<div align="center">
|
||||||
|
|
||||||
<div style="display: flex; flex-direction: column; align-items: center; justify-content: center; text-align: center; font-size: 16px; font-weight: bold; margin-top: 50px;">
|
|
||||||
|
|
||||||
<div>
|
<img src="assets/images/logo.png" width="auto" alt="Logo">
|
||||||
<a href="#english" style="text-decoration: none; margin: 0 10px; color: blue;">English</a> |
|
<p>
|
||||||
<a href="#chinese" style="text-decoration: none; margin: 0 10px; color: blue;">中文</a>
|
<strong>A lightweight Transformer training & inference framework</strong>
|
||||||
</div>
|
</p>
|
||||||
|
|
||||||
<h1 style="margin: 20px 0 0 0; font-size: 2.5em; font-weight: bold;">KHAOSZ </h1>
|
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<h2 id="english">English Version</h2>
|
<div align="center">
|
||||||
|
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
||||||
|
<img src="https://img.shields.io/badge/license-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>
|
||||||
|
|
||||||
This is a Chinese-English bilingual Transformer model supporting both languages. It contains model configurations and training workflows, completing training by loading parameters defined in `param_path/config.json`. The training script `train.py` parses command-line arguments, including dataset root directory, number of training epochs, batch size, checkpoint interval, and checkpoint directory.
|
<div align="center">
|
||||||
|
<a href="#english">English</a> •
|
||||||
|
<a href="assets/docs/README-zh-CN.md">中文</a> •
|
||||||
|
<a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> •
|
||||||
|
<a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> •
|
||||||
|
<a href="https://huggingface.co/ViperEk/">HuggingFace</a>
|
||||||
|
</div>
|
||||||
|
|
||||||
**Model Download Options (Choose One):**
|
<br>
|
||||||
|
|
||||||
1. Visit [HuggingFace](https://huggingface.co/ViperEk/KHAOSZ) to access **Files and versions**
|
## 📖 Table of Contents
|
||||||
2. Run `scripts/download.py` to download parameters
|
|
||||||
|
|
||||||
**Demo Video:** [bilibili](https://www.bilibili.com/video/BV1z5RPYHEkd)
|
- [Features](#features)
|
||||||
|
- [Quick Start](#quick-start)
|
||||||
|
- [Documentation](#documentation)
|
||||||
|
- [Contributing](#contributing)
|
||||||
|
- [Community](#community)
|
||||||
|
- [License](#license)
|
||||||
|
|
||||||
Training dataset sources are listed in the **Model Card** section of the HuggingFace download link.
|
---
|
||||||
|
|
||||||
**License:** Code follows Apache-2.0 protocol. Please credit the source code when used.
|
<a id="english"></a>
|
||||||
|
## English
|
||||||
|
|
||||||
- **📊 Device Selection:** Code defaults to CUDA training
|
### Features
|
||||||
- **🌐 Performance Optimization:** `dtype=torch.bfloat16` is enabled to accelerate training and reduce memory usage. Ensure hardware supports this feature.
|
|
||||||
- **🤖 Language Support:** Model supports Chinese and English training. The BBPE tokenizer was trained without multilingual text, so OOV (out-of-vocabulary) issues are minimized for these languages but may exist for others.
|
|
||||||
|
|
||||||
### 📌 Training Guide
|
- 🚀 **High Performance**: Optimized for both training and inference with efficient parallelization.
|
||||||
|
- 🔧 **Flexible**: Support for seq/sft/dpo/grpo training, customizable model architectures.
|
||||||
|
- 💡 **Easy to Use**: Simple API with comprehensive examples and demos.
|
||||||
|
- 📦 **Lightweight**: Minimal dependencies, easy to deploy.
|
||||||
|
- 🔬 **Research‑Friendly**: Modular design, easy to experiment with new ideas.
|
||||||
|
- 🤗 **HuggingFace-Style API**: AutoModel/AutoTokenizer APIs inspired by HuggingFace for easy model and tokenizer loading.
|
||||||
|
- 🔌 **Dual API Compatibility**: Supports both OpenAI and Anthropic chat completion APIs out of the box.
|
||||||
|
|
||||||
To train this Transformer model, follow these steps:
|
### Quick Start
|
||||||
|
|
||||||
**(1). Prepare Dataset:**
|
#### Installation
|
||||||
|
|
||||||
Place datasets in the designated root directory. Files should be text documents in Chinese, English, or mixed. Format should align with model input requirements - preferably pre-tokenized token_ids stored as `torch.Tensor` (using `torch.Tensor` saves memory compared to Python lists, which default to 64-bit precision).
|
|
||||||
|
|
||||||
**(2). Install Dependencies:**
|
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
pip install -r requirements.txt
|
git clone https://github.com/ViperEkura/AstrAI.git
|
||||||
pip install .
|
cd AstrAI
|
||||||
|
pip install -e .
|
||||||
```
|
```
|
||||||
|
|
||||||
**(3). Run Training Script:**
|
For development dependencies:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python train.py \
|
pip install -e ".[dev]"
|
||||||
--train_type=train_type[seq, sft, dpo] \
|
|
||||||
--data_root_path=/path/to/dataset \
|
|
||||||
--param_path=/path/to/param_path \
|
|
||||||
--n_epoch=5 \
|
|
||||||
--batch_size=8 \
|
|
||||||
--max_lr=2e-4 \
|
|
||||||
--checkpoint_interval=10000 \
|
|
||||||
--checkpoint_dir=checkpoints
|
|
||||||
```
|
```
|
||||||
|
|
||||||
**Parameters Explanation:**
|
#### Train a Model
|
||||||
- `--train_type`: Training type (seq, sft, dpo)
|
|
||||||
- `--data_root_path`: Root directory of the dataset
|
|
||||||
- `--param_path`: Path to the model training parameters
|
|
||||||
- `--n_epoch`: Total number of training epochs
|
|
||||||
- `--batch_size`: Batch size
|
|
||||||
- `--accumulation_steps`: Number of batches per training step
|
|
||||||
- `--warmup_steps`: Number of warmup steps
|
|
||||||
- `--max_lr`: Maximum learning rate (using warmup + cosine decay)
|
|
||||||
- `--checkpoint_interval`: Checkpoint saving interval
|
|
||||||
- `--checkpoint_dir`: Directory to save checkpoints
|
|
||||||
- `--resume_dir`: Resume training from the specified path
|
|
||||||
|
|
||||||
Training logs will be saved in `train_log.txt`. Checkpoints will be saved in the specified directory for resuming training or evaluation.
|
|
||||||
|
|
||||||
### 👉 Usage Guide
|
|
||||||
|
|
||||||
**(1). Chatting with the Model:**
|
|
||||||
|
|
||||||
Open `chat.py` or use streaming/non-streaming interfaces:
|
|
||||||
|
|
||||||
**Streaming Output:**
|
|
||||||
```python
|
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
model_dir = "your_model_parameter_dir"
|
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
history = []
|
|
||||||
|
|
||||||
while True:
|
|
||||||
query = input(">> ")
|
|
||||||
if query == "!exit":
|
|
||||||
break
|
|
||||||
|
|
||||||
response_size = 0
|
|
||||||
for response, history in model.stream_generate(
|
|
||||||
query=query,
|
|
||||||
history=history,
|
|
||||||
temperature=0.85,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50
|
|
||||||
):
|
|
||||||
print(response[response_size:], end="")
|
|
||||||
response_size = len(response)
|
|
||||||
```
|
|
||||||
|
|
||||||
**Non-streaming Output:**
|
|
||||||
```python
|
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
model_dir = "your_model_parameter_dir"
|
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
history = []
|
|
||||||
|
|
||||||
while True:
|
|
||||||
query = input(">> ")
|
|
||||||
if query == "!exit":
|
|
||||||
break
|
|
||||||
|
|
||||||
response = model.generate(
|
|
||||||
query=query,
|
|
||||||
history=history,
|
|
||||||
temperature=0.85,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50
|
|
||||||
)
|
|
||||||
print(response)
|
|
||||||
```
|
|
||||||
|
|
||||||
**(2) Retrieval-Augmented Generation (RAG):**
|
|
||||||
|
|
||||||
```python
|
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
model_dir = "your_model_parameter_dir"
|
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
|
|
||||||
retrieved_content = model.retrieve_generate(
|
|
||||||
query=query,
|
|
||||||
retrieve_top_k=5,
|
|
||||||
temperature=0.6,
|
|
||||||
top_k=30,
|
|
||||||
top_p=0.95
|
|
||||||
)
|
|
||||||
print(retrieved_content)
|
|
||||||
```
|
|
||||||
|
|
||||||
### 📌 Model Specifications
|
|
||||||
|
|
||||||
This model is based on a 24-layer Transformer with parameters defined in `config.json`, totaling approximately 1.0 billion (1.0B) parameters.
|
|
||||||
|
|
||||||
**Key Design Choices:**
|
|
||||||
- Weight tying between embedding and final linear layers (standard for small models to save parameters)
|
|
||||||
- Embedding layer optimization: Without weight tying, a 10,000-word vocabulary would consume ~102M parameters (0.1B)
|
|
||||||
|
|
||||||
**Limitations:**
|
|
||||||
- May struggle with complex language phenomena due to smaller parameter size
|
|
||||||
- Prone to overfitting on specialized datasets
|
|
||||||
- Limited multilingual capabilities
|
|
||||||
|
|
||||||
**Advantages:**
|
|
||||||
- Runs efficiently on lower-spec hardware
|
|
||||||
- Shorter training time compared to larger models
|
|
||||||
|
|
||||||
**Training Pipeline:**
|
|
||||||
The model has completed pre-training + SFT (Supervised Fine-Tuning) + DPO (Direct Preference Optimization) workflows. All corresponding training code is included in the repository.
|
|
||||||
|
|
||||||
|
|
||||||
<h2 id="chinese">中文版本</h2>
|
|
||||||
这是一个支持中英文双语的 Transformer 模型,能够处理两种语言。模型包含配置文件和训练流程,通过加载 `param_path/config.json` 中定义的参数完成训练。训练脚本 `train.py` 支持命令行参数解析,包括数据集根目录、训练轮数(epochs)、批量大小(batch size)、检查点保存间隔、检查点目录等。
|
|
||||||
|
|
||||||
**模型下载选项(任选其一):**
|
|
||||||
|
|
||||||
1. 访问 [HuggingFace](https://huggingface.co/ViperEk/KHAOSZ) 查看 **Files and versions**
|
|
||||||
2. 运行 `scripts/download.py` 下载模型参数
|
|
||||||
|
|
||||||
**演示视频:** [bilibili](https://www.bilibili.com/video/BV1z5RPYHEkd)
|
|
||||||
|
|
||||||
训练数据来源请参见 HuggingFace 下载页面中的 **Model Card** 部分。
|
|
||||||
|
|
||||||
**许可证:** 代码遵循 Apache-2.0 协议,使用时请注明出处。
|
|
||||||
|
|
||||||
- **📊 设备选择:** 默认使用 CUDA 进行训练
|
|
||||||
- **🌐 性能优化:** 启用 `dtype=torch.bfloat16` 以加速训练并减少内存占用,请确保硬件支持该特性
|
|
||||||
- **🤖 语言支持:** 模型支持中文和英文训练。由于 BBPE 分词器未使用多语言文本训练,因此中英文的 OOV(未登录词)问题较少,其他语言可能存在 OOV 问题
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
### 📌 训练指南
|
|
||||||
|
|
||||||
要训练该 Transformer 模型,请按照以下步骤操作:
|
|
||||||
|
|
||||||
#### **(1). 准备数据集:**
|
|
||||||
|
|
||||||
将数据集放置在指定的根目录下。文件应为包含中文、英文或混合文本的文本文档。格式应符合模型输入要求——建议使用预分词后的 `token_ids` 并以 `torch.Tensor` 格式保存(使用 `torch.Tensor` 相比 Python 列表更节省内存,列表默认为 64 位精度)。
|
|
||||||
|
|
||||||
#### **(2). 安装依赖:**
|
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
pip install -r requirements.txt
|
CUDA_VISIBLE_DEVICES=0,1,2,3 python scripts/tools/train.py \
|
||||||
pip install .
|
--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
|
||||||
```
|
```
|
||||||
|
|
||||||
#### **(3). 运行训练脚本:**
|
Full reference at [Parameter Guide](assets/docs/params.md).
|
||||||
|
|
||||||
|
#### Generate Text
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python train.py \
|
python scripts/tools/generate.py \
|
||||||
--train_type=train_type[seq, sft, dpo] \
|
--param_path /path/to/model \
|
||||||
--data_root_path=/path/to/dataset \
|
--input_json_file /path/to/input.json \
|
||||||
--param_path=/path/to/param_path \
|
--output_json_file /path/to/output.json
|
||||||
--n_epoch=5 \
|
|
||||||
--batch_size=8 \
|
|
||||||
--max_lr=2e-4 \
|
|
||||||
--checkpoint_interval=10000 \
|
|
||||||
--checkpoint_dir=checkpoints
|
|
||||||
```
|
```
|
||||||
|
|
||||||
**参数说明:**
|
#### Docker
|
||||||
- `--train_type`: 训练类型(seq, sft, dpo)
|
|
||||||
- `--data_root_path`: 数据集根目录
|
|
||||||
- `--param_path`: 模型训练参数路径
|
|
||||||
- `--n_epoch`: 总训练轮数
|
|
||||||
- `--batch_size`: 批量大小
|
|
||||||
- `--accumulation_steps`: 每个训练步骤的 batch 数量
|
|
||||||
- `--warmup_steps`: 预热步数(warmup steps)
|
|
||||||
- `--max_lr`: 最大学习率(使用预热 + 余弦衰减)
|
|
||||||
- `--checkpoint_interval`: 检查点保存间隔
|
|
||||||
- `--checkpoint_dir`: 检查点保存目录
|
|
||||||
- `--resume_dir`: 从指定路径恢复训练
|
|
||||||
|
|
||||||
训练日志将保存在 `train_log.txt` 中。检查点将保存在指定目录,用于恢复训练或评估。
|
Build and run with Docker (recommended for GPU environments):
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Build image
|
||||||
|
docker build -t astrai:latest .
|
||||||
|
|
||||||
|
# Run with GPU support
|
||||||
|
docker run --gpus all -it astrai:latest
|
||||||
|
|
||||||
### 👉 使用指南
|
# Run with specific GPUs
|
||||||
|
docker run --gpus '"device=0,1"' -it astrai:latest
|
||||||
|
|
||||||
#### **(1). 与模型对话:**
|
# Run inference server
|
||||||
|
docker run --gpus all -p 8000:8000 astrai:latest \
|
||||||
|
python -m scripts.tools.server --port 8000 --device cuda
|
||||||
|
|
||||||
打开 `chat.py` 或使用流式/非流式接口:
|
# Run with volume mount for data
|
||||||
|
docker run --gpus all -v /path/to/data:/data -it astrai:latest
|
||||||
|
|
||||||
**流式输出:**
|
# Docker Compose (GPU, default)
|
||||||
```python
|
docker compose up -d
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
model_dir = "your_model_parameter_dir"
|
# Docker Compose (CPU only)
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
docker compose --profile cpu up -d
|
||||||
history = []
|
|
||||||
|
|
||||||
while True:
|
|
||||||
query = input(">> ")
|
|
||||||
if query == "!exit":
|
|
||||||
break
|
|
||||||
|
|
||||||
response_size = 0
|
|
||||||
for response, history in model.stream_generate(
|
|
||||||
query=query,
|
|
||||||
history=history,
|
|
||||||
temperature=0.85,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50
|
|
||||||
):
|
|
||||||
print(response[response_size:], end="")
|
|
||||||
response_size = len(response)
|
|
||||||
```
|
```
|
||||||
|
|
||||||
**非流式输出:**
|
> **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`.
|
||||||
```python
|
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
model_dir = "your_model_parameter_dir"
|
#### Start HTTP Server
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
history = []
|
|
||||||
|
|
||||||
while True:
|
Start the inference server with OpenAI and Anthropic-compatible HTTP API:
|
||||||
query = input(">> ")
|
|
||||||
if query == "!exit":
|
```bash
|
||||||
break
|
python -m scripts.tools.server --port 8000 --device cuda
|
||||||
|
|
||||||
response = model.generate(
|
|
||||||
query=query,
|
|
||||||
history=history,
|
|
||||||
temperature=0.85,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50
|
|
||||||
)
|
|
||||||
print(response)
|
|
||||||
```
|
```
|
||||||
|
|
||||||
#### **(2). 基于检索的生成(RAG):**
|
Make requests:
|
||||||
|
|
||||||
```python
|
```bash
|
||||||
import torch
|
# OpenAI-compatible
|
||||||
from khaosz import Khaosz
|
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"messages": [{"role": "user", "content": "Hello"}],
|
||||||
|
"max_tokens": 512
|
||||||
|
}'
|
||||||
|
|
||||||
model_dir = "your_model_parameter_dir"
|
# OpenAI-compatible streaming
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"messages": [{"role": "user", "content": "Tell a story"}],
|
||||||
|
"stream": true,
|
||||||
|
"max_tokens": 500
|
||||||
|
}'
|
||||||
|
|
||||||
retrieved_content = model.retrieve_generate(
|
# Anthropic-compatible
|
||||||
query=query,
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
retrieve_top_k=5,
|
-H "Content-Type: application/json" \
|
||||||
temperature=0.6,
|
-d '{
|
||||||
top_k=30,
|
"model": "astrai",
|
||||||
top_p=0.95
|
"system": "You are a helpful assistant.",
|
||||||
)
|
"messages": [{"role": "user", "content": "Hello"}],
|
||||||
print(retrieved_content)
|
"max_tokens": 512
|
||||||
|
}'
|
||||||
|
|
||||||
|
# Anthropic-compatible streaming with stop sequences
|
||||||
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"model": "astrai",
|
||||||
|
"messages": [{"role": "user", "content": "Write a story"}],
|
||||||
|
"max_tokens": 500,
|
||||||
|
"stream": true,
|
||||||
|
"stop_sequences": ["The end"]
|
||||||
|
}'
|
||||||
|
|
||||||
|
# Health check
|
||||||
|
curl http://localhost:8000/health
|
||||||
```
|
```
|
||||||
|
|
||||||
|
#### Demo
|
||||||
|
|
||||||
|
Check out the demos in the `scripts/demo/` folder:
|
||||||
|
|
||||||
### 📌 模型规格说明(重复部分)
|
```bash
|
||||||
|
# Download pre‑processed data (required before running demos)
|
||||||
|
python scripts/demo/download.py
|
||||||
|
|
||||||
该模型基于一个 24 层的 Transformer 架构,参数配置定义在 `config.json` 中,总参数量约为 10 亿(1.0B)。
|
# Interactive streaming chat
|
||||||
|
python scripts/demo/stream_chat.py
|
||||||
|
|
||||||
**关键设计选择:**
|
# Batch generation
|
||||||
- 在嵌入层(embedding)与最终线性层之间进行权重绑定(weight tying),这是小型模型中常见的节省参数量的做法
|
python scripts/demo/generate_batch.py
|
||||||
- 嵌入层优化:若不进行权重绑定,一个包含 10,000 个词的词汇表将消耗约 1.02 亿(0.1B)参数
|
|
||||||
|
|
||||||
**局限性:**
|
# Auto‑regressive generation
|
||||||
- 由于参数规模较小,可能在处理复杂语言现象时表现受限
|
python scripts/demo/generate_ar.py
|
||||||
- 在特定领域的数据集上容易出现过拟合
|
```
|
||||||
- 多语言能力有限
|
|
||||||
|
|
||||||
**优势:**
|
Watch a video walkthrough on [bilibili](https://www.bilibili.com/video/BV1z5RPYHEkd).
|
||||||
- 可在低配置硬件上高效运行
|
|
||||||
- 相较于大型模型,训练时间更短
|
|
||||||
|
|
||||||
**训练流程:**
|
### Documentation
|
||||||
该模型已完成预训练(pre-training)+ 监督微调(SFT, Supervised Fine-Tuning)+ 直接偏好优化(DPO, Direct Preference Optimization)的全流程。所有相关的训练代码均已包含在代码库中。
|
|
||||||
|
| Document | Description |
|
||||||
|
|----------|-------------|
|
||||||
|
| [Parameter Guide](./assets/docs/params.md) | Training & inference parameters |
|
||||||
|
| [Design Document](./assets/docs/design.md) | Framework architecture & module design |
|
||||||
|
| [Data Flow](./assets/docs/dataflow.md) | Data processing pipeline details |
|
||||||
|
| [Model Introduction](./assets/docs/introduction.md) | Model architecture & technical details |
|
||||||
|
|
||||||
|
### Contributing
|
||||||
|
|
||||||
|
We welcome contributions! Please see our [Contributing Guidelines](CONTRIBUTING.md) for details.
|
||||||
|
|
||||||
|
1. Fork the repository.
|
||||||
|
2. Create a feature branch.
|
||||||
|
3. Commit your changes.
|
||||||
|
4. Open a Pull Request.
|
||||||
|
|
||||||
|
For major changes, please open an issue first to discuss what you would like to change.
|
||||||
|
|
||||||
|
### Community
|
||||||
|
|
||||||
|
- **GitHub Issues**: [Issue Tracker](https://github.com/ViperEkura/AstrAI/issues)
|
||||||
|
- **Discussions**: [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions)
|
||||||
|
- **HuggingFace**: [Model Hub](https://huggingface.co/ViperEk)
|
||||||
|
|
||||||
|
### License
|
||||||
|
|
||||||
|
This project is licensed under the [GPL-3.0 License](LICENSE).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
<div align="center">
|
||||||
|
<em>A lightweight Transformer framework designed for both high performance and ease of use.</em>
|
||||||
|
</div>
|
||||||
@@ -0,0 +1,236 @@
|
|||||||
|
<div align="center">
|
||||||
|
|
||||||
|
<img src="../images/logo.png" width="auto" alt="Logo">
|
||||||
|
|
||||||
|
<div>
|
||||||
|
<a href="../../README.md">English</a> •
|
||||||
|
<a href="#chinese">中文</a>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<p>
|
||||||
|
<strong>轻量级 Transformer 训练与推理框架</strong>
|
||||||
|
</p>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<div align="center">
|
||||||
|
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
||||||
|
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license">
|
||||||
|
<img src="https://img.shields.io/github/v/release/ViperEkura/AstrAI?color=76bad9" alt="release">
|
||||||
|
<img src="https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fapi.github.com%2Frepos%2FViperEkura%2FAstrAI&query=%24.stargazers_count&label=stars&suffix=%20stars&color=76bad9" alt="stars">
|
||||||
|
<img src="https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fapi.github.com%2Frepos%2FViperEkura%2FAstrAI&query=%24.forks_count&label=forks&suffix=%20forks&color=76bad9" alt="forks">
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<br>
|
||||||
|
|
||||||
|
<div align="center">
|
||||||
|
<a href="../../README.md">English</a> •
|
||||||
|
<a href="#chinese">中文</a> •
|
||||||
|
<a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> •
|
||||||
|
<a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> •
|
||||||
|
<a href="https://huggingface.co/ViperEk">HuggingFace</a>
|
||||||
|
</div>
|
||||||
|
<br>
|
||||||
|
|
||||||
|
## 📖 目录
|
||||||
|
|
||||||
|
- [特性](#特性)
|
||||||
|
- [快速开始](#快速开始)
|
||||||
|
- [文档](#文档)
|
||||||
|
- [贡献](#贡献)
|
||||||
|
- [社区](#社区)
|
||||||
|
- [许可证](#许可证)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
<a id="chinese"></a>
|
||||||
|
## 中文
|
||||||
|
|
||||||
|
### 特性
|
||||||
|
|
||||||
|
- 🚀 **高性能**: 训练与推理双向优化,高效并行。
|
||||||
|
- 🔧 **灵活**: 支持 seq/sft/dpo/grpo 多种训练方式,可定制模型架构。
|
||||||
|
- 💡 **易用**: 简洁的 API 与丰富的示例、演示。
|
||||||
|
- 📦 **轻量**: 依赖少,部署简单。
|
||||||
|
- 🔬 **研究友好**: 模块化设计,便于实验新想法。
|
||||||
|
- 🤗 **HuggingFace 风格 API**: 类 HuggingFace 的 AutoModel/AutoTokenizer 接口,方便加载模型和分词器。
|
||||||
|
- 🔌 **双 API 兼容**: 同时支持 OpenAI 和 Anthropic 聊天补全 API,开箱即用。
|
||||||
|
|
||||||
|
### 快速开始
|
||||||
|
|
||||||
|
#### 安装
|
||||||
|
|
||||||
|
```bash
|
||||||
|
git clone https://github.com/ViperEkura/AstrAI.git
|
||||||
|
cd AstrAI
|
||||||
|
pip install -e .
|
||||||
|
```
|
||||||
|
|
||||||
|
安装开发依赖:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
pip install -e ".[dev]"
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 训练模型
|
||||||
|
|
||||||
|
```bash
|
||||||
|
CUDA_VISIBLE_DEVICES=0,1,2,3 python scripts/tools/train.py \
|
||||||
|
--train_type seq \
|
||||||
|
--data_root_path /path/to/dataset \
|
||||||
|
--param_path /path/to/model \
|
||||||
|
--batch_size 4 \
|
||||||
|
--accumulation_steps 8 \
|
||||||
|
--max_lr 3e-4 \
|
||||||
|
--warmup_steps 1000 \
|
||||||
|
--n_epoch 1
|
||||||
|
```
|
||||||
|
|
||||||
|
完整参数列表见[参数说明](./params.md)。
|
||||||
|
|
||||||
|
#### 文本生成
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/tools/generate.py \
|
||||||
|
--param_path /path/to/model \
|
||||||
|
--input_json_file /path/to/input.json \
|
||||||
|
--output_json_file /path/to/output.json
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Docker
|
||||||
|
|
||||||
|
使用 Docker 构建和运行(推荐用于 GPU 环境):
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 构建镜像
|
||||||
|
docker build -t astrai:latest .
|
||||||
|
|
||||||
|
# 启用 GPU 运行
|
||||||
|
docker run --gpus all -it astrai:latest
|
||||||
|
|
||||||
|
# 指定特定 GPU
|
||||||
|
docker run --gpus '"device=0,1"' -it astrai:latest
|
||||||
|
|
||||||
|
# 运行推理服务
|
||||||
|
docker run --gpus all -p 8000:8000 astrai:latest \
|
||||||
|
python -m scripts.tools.server --port 8000 --device cuda
|
||||||
|
|
||||||
|
# 挂载数据卷
|
||||||
|
docker run --gpus all -v /path/to/data:/data -it astrai:latest
|
||||||
|
|
||||||
|
# Docker Compose(GPU,默认)
|
||||||
|
docker compose up -d
|
||||||
|
|
||||||
|
# Docker Compose(仅 CPU)
|
||||||
|
docker compose --profile cpu up -d
|
||||||
|
```
|
||||||
|
|
||||||
|
> **注意**: 必须使用 `--gpus all` 才能启用 CUDA 支持,否则 `torch.cuda.is_available()` 将返回 `False`。
|
||||||
|
|
||||||
|
#### 启动 HTTP 服务
|
||||||
|
|
||||||
|
启动推理服务器,支持 OpenAI 和 Anthropic 兼容的 HTTP API:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python -m scripts.tools.server --port 8000 --device cuda
|
||||||
|
```
|
||||||
|
|
||||||
|
发起请求:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# OpenAI 兼容
|
||||||
|
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"messages": [{"role": "user", "content": "你好"}],
|
||||||
|
"max_tokens": 512
|
||||||
|
}'
|
||||||
|
|
||||||
|
# OpenAI 兼容流式
|
||||||
|
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"messages": [{"role": "user", "content": "讲个故事"}],
|
||||||
|
"stream": true,
|
||||||
|
"max_tokens": 500
|
||||||
|
}'
|
||||||
|
|
||||||
|
# Anthropic 兼容
|
||||||
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"model": "astrai",
|
||||||
|
"system": "你是一个乐于助人的助手。",
|
||||||
|
"messages": [{"role": "user", "content": "你好"}],
|
||||||
|
"max_tokens": 512
|
||||||
|
}'
|
||||||
|
|
||||||
|
# Anthropic 兼容流式并设置停止序列
|
||||||
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"model": "astrai",
|
||||||
|
"messages": [{"role": "user", "content": "写个故事"}],
|
||||||
|
"max_tokens": 500,
|
||||||
|
"stream": true,
|
||||||
|
"stop_sequences": ["结束"]
|
||||||
|
}'
|
||||||
|
|
||||||
|
# 健康检查
|
||||||
|
curl http://localhost:8000/health
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 演示
|
||||||
|
|
||||||
|
查看 `scripts/demo/` 文件夹中的演示:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 下载预处理数据(运行演示前必需)
|
||||||
|
python scripts/demo/download.py
|
||||||
|
|
||||||
|
# 交互式流式聊天
|
||||||
|
python scripts/demo/stream_chat.py
|
||||||
|
|
||||||
|
# 批量生成
|
||||||
|
python scripts/demo/generate_batch.py
|
||||||
|
|
||||||
|
# 自回归生成
|
||||||
|
python scripts/demo/generate_ar.py
|
||||||
|
```
|
||||||
|
|
||||||
|
观看 [bilibili](https://www.bilibili.com/video/BV1z5RPYHEkd) 上的视频演示。
|
||||||
|
|
||||||
|
### 文档
|
||||||
|
|
||||||
|
| 文档 | 说明 |
|
||||||
|
|------|------|
|
||||||
|
| [参数说明](./params.md) | 训练与推理参数配置 |
|
||||||
|
| [设计文档](./design.md) | 系统架构与模块设计 |
|
||||||
|
| [数据流程](./dataflow.md) | 数据处理管道详解 |
|
||||||
|
| [模型介绍](./introduction.md) | 模型架构与技术细节 |
|
||||||
|
|
||||||
|
### 贡献
|
||||||
|
|
||||||
|
我们欢迎贡献!请参阅[贡献指南](../../CONTRIBUTING.md)了解详情。
|
||||||
|
|
||||||
|
1. Fork 本仓库。
|
||||||
|
2. 创建功能分支。
|
||||||
|
3. 提交更改。
|
||||||
|
4. 发起 Pull Request。
|
||||||
|
|
||||||
|
重大更改请先开 issue 讨论。
|
||||||
|
|
||||||
|
### 社区
|
||||||
|
|
||||||
|
- **GitHub Issues**: [问题追踪](https://github.com/ViperEkura/AstrAI/issues)
|
||||||
|
- **Discussions**: [GitHub 讨论区](https://github.com/ViperEkura/AstrAI/discussions)
|
||||||
|
- **HuggingFace**: [模型中心](https://huggingface.co/ViperEk)
|
||||||
|
|
||||||
|
### 许可证
|
||||||
|
|
||||||
|
本项目采用 [GPL-3.0 许可证](../../LICENSE)。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
<div align="center">
|
||||||
|
<em>专为高性能与易用性设计的轻量级 Transformer 框架。</em>
|
||||||
|
</div>
|
||||||
@@ -0,0 +1,237 @@
|
|||||||
|
# AstrAI Data Flow Documentation
|
||||||
|
|
||||||
|
This document describes the data flow of the AstrAI project (a training and inference framework for autoregressive Transformer language models). It covers the complete flow from raw data to model training and inference.
|
||||||
|
|
||||||
|
## Overview
|
||||||
|
|
||||||
|
AstrAI adopts a modular design with the following main components:
|
||||||
|
- **Dataset Module** (`astrai/dataset/`): Dataset, sampler, serialization tools
|
||||||
|
- **Model Module** (`astrai/model/`): AutoModel, Transformer model and its submodules
|
||||||
|
- **Training Module** (`astrai/trainer/`): Trainer, training context, strategies, schedulers, callbacks, metric utilities
|
||||||
|
- **Inference Module** (`astrai/inference/`): Inference engine with continuous batching, streaming generation
|
||||||
|
- **Config Module** (`astrai/config/`): ModelConfig, TrainConfig
|
||||||
|
- **Factory Module** (`astrai/factory/`): Registry, BaseFactory for component registration
|
||||||
|
- **Parallel Module** (`astrai/parallel/`): Distributed training support
|
||||||
|
- **Serialization** (`astrai/serialization.py`): HDF5 data loading, checkpoint management
|
||||||
|
|
||||||
|
## Data Flow Diagram
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
flowchart LR
|
||||||
|
subgraph A[Data Preparation]
|
||||||
|
direction TB
|
||||||
|
A1[Raw Text] --> A2[AutoTokenizer]
|
||||||
|
A2 --> A3[Tokenized .h5 files]
|
||||||
|
A3 --> A4[BaseDataset]
|
||||||
|
A4 --> A5[ResumableDistributedSampler]
|
||||||
|
A5 --> A6[DataLoader]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph B[Training]
|
||||||
|
direction TB
|
||||||
|
B1[DataLoader] --> B2[BaseStrategy]
|
||||||
|
B2 --> B3[Transformer Forward]
|
||||||
|
B3 --> B4[Loss + Backward]
|
||||||
|
B4 --> B5[Gradient Accumulation]
|
||||||
|
B5 -->|every accum_steps| B6[Optimizer Step]
|
||||||
|
B6 --> B7[LR Scheduler]
|
||||||
|
B7 -->|next batch| B2
|
||||||
|
B6 --> B8[CheckpointCallback]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph C[Inference]
|
||||||
|
direction TB
|
||||||
|
C1[Checkpoint] --> C2[AutoModel]
|
||||||
|
C1 --> C3[AutoTokenizer]
|
||||||
|
C2 --> C4[InferenceEngine]
|
||||||
|
C3 --> C4
|
||||||
|
C4 --> C5[InferenceScheduler]
|
||||||
|
C5 --> C6[Transformer Forward]
|
||||||
|
C6 --> C7[sample]
|
||||||
|
C7 --> C8{End?}
|
||||||
|
C8 -->|No| C6
|
||||||
|
C8 -->|Yes| C9[Generated Text]
|
||||||
|
end
|
||||||
|
|
||||||
|
A --> B
|
||||||
|
B --> C
|
||||||
|
```
|
||||||
|
|
||||||
|
## Detailed Module Descriptions
|
||||||
|
|
||||||
|
### 1. Serialization (`astrai/serialization.py`)
|
||||||
|
|
||||||
|
- **`save_h5`**: Saves tensors by groups as HDF5 files (`.h5`), each key maps to a list of tensors
|
||||||
|
- **`load_h5`**: Loads `.h5` files, returns `Dict[str, List[Tensor]]`, supports shared memory
|
||||||
|
- **`Checkpoint`**: Encapsulates model state dict + epoch + iteration; uses safetensors
|
||||||
|
|
||||||
|
### 2. Dataset Module
|
||||||
|
|
||||||
|
#### 2.1 Dataset (`dataset.py`)
|
||||||
|
- **`BaseDataset`**: Abstract base class for windowed sequence sampling
|
||||||
|
- **`BaseSegmentFetcher` / `MultiSegmentFetcher`**: Fetch tensor segments by index range
|
||||||
|
- **`DatasetFactory`**: Creates dataset instances by `train_type` (`seq`, `sft`, `dpo`, `grpo`)
|
||||||
|
- Data keys: `"sequence"` (SEQ), `"loss_mask"` (SFT), `"chosen_mask"/"rejected_mask"` (DPO), `"masks"` (GRPO)
|
||||||
|
|
||||||
|
#### 2.2 Sampler (`sampler.py`)
|
||||||
|
- **`ResumableDistributedSampler`**: Tracks `epoch` and `iter` for breakpoint resume; supports shuffle and drop_last
|
||||||
|
|
||||||
|
### 3. Model Module
|
||||||
|
|
||||||
|
#### 3.1 Transformer / AutoModel
|
||||||
|
- **`AutoModel`**: Base class with `from_pretrained()` / `save_pretrained()`
|
||||||
|
- **`Transformer`**: Decoder-only architecture, registered via `@AutoModel.register('transformer')`
|
||||||
|
- Embedding → N×DecoderBlock → RMSNorm → Linear lm_head
|
||||||
|
- RoPE position encoding, optional weight tying
|
||||||
|
|
||||||
|
#### 3.2 Submodules (`module.py`)
|
||||||
|
- **`DecoderBlock`**: GQA attention + residual + MLP + RMSNorm
|
||||||
|
- **`GQA`**: Grouped Query Attention (also `MLA` for multi-latent attention)
|
||||||
|
- **`MLP`**: `SiLU(gate(x)) * up(x)` → down projection
|
||||||
|
- **`RotaryEmbedding`**: RoPE cos/sin cache
|
||||||
|
- **`RMSNorm`**: Layer normalization
|
||||||
|
|
||||||
|
### 4. Training Module
|
||||||
|
|
||||||
|
#### 4.1 Training Context (`train_context.py`)
|
||||||
|
- **`TrainContext`**: Dataclass holding model, optimizer, dataloader, strategy, scheduler, checkpoint state
|
||||||
|
- **`TrainContextBuilder`**: Builder pattern — takes checkpoint for resume, builds all components
|
||||||
|
|
||||||
|
#### 4.2 Trainer (`trainer.py`)
|
||||||
|
|
||||||
|
The training loop is nested: **epoch** → **batch** (with step phase interspersed):
|
||||||
|
|
||||||
|
```
|
||||||
|
on_train_begin
|
||||||
|
on_epoch_begin
|
||||||
|
for each batch:
|
||||||
|
if iteration % accumulation_steps == 0: ← step phase
|
||||||
|
on_step_begin → optimizer.step() → zero_grad → on_step_end
|
||||||
|
← batch phase
|
||||||
|
on_batch_begin → strategy(batch) → loss → backward → on_batch_end
|
||||||
|
iteration += 1
|
||||||
|
|
||||||
|
on_epoch_end
|
||||||
|
on_train_end
|
||||||
|
```
|
||||||
|
|
||||||
|
Key points:
|
||||||
|
- `on_step_*` wraps optimizer step (fires every `accumulation_steps` batches)
|
||||||
|
- `on_batch_*` wraps loss computation (fires every batch)
|
||||||
|
- `SchedulerCallback` fires on `on_batch_end` — LR scheduler steps every batch
|
||||||
|
- `GradientClippingCallback` fires on `on_step_begin`
|
||||||
|
|
||||||
|
#### 4.3 Strategy (`strategy.py`)
|
||||||
|
- **`SEQStrategy`**: Next-token prediction, cross-entropy with label smoothing
|
||||||
|
- **`SFTStrategy`**: Supervised fine-tuning with loss masking
|
||||||
|
- **`DPOStrategy`**: Direct Preference Optimization with reference model
|
||||||
|
- **`GRPOStrategy`**: Group Relative Policy Optimization with clipped ratio
|
||||||
|
|
||||||
|
#### 4.4 Scheduler (`schedule.py`)
|
||||||
|
- **`CosineScheduler`**: Cosine decay + linear warmup
|
||||||
|
- **`SGDRScheduler`**: Cosine annealing with warm restarts
|
||||||
|
- Created by `SchedulerFactory` and bound to optimizer
|
||||||
|
|
||||||
|
#### 4.5 Callbacks
|
||||||
|
- **`CheckpointCallback`**: Saves safetensors at `ckpt_interval` iterations
|
||||||
|
- **`ProgressBarCallback`**: tqdm progress display
|
||||||
|
- **`MetricLoggerCallback`**: Writes JSONL metrics to `{ckpt_dir}/logs/`
|
||||||
|
- **`GradientClippingCallback`**: `clip_grad_norm_` on `on_step_begin`
|
||||||
|
- **`SchedulerCallback`**: `scheduler.step()` on `on_batch_end`
|
||||||
|
|
||||||
|
### 5. Inference Module
|
||||||
|
|
||||||
|
#### 5.1 Inference Engine (`engine.py`)
|
||||||
|
- **`InferenceEngine`**: Facade over scheduler; provides `generate()`, `generate_with_request()`, `generate_async()`
|
||||||
|
- Accepts `prompt: str | List[str]`, returns generator (stream) or string (non-stream)
|
||||||
|
|
||||||
|
#### 5.2 Scheduler 4-Phase Loop (`scheduler.py`)
|
||||||
|
|
||||||
|
Background thread runs continuously:
|
||||||
|
|
||||||
|
```
|
||||||
|
1. Cleanup → Remove finished tasks, free KV cache pages
|
||||||
|
2. Refill → Pop from waiting_queue, alloc pages, add to active
|
||||||
|
3. Prefill → Group active tasks by prompt_len, run full forward pass
|
||||||
|
4. Decode → Pick largest same-position group, run single-token forward
|
||||||
|
```
|
||||||
|
|
||||||
|
- **`Task`**: Tracks prompt_ids, output_ids, page_table, status (PENDING/RUNNING/FINISHED/ABORTED)
|
||||||
|
- **`PagedCache`**: Bitmask-based page allocator with page-table-indirected read/write
|
||||||
|
- **`CacheView`**: Batch view bundling cache + page table for attention layers
|
||||||
|
- **`sample()`**: Temperature → top-k → top-p → multinomial
|
||||||
|
|
||||||
|
#### 5.3 Server (`server.py`)
|
||||||
|
- FastAPI with OpenAI `/v1/chat/completions` and Anthropic `/v1/messages` endpoints
|
||||||
|
- Streaming via SSE, health check at `/health`, stats at `/stats`
|
||||||
|
|
||||||
|
### 6. Tokenizer Module
|
||||||
|
|
||||||
|
- **`AutoTokenizer`**: Wraps HuggingFace tokenizers (BBPE); `encode`/`decode`/`apply_chat_template`
|
||||||
|
- **`ChatTemplate`**: Jinja2-based template rendering for multi-turn chat
|
||||||
|
|
||||||
|
### 7. Factory & Parallel
|
||||||
|
|
||||||
|
- **`Registry` / `BaseFactory`**: Decorator-based component registration
|
||||||
|
- **`spawn_parallel_fn`**: Multi-process DDP launcher with NCCL backend
|
||||||
|
- **`ParallelModel` / `ColumnParallelLinear` / `RowParallelLinear`**: Tensor model parallelism
|
||||||
|
|
||||||
|
## Training Data Flow — Detailed Steps
|
||||||
|
|
||||||
|
1. **Data Preparation**
|
||||||
|
- Raw text → token IDs via `AutoTokenizer.encode()`
|
||||||
|
- Save as `.h5` files (groups of tensor lists per data key)
|
||||||
|
|
||||||
|
2. **Dataset Loading**
|
||||||
|
- `BaseDataset.load()` calls `load_h5()`, builds `MultiSegmentFetcher`
|
||||||
|
- Sliding window of `window_size` with `stride` determines sample boundaries
|
||||||
|
|
||||||
|
3. **Sampling & Batching**
|
||||||
|
- `ResumableDistributedSampler` produces shuffled index sequences
|
||||||
|
- `DataLoader` fetches `[batch_size, window_size]` tensors via `__getitem__`
|
||||||
|
|
||||||
|
4. **Strategy Forward**
|
||||||
|
- Strategy receives batch, calls `Transformer.forward()` for logits
|
||||||
|
- Computes task-specific loss (cross-entropy, DPO, GRPO)
|
||||||
|
|
||||||
|
5. **Backward & Accumulation**
|
||||||
|
- `loss = raw_loss / accumulation_steps`
|
||||||
|
- `loss.backward()` accumulates gradients
|
||||||
|
- Every `accumulation_steps` batches: `optimizer.step()` → `zero_grad()`
|
||||||
|
- Every batch: `scheduler.step()` updates learning rate
|
||||||
|
|
||||||
|
6. **Checkpoint**
|
||||||
|
- `CheckpointCallback` saves `model.state_dict()` + metadata to safetensors at `ckpt_interval` iterations
|
||||||
|
- Does NOT save optimizer/scheduler state (resume resets those)
|
||||||
|
|
||||||
|
## Inference Data Flow — Detailed Steps
|
||||||
|
|
||||||
|
1. **Model Loading**
|
||||||
|
- `AutoModel.from_pretrained(path)` loads weights from safetensors
|
||||||
|
- `torch.inference_mode()` wraps generation
|
||||||
|
|
||||||
|
2. **Prompt Construction**
|
||||||
|
- Messages → `apply_chat_template(messages, tokenize=False)` → prompt string
|
||||||
|
- `tokenizer.encode(prompt)` → token IDs (truncated to `max_prompt_len`)
|
||||||
|
|
||||||
|
3. **Continuous Batching Loop**
|
||||||
|
- **Cleanup**: Finished tasks → `stream_callback(STOP)`, free KV pages
|
||||||
|
- **Refill**: Pop from waiting queue, `PagedCache.alloc_n()` for prompt pages
|
||||||
|
- **Prefill**: Group by prompt length, run full forward with `start_pos=0`
|
||||||
|
- **Decode**: Pick position group with most tasks, single-token forward:
|
||||||
|
- Model forward → `logits` → `sample()` → next token ID
|
||||||
|
- Append to `output_ids`, update `output_tokens`
|
||||||
|
- `_maybe_alloc_page()` grows page table as needed
|
||||||
|
- `stream_callback(token)` for streaming clients
|
||||||
|
|
||||||
|
4. **Output**
|
||||||
|
- `tokenizer.decode(output_ids)` → text
|
||||||
|
- Return to caller (streaming: token-by-token; non-streaming: complete string)
|
||||||
|
|
||||||
|
## Checkpoint & Serialization
|
||||||
|
|
||||||
|
- **Training Checkpoint**: safetensors weights + epoch/iteration metadata. Optimizer/scheduler state is NOT persisted.
|
||||||
|
- **Inference Loading**: `AutoModel.from_pretrained()` loads from the same safetensors format.
|
||||||
|
- **Dataset Serialization**: HDF5 with shared memory support for large-scale pre-training data.
|
||||||
|
|
||||||
|
> Document Update Time: 2026-05-09
|
||||||
@@ -0,0 +1,719 @@
|
|||||||
|
## 1. Why I Created This Project
|
||||||
|
|
||||||
|
There are many large language models on the market today, such as GPT, LLaMA, and others, with tens of billions or even hundreds of billions of parameters. But honestly, these models have extremely high hardware requirements, making them inaccessible for ordinary developers. I thought: **Can we create a model that is both useful and can run on ordinary computers?** This is also what most people currently hope for - a locally deployable AI project that achieves complete privatization while maintaining some level of intelligence.
|
||||||
|
|
||||||
|
Thus, the AstrAI project was born - 1B parameters, Chinese-English bilingual, supporting dialogue, text generation, and the training code is open source!
|
||||||
|
|
||||||
|
## 2. System Architecture
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
classDiagram
|
||||||
|
namespace config {
|
||||||
|
class ModelConfig {
|
||||||
|
+int vocab_size
|
||||||
|
+int dim
|
||||||
|
+int n_layers
|
||||||
|
+float norm_eps
|
||||||
|
+int dim_ffn
|
||||||
|
+bool tie_weight
|
||||||
|
+int max_len
|
||||||
|
+float rope_theta
|
||||||
|
+int n_heads
|
||||||
|
+int n_kv_heads
|
||||||
|
+bool use_qk_norm
|
||||||
|
+bool use_gated_attention
|
||||||
|
+load(config_path) ModelConfig
|
||||||
|
+save(config_path)
|
||||||
|
}
|
||||||
|
|
||||||
|
class TrainConfig {
|
||||||
|
+nn.Module model
|
||||||
|
+str strategy
|
||||||
|
+Dataset dataset
|
||||||
|
+Callable optimizer_fn
|
||||||
|
+Callable scheduler_fn
|
||||||
|
+int n_epoch
|
||||||
|
+int batch_size
|
||||||
|
+int accumulation_steps
|
||||||
|
+float max_grad_norm
|
||||||
|
+int start_epoch
|
||||||
|
+int start_batch
|
||||||
|
+str ckpt_dir
|
||||||
|
+int ckpt_interval
|
||||||
|
+int random_seed
|
||||||
|
+int num_workers
|
||||||
|
+int prefetch_factor
|
||||||
|
+bool pin_memory
|
||||||
|
+int nprocs
|
||||||
|
+str backend
|
||||||
|
+str master_addr
|
||||||
|
+str master_port
|
||||||
|
+Callable parallel_wrapper
|
||||||
|
+Callable state_dict_fn
|
||||||
|
+str device_type
|
||||||
|
+dict extra_kwargs
|
||||||
|
+validate()
|
||||||
|
}
|
||||||
|
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace dataset {
|
||||||
|
class BaseDataset {
|
||||||
|
+int window_size
|
||||||
|
+int stride
|
||||||
|
+MultiSegmentFetcher fetcher
|
||||||
|
+load(load_path)
|
||||||
|
+__getitem__(index)
|
||||||
|
+__len__()
|
||||||
|
}
|
||||||
|
|
||||||
|
class SEQDataset {
|
||||||
|
+__getitem__(index) Dict
|
||||||
|
}
|
||||||
|
|
||||||
|
class SFTDataset {
|
||||||
|
+__getitem__(index) Dict
|
||||||
|
}
|
||||||
|
|
||||||
|
class DPODataset {
|
||||||
|
+__getitem__(index) Dict
|
||||||
|
}
|
||||||
|
|
||||||
|
class GRPODataset {
|
||||||
|
+__getitem__(index) Dict
|
||||||
|
}
|
||||||
|
|
||||||
|
class BaseSegmentFetcher {
|
||||||
|
+List[Tensor] segments
|
||||||
|
+List[int] cum_lengths
|
||||||
|
+int total_length
|
||||||
|
+fetch_data(begin_idx, end_idx) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class MultiSegmentFetcher {
|
||||||
|
+Dict multi_fetchers
|
||||||
|
+List multi_keys
|
||||||
|
+key_fetch(begin_idx, end_idx, keys) Dict
|
||||||
|
+fetch_data(begin_idx, end_idx) Dict
|
||||||
|
}
|
||||||
|
|
||||||
|
class ResumableDistributedSampler {
|
||||||
|
+int epoch
|
||||||
|
+int iter
|
||||||
|
}
|
||||||
|
|
||||||
|
class DatasetFactory {
|
||||||
|
+Registry _registry
|
||||||
|
+register(name) decorator
|
||||||
|
+create(train_type, window_size, stride) BaseDataset
|
||||||
|
+load(train_type, load_path, window_size, stride) BaseDataset
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace serialization {
|
||||||
|
class Checkpoint {
|
||||||
|
+dict state_dict
|
||||||
|
+int epoch
|
||||||
|
+int iteration
|
||||||
|
+save(save_dir)
|
||||||
|
+load(save_dir) Checkpoint
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace model {
|
||||||
|
class AutoModel {
|
||||||
|
+ModelConfig config
|
||||||
|
+Registry _registry
|
||||||
|
+register(model_type) decorator
|
||||||
|
+get_model_class(model_type) Type
|
||||||
|
+from_pretrained(path, disable_random_init) nn.Module
|
||||||
|
+save_pretrained(save_directory)
|
||||||
|
+to(*args, **kwargs) Self
|
||||||
|
}
|
||||||
|
|
||||||
|
class Transformer {
|
||||||
|
+ModelConfig config
|
||||||
|
+RotaryEmbedding rotary_embedding
|
||||||
|
+Embedding embed_tokens
|
||||||
|
+ModuleList layers
|
||||||
|
+RMSNorm norm
|
||||||
|
+Linear lm_head
|
||||||
|
+forward(input_ids, input_mask, paged_cache, start_pos) Dict
|
||||||
|
+load_state_dict(state_dict)
|
||||||
|
+state_dict()
|
||||||
|
}
|
||||||
|
|
||||||
|
class DecoderBlock {
|
||||||
|
+GQA attention
|
||||||
|
+RMSNorm input_norm
|
||||||
|
+MLP mlp
|
||||||
|
+RMSNorm post_attention_norm
|
||||||
|
+forward(x, rotary_emb, attention_mask, paged_cache, start_pos) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class GQA {
|
||||||
|
+int n_heads
|
||||||
|
+int n_kv_heads
|
||||||
|
+int head_dim
|
||||||
|
+Linear q_proj, k_proj, v_proj, o_proj
|
||||||
|
+RMSNorm q_norm, k_norm
|
||||||
|
+forward(x, rotary_emb, mask, paged_cache, start_pos) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class MLA {
|
||||||
|
+int n_heads
|
||||||
|
+int n_kv_heads
|
||||||
|
+int head_dim
|
||||||
|
+int kv_lora_rank
|
||||||
|
+int qk_nope_head_dim
|
||||||
|
+int qk_rope_head_dim
|
||||||
|
+Linear q_proj, kv_a_proj, kv_b_proj
|
||||||
|
+Linear o_proj
|
||||||
|
+RMSNorm kv_norm
|
||||||
|
+forward(x, rotary_emb, mask, paged_cache, start_pos) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class MLP {
|
||||||
|
+Linear up, gate, down
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class RMSNorm {
|
||||||
|
+Parameter weight
|
||||||
|
+float norm_eps
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class Linear {
|
||||||
|
+Parameter weight
|
||||||
|
+Parameter bias
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class RotaryEmbedding {
|
||||||
|
+int dim
|
||||||
|
+int max_len
|
||||||
|
+float base
|
||||||
|
+forward(x, start_pos) Tuple[Tensor, Tensor]
|
||||||
|
}
|
||||||
|
|
||||||
|
class Embedding {
|
||||||
|
+Parameter weight
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace tokenize {
|
||||||
|
class AutoTokenizer {
|
||||||
|
+List[int] stop_ids
|
||||||
|
+int bos_id
|
||||||
|
+int eos_id
|
||||||
|
+int pad_id
|
||||||
|
+vocab_size int
|
||||||
|
+encode(tokens, out_ids, add_special_tokens) List[int]
|
||||||
|
+decode(tokens, skip_special_tokens) str
|
||||||
|
+apply_chat_template(messages, tokenize) Union[str, List[int]]
|
||||||
|
+set_chat_template(template)
|
||||||
|
+load(path)
|
||||||
|
+from_pretrained(path) AutoTokenizer
|
||||||
|
+save_pretrained(save_path)
|
||||||
|
}
|
||||||
|
|
||||||
|
class ChatTemplate {
|
||||||
|
+String template_str
|
||||||
|
+render(messages, system_prompt, **extra_variables) str
|
||||||
|
+from_string(template) ChatTemplate
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace factory {
|
||||||
|
class Registry {
|
||||||
|
+Dict _entries
|
||||||
|
+register(name, component_cls, category, priority)
|
||||||
|
+get(name) Type
|
||||||
|
+list_names() List[str]
|
||||||
|
}
|
||||||
|
|
||||||
|
class BaseFactory {
|
||||||
|
+Registry _registry
|
||||||
|
+register(name, category, priority) decorator
|
||||||
|
+create(name, *args, **kwargs) T
|
||||||
|
+list_registered() list
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace trainer {
|
||||||
|
class Trainer {
|
||||||
|
+TrainConfig train_config
|
||||||
|
+List[TrainCallback] callbacks
|
||||||
|
+train(checkpoint)
|
||||||
|
+_build_context(checkpoint) TrainContext
|
||||||
|
+_get_default_callbacks() List[TrainCallback]
|
||||||
|
}
|
||||||
|
|
||||||
|
class TrainContext {
|
||||||
|
+nn.Module model
|
||||||
|
+BaseStrategy strategy
|
||||||
|
+DataLoader dataloader
|
||||||
|
+Optimizer optimizer
|
||||||
|
+LRScheduler scheduler
|
||||||
|
+Checkpoint checkpoint
|
||||||
|
+int epoch
|
||||||
|
+int iteration
|
||||||
|
+float loss
|
||||||
|
+int world_size
|
||||||
|
+int rank
|
||||||
|
}
|
||||||
|
|
||||||
|
class TrainContextBuilder {
|
||||||
|
+TrainConfig config
|
||||||
|
+with_checkpoint(checkpoint) TrainContextBuilder
|
||||||
|
+build() TrainContext
|
||||||
|
}
|
||||||
|
|
||||||
|
class BaseStrategy {
|
||||||
|
+nn.Module model
|
||||||
|
+str device
|
||||||
|
+compute_loss(batch) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class StrategyFactory {
|
||||||
|
+Registry _registry
|
||||||
|
+register(name) decorator
|
||||||
|
+create(model, train_type, device, **kwargs) BaseStrategy
|
||||||
|
}
|
||||||
|
|
||||||
|
class SEQStrategy {
|
||||||
|
+float label_smoothing
|
||||||
|
+compute_loss(batch) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class SFTStrategy {
|
||||||
|
+float label_smoothing
|
||||||
|
+compute_loss(batch) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class DPOStrategy {
|
||||||
|
+nn.Module ref_model
|
||||||
|
+float beta
|
||||||
|
+str reduction
|
||||||
|
+compute_loss(batch) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class GRPOStrategy {
|
||||||
|
+nn.Module ref_model
|
||||||
|
+float clip_eps
|
||||||
|
+float kl_coef
|
||||||
|
+int group_size
|
||||||
|
+compute_loss(batch) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class BaseScheduler {
|
||||||
|
+get_lr() List[float]
|
||||||
|
+step()
|
||||||
|
}
|
||||||
|
|
||||||
|
class SchedulerFactory {
|
||||||
|
+Registry _registry
|
||||||
|
+register(name) decorator
|
||||||
|
+create(optimizer, schedule_type, **kwargs) BaseScheduler
|
||||||
|
}
|
||||||
|
|
||||||
|
class CosineScheduler {
|
||||||
|
+int warmup_steps
|
||||||
|
+int lr_decay_steps
|
||||||
|
+float min_rate
|
||||||
|
}
|
||||||
|
|
||||||
|
class SGDRScheduler {
|
||||||
|
+int warmup_steps
|
||||||
|
+int cycle_length
|
||||||
|
+float min_rate
|
||||||
|
+int t_mult
|
||||||
|
}
|
||||||
|
|
||||||
|
class TrainCallback {
|
||||||
|
+on_train_begin(context)
|
||||||
|
+on_train_end(context)
|
||||||
|
+on_epoch_begin(context)
|
||||||
|
+on_epoch_end(context)
|
||||||
|
+on_step_begin(context)
|
||||||
|
+on_step_end(context)
|
||||||
|
+on_batch_begin(context)
|
||||||
|
+on_batch_end(context)
|
||||||
|
+on_error(context)
|
||||||
|
}
|
||||||
|
|
||||||
|
class GradientClippingCallback {
|
||||||
|
+float max_grad_norm
|
||||||
|
+on_step_begin(context)
|
||||||
|
}
|
||||||
|
|
||||||
|
class SchedulerCallback {
|
||||||
|
+on_train_begin(context)
|
||||||
|
+on_batch_end(context)
|
||||||
|
}
|
||||||
|
|
||||||
|
class CheckpointCallback {
|
||||||
|
+str save_dir
|
||||||
|
+int interval
|
||||||
|
+_save_checkpoint(context)
|
||||||
|
+on_batch_end(context)
|
||||||
|
+on_train_end(context)
|
||||||
|
+on_error(context)
|
||||||
|
}
|
||||||
|
|
||||||
|
class ProgressBarCallback {
|
||||||
|
+int num_epoch
|
||||||
|
+on_epoch_begin(context)
|
||||||
|
+on_batch_end(context)
|
||||||
|
+on_epoch_end(context)
|
||||||
|
}
|
||||||
|
|
||||||
|
class MetricLoggerCallback {
|
||||||
|
+str log_dir
|
||||||
|
+int save_interval
|
||||||
|
+on_batch_end(context)
|
||||||
|
+on_train_end(context)
|
||||||
|
}
|
||||||
|
|
||||||
|
class CallbackFactory {
|
||||||
|
+Registry _registry
|
||||||
|
+register(name) decorator
|
||||||
|
+create(name, **kwargs) TrainCallback
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace inference {
|
||||||
|
class InferenceEngine {
|
||||||
|
+nn.Module model
|
||||||
|
+AutoTokenizer tokenizer
|
||||||
|
+InferenceScheduler scheduler
|
||||||
|
+int max_batch_size
|
||||||
|
+Optional int max_seq_len
|
||||||
|
+generate(prompt, stream, max_tokens, temperature, top_p, top_k) Union[Generator, str, List[str]]
|
||||||
|
+generate_with_request(request) Union[Generator, str, List[str]]
|
||||||
|
+generate_async(prompt, max_tokens, temperature, top_p, top_k) AsyncGenerator
|
||||||
|
+get_stats() Dict
|
||||||
|
+shutdown()
|
||||||
|
}
|
||||||
|
|
||||||
|
class InferenceScheduler {
|
||||||
|
+nn.Module model
|
||||||
|
+AutoTokenizer tokenizer
|
||||||
|
+PagedCache page_cache
|
||||||
|
+int max_batch_size
|
||||||
|
+int max_seq_len
|
||||||
|
+int max_prompt_len
|
||||||
|
+int page_size
|
||||||
|
+List waiting_queue
|
||||||
|
+List active_tasks
|
||||||
|
+add_task(prompt, max_tokens, temperature, top_p, top_k, stream_callback) str
|
||||||
|
+remove_task(task_id)
|
||||||
|
+start()
|
||||||
|
+stop()
|
||||||
|
+get_stats() Dict
|
||||||
|
}
|
||||||
|
|
||||||
|
class PagedCache {
|
||||||
|
+int page_size
|
||||||
|
+int _free_mask
|
||||||
|
+List[int] _refs
|
||||||
|
+Tensor k_cache
|
||||||
|
+Tensor v_cache
|
||||||
|
+alloc() int
|
||||||
|
+alloc_n(n) List[int]
|
||||||
|
+free(idx)
|
||||||
|
+bind(page_table, total_len) CacheView
|
||||||
|
+write(layer_id, page_table, start_pos, k, v)
|
||||||
|
+gather(layer_id, page_table) Tuple[Tensor, Tensor]
|
||||||
|
}
|
||||||
|
|
||||||
|
class CacheView {
|
||||||
|
+PagedCache _cache
|
||||||
|
+Tensor _page_table
|
||||||
|
+int _total_len
|
||||||
|
+write(layer_id, start_pos, k, v)
|
||||||
|
+gather(layer_id) Tuple[Tensor, Tensor]
|
||||||
|
}
|
||||||
|
|
||||||
|
class Task {
|
||||||
|
+str task_id
|
||||||
|
+List prompt_ids
|
||||||
|
+int max_tokens
|
||||||
|
+float temperature
|
||||||
|
+float top_p
|
||||||
|
+int top_k
|
||||||
|
+TaskStatus status
|
||||||
|
+List output_ids
|
||||||
|
+int input_tokens
|
||||||
|
+int output_tokens
|
||||||
|
+List[int] page_table
|
||||||
|
+int n_pages
|
||||||
|
+float arrival_time
|
||||||
|
+float finish_time
|
||||||
|
+Callable stream_callback
|
||||||
|
+int next_pos
|
||||||
|
+is_finished(stop_ids) bool
|
||||||
|
}
|
||||||
|
|
||||||
|
class TaskStatus {
|
||||||
|
<<enumeration>>
|
||||||
|
PENDING
|
||||||
|
RUNNING
|
||||||
|
FINISHED
|
||||||
|
ABORTED
|
||||||
|
}
|
||||||
|
|
||||||
|
class GenerationRequest {
|
||||||
|
+List[Dict] messages
|
||||||
|
+GenerationParams params
|
||||||
|
+bool stream
|
||||||
|
}
|
||||||
|
|
||||||
|
class GenerationParams {
|
||||||
|
<<value object>>
|
||||||
|
+int top_k
|
||||||
|
+float top_p
|
||||||
|
+float temperature
|
||||||
|
+int max_tokens
|
||||||
|
}
|
||||||
|
|
||||||
|
class BaseSamplingStrategy {
|
||||||
|
<<abstract>>
|
||||||
|
+apply(logits, filter_value) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class TemperatureStrategy {
|
||||||
|
+float temperature
|
||||||
|
+apply(logits, filter_value) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class TopKStrategy {
|
||||||
|
+int top_k
|
||||||
|
+apply(logits, filter_value) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class TopPStrategy {
|
||||||
|
+float top_p
|
||||||
|
+apply(logits, filter_value) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class SamplingPipeline {
|
||||||
|
+List strategies
|
||||||
|
+apply(logits, filter_value) Tensor
|
||||||
|
+sample(logits, filter_value) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class _Result {
|
||||||
|
+List[str] tokens
|
||||||
|
+List[str] results
|
||||||
|
+List[bool] _done
|
||||||
|
+append(token, idx)
|
||||||
|
+get_results() List[str]
|
||||||
|
+pop_all() List[str]
|
||||||
|
+wait(timeout) bool
|
||||||
|
}
|
||||||
|
|
||||||
|
class ChatMessage {
|
||||||
|
+str role
|
||||||
|
+str content
|
||||||
|
}
|
||||||
|
|
||||||
|
class ChatCompletionRequest {
|
||||||
|
+List[ChatMessage] messages
|
||||||
|
+float temperature
|
||||||
|
+float top_p
|
||||||
|
+int top_k
|
||||||
|
+int max_tokens
|
||||||
|
+bool stream
|
||||||
|
+Optional[str] stop
|
||||||
|
+Optional[int] n
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace parallel {
|
||||||
|
class ParallelFunctions {
|
||||||
|
+spawn_parallel_fn(fn, nprocs)
|
||||||
|
+setup_parallel(rank, world_size, backend, master_addr, master_port, device_type)
|
||||||
|
}
|
||||||
|
|
||||||
|
class ParallelModel {
|
||||||
|
+dist.ProcessGroup process_group
|
||||||
|
+int rank
|
||||||
|
+int world_size
|
||||||
|
}
|
||||||
|
|
||||||
|
class ColumnParallelLinear {
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class RowParallelLinear {
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
%% Relationships
|
||||||
|
TrainConfig --> ModelConfig : uses
|
||||||
|
TrainConfig --> BaseDataset : uses
|
||||||
|
TrainConfig --> StrategyFactory : selects
|
||||||
|
StrategyFactory ..> BaseStrategy : creates
|
||||||
|
BaseStrategy <|-- SEQStrategy
|
||||||
|
BaseStrategy <|-- SFTStrategy
|
||||||
|
BaseStrategy <|-- DPOStrategy
|
||||||
|
BaseStrategy <|-- GRPOStrategy
|
||||||
|
DPOStrategy --> Transformer : uses
|
||||||
|
GRPOStrategy --> Transformer : uses
|
||||||
|
Trainer --> TrainConfig : configures
|
||||||
|
Trainer --> TrainContextBuilder : builds
|
||||||
|
Trainer --> TrainCallback : manages
|
||||||
|
TrainContextBuilder --> TrainContext : creates
|
||||||
|
Checkpoint ..> Checkpoint : saves/loads
|
||||||
|
TrainContext --> Checkpoint : manages
|
||||||
|
TrainContext --> BaseStrategy : uses
|
||||||
|
TrainContext --> BaseScheduler : uses
|
||||||
|
SchedulerFactory ..> BaseScheduler : creates
|
||||||
|
BaseScheduler <|-- CosineScheduler
|
||||||
|
BaseScheduler <|-- SGDRScheduler
|
||||||
|
CallbackFactory ..> TrainCallback : creates
|
||||||
|
TrainCallback <|-- GradientClippingCallback
|
||||||
|
TrainCallback <|-- SchedulerCallback
|
||||||
|
TrainCallback <|-- CheckpointCallback
|
||||||
|
TrainCallback <|-- ProgressBarCallback
|
||||||
|
TrainCallback <|-- MetricLoggerCallback
|
||||||
|
InferenceEngine --> InferenceScheduler : uses
|
||||||
|
InferenceEngine --> GenerationRequest : uses
|
||||||
|
GenerationRequest --> GenerationParams : contains
|
||||||
|
InferenceScheduler --> Task : manages
|
||||||
|
Task --> TaskStatus : uses
|
||||||
|
InferenceScheduler --> TaskStatus : uses
|
||||||
|
InferenceScheduler --> PagedCache : uses
|
||||||
|
InferenceScheduler --> Transformer : uses
|
||||||
|
InferenceEngine --> Transformer : uses
|
||||||
|
InferenceEngine --> _Result : uses
|
||||||
|
BaseSamplingStrategy <|-- TemperatureStrategy
|
||||||
|
BaseSamplingStrategy <|-- TopKStrategy
|
||||||
|
BaseSamplingStrategy <|-- TopPStrategy
|
||||||
|
SamplingPipeline --> BaseSamplingStrategy : composes
|
||||||
|
BaseDataset <|-- SEQDataset
|
||||||
|
BaseDataset <|-- SFTDataset
|
||||||
|
BaseDataset <|-- DPODataset
|
||||||
|
BaseDataset <|-- GRPODataset
|
||||||
|
DatasetFactory ..> BaseDataset : creates
|
||||||
|
MultiSegmentFetcher --> BaseSegmentFetcher : uses
|
||||||
|
BaseDataset --> MultiSegmentFetcher : uses
|
||||||
|
AutoModel <|-- Transformer
|
||||||
|
AutoModel --> ModelConfig : contains
|
||||||
|
Transformer --> DecoderBlock : uses
|
||||||
|
Transformer --> RotaryEmbedding : uses
|
||||||
|
Transformer --> Embedding : uses
|
||||||
|
DecoderBlock --> GQA : uses
|
||||||
|
DecoderBlock --> MLP : uses
|
||||||
|
DecoderBlock --> RMSNorm : uses
|
||||||
|
TrainContextBuilder --> ResumableDistributedSampler : creates
|
||||||
|
ResumableDistributedSampler --> BaseDataset : samples
|
||||||
|
ParallelModel <|-- RowParallelLinear
|
||||||
|
ParallelModel <|-- ColumnParallelLinear
|
||||||
|
AutoTokenizer --> ChatTemplate : uses
|
||||||
|
TrainConfig --> DatasetFactory : selects
|
||||||
|
TrainConfig --> SchedulerFactory : selects
|
||||||
|
TrainConfig --> CallbackFactory : selects
|
||||||
|
AutoModel ..> AutoTokenizer : loads with
|
||||||
|
BaseFactory <|-- DatasetFactory
|
||||||
|
BaseFactory <|-- StrategyFactory
|
||||||
|
BaseFactory <|-- SchedulerFactory
|
||||||
|
BaseFactory <|-- CallbackFactory
|
||||||
|
```
|
||||||
|
|
||||||
|
### Module Overview
|
||||||
|
|
||||||
|
| Module | Components | Description |
|
||||||
|
|--------|------------|-------------|
|
||||||
|
| **astrai.config** | ModelConfig, TrainConfig | Configuration management |
|
||||||
|
| **astrai.dataset** | BaseDataset, SEQDataset, SFTDataset, DPODataset, GRPODataset, BaseSegmentFetcher, MultiSegmentFetcher, ResumableDistributedSampler, DatasetFactory | Dataset loading and management |
|
||||||
|
| **astrai.serialization** | Checkpoint, save_h5, load_h5 | Model serialization and checkpoint management |
|
||||||
|
| **astrai.model** | AutoModel, Transformer, DecoderBlock, GQA, MLA, MLP, RMSNorm, Linear, RotaryEmbedding, Embedding | Neural network model |
|
||||||
|
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
||||||
|
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy, StrategyFactory, BaseScheduler, SchedulerFactory, TrainCallback, CallbackFactory | Training workflow management |
|
||||||
|
| **astrai.inference** | InferenceEngine, InferenceScheduler, PagedCache, CacheView, Task, TaskStatus, GenerationParams, GenerationRequest, BaseSamplingStrategy, TemperatureStrategy, TopKStrategy, TopPStrategy, SamplingPipeline, ChatMessage, ChatCompletionRequest | Inference service with continuous batching and paged KV cache |
|
||||||
|
| **astrai.parallel** | ParallelFunctions, ParallelModel, ColumnParallelLinear, RowParallelLinear | Distributed parallel |
|
||||||
|
| **astrai.factory** | Registry, BaseFactory | Generic component registration |
|
||||||
|
|
||||||
|
### Design Patterns
|
||||||
|
|
||||||
|
| Pattern | Classes | Purpose |
|
||||||
|
|---------|---------|---------|
|
||||||
|
| **Strategy** | `BaseStrategy`, `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy`, `StrategyFactory` | Flexible training strategy switching, supports SEQ/SFT/DPO/GRPO |
|
||||||
|
| **Builder** | `TrainContextBuilder` | Chain-building training context, step-by-step initialization of components |
|
||||||
|
| **Factory** | `StrategyFactory`, `SchedulerFactory`, `DatasetFactory`, `CallbackFactory`, `BaseFactory` | Decorator registration mechanism, dynamically create training strategies, schedulers, datasets, and callbacks |
|
||||||
|
| **Observer** | `TrainCallback`, `CallbackFactory` | Callback mechanism for training process monitoring (checkpoint, early stopping, metrics) |
|
||||||
|
| **Context** | `TrainContext` | Training process state container with model, optimizer, scheduler and checkpoint |
|
||||||
|
| **Registry** | `BaseFactory`, `Registry` | Generic component registration with category and priority support |
|
||||||
|
| **Object Pool** | `PagedCache` | Page-based KV cache with O(1) alloc/free via bitmask |
|
||||||
|
| **Strategy (Sampling)** | `BaseSamplingStrategy`, `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations with temperature, top-k, top-p |
|
||||||
|
| **Producer-Consumer** | `InferenceScheduler`, `Task`, `waiting_queue`, `active_tasks` | Continuous batching with dynamic task queue management |
|
||||||
|
| **Event-Driven** | `threading.Event`, `_task_event` | Non-blocking wait mechanism for task scheduling using Python's `threading` module |
|
||||||
|
| **AutoModel Registry** | `AutoModel`, `Transformer` | Model type registration and dynamic loading via decorator pattern |
|
||||||
|
| **Generator Pattern** | `_Result`, `GenerationRequest` | Event-based result notification for streaming/non-streaming generation |
|
||||||
|
|
||||||
|
### Core Relationships
|
||||||
|
|
||||||
|
1. **Configuration → Training**: `TrainConfig` contains `ModelConfig`, holds model, dataset, optimizer and other references
|
||||||
|
2. **Training Flow**: `Trainer` → `TrainContextBuilder` → `TrainContext`, uses `BaseStrategy` to compute loss
|
||||||
|
3. **Strategy Selection**: `StrategyFactory` creates corresponding strategy instance based on `train_type`
|
||||||
|
4. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `Transformer`, uses `PagedCache` for paged KV cache management and `SamplingPipeline` for efficient continuous batching with streaming/non-streaming
|
||||||
|
5. **Distributed Support**: `spawn_parallel_fn` and `setup_parallel` provide multi-process training capability for `Trainer`
|
||||||
|
6. **Dataset Loading**: `DatasetFactory` creates datasets (SEQDataset, SFTDataset, DPODataset, GRPODataset), supports HDF5 loading via `BaseSegmentFetcher` and `MultiSegmentFetcher`
|
||||||
|
7. **Checkpoint Management**: `Checkpoint` handles model state serialization/deserialization with safetensors
|
||||||
|
8. **Scheduler Support**: `SchedulerFactory` creates learning rate schedulers (CosineScheduler, SGDRScheduler)
|
||||||
|
9. **AutoModel Loading**: `AutoModel.from_pretrained()` dynamically loads model based on `config.json` model_type, uses `Registry` pattern for model type registration
|
||||||
|
|
||||||
|
## 3. Training Process
|
||||||
|
|
||||||
|
The common training process for large language models (LLM) typically includes three stages: **Pre-training (SEQ)**, **Supervised Fine-Tuning (SFT)**, and **Reinforcement Learning from Human Feedback (DPO/GRPO)**. This system is designed to support seamless end-to-end flow, achieving efficient switching and state management of different training stages through modular strategies.
|
||||||
|
|
||||||
|
### Core Formulas
|
||||||
|
|
||||||
|
**Pre-training (SEQ):**
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{\text{PT}} = - \sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta)
|
||||||
|
$$
|
||||||
|
|
||||||
|
**SFT:**
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{\text{SFT}} = - \sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
|
||||||
|
$$
|
||||||
|
|
||||||
|
**DPO:**
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{\text{DPO}} = -\mathbb{E}_{(x, y_w, y_l) \sim D} \left[ \log \sigma\left( \beta \log \frac{\pi_\theta(y_w \mid x)}{\pi_{\text{ref}}(y_w \mid x)} - \beta \log \frac{\pi_\theta(y_l \mid x)}{\pi_{\text{ref}}(y_l \mid x)} \right) \right]
|
||||||
|
$$
|
||||||
|
|
||||||
|
**GRPO:**
|
||||||
|
|
||||||
|
GRPO (Group Relative Policy Optimization) computes advantages from multiple responses to the same prompt, then optimizes using a PPO-style clipped objective:
|
||||||
|
|
||||||
|
$$
|
||||||
|
\text{Advantage}_i = \frac{r_i - \mu}{\sigma + \epsilon}
|
||||||
|
$$
|
||||||
|
|
||||||
|
Where $r_i$ is the reward for the $i$-th response, $\mu$ and $\sigma$ are the mean and standard deviation of group rewards.
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{\text{GRPO}} = -\mathbb{E} \left[ \min\left( \frac{\pi_\theta(a|s)}{\pi_{\text{ref}}(a|s)} \cdot A, \text{clip}\left(\frac{\pi_\theta(a|s)}{\pi_{\text{ref}}(a|s)}, 1-\epsilon, 1+\epsilon\right) \cdot A \right) \right] + \lambda \cdot D_{KL}
|
||||||
|
$$
|
||||||
|
|
||||||
|
The KL divergence term uses mean squared error approximation:
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{KL} = \lambda \cdot \mathbb{E} \left[ (\log \pi_\theta - \log \pi_{\text{ref}})^2 \right]
|
||||||
|
$$
|
||||||
|
|
||||||
|
The final loss is the sum of both: $L = L_{\text{policy}} + L_{KL}$
|
||||||
|
|
||||||
|
Through the above three-stage progressive training, the model completes its evolution from a general language foundation to a specialized, highly-aligned dialogue intelligence.
|
||||||
|
|
||||||
|
> Document Update Time: 2026-04-09
|
||||||
+289
-44
@@ -1,50 +1,83 @@
|
|||||||
## 模型介绍
|
## Model Introduction
|
||||||
|
|
||||||
|
### 1. Model Architecture
|
||||||
|
|
||||||
|
This model uses the Transformer architecture with GQA mechanism (q_head=24, kv_head=4), which saves KV cache memory compared to traditional MHA. The model is built by stacking 24 layers of Transformer blocks, with 1.0 billion parameters. Transformer is an autoregressive model that calculates the relationship between all previous tokens to obtain the probability distribution of the next token.
|
||||||
|
|
||||||
### 1. 模型搭建
|
The model now uses the **AutoModel** base class for flexible loading and saving:
|
||||||
|
|
||||||
本模型采用Transformer架构, 使用GQA(q_head=24, kv_head=4) 机制,相较于传统的MHA可以节省KV cache 的显存占用(但是目前没有做KV cache),通过堆叠24层Transformer实现模型的搭建, 参数量为1.0b。Transformer 是自回归模型, 是通过计算前面所有的token的关系得到下一个token的概率分布
|
```python
|
||||||
|
from astrai.model import AutoModel
|
||||||
|
|
||||||

|
# Load model from checkpoint
|
||||||
|
model = AutoModel.from_pretrained("path/to/model")
|
||||||
|
|
||||||
什么是自回归模型呢, 在把句子拆分成token之后, 模型会预测下一个token的概率分布。这意味着模型会根据给定的上下文(即已经出现的tokens序列),计算出下一个可能的token及其对应的概率。
|
# Save model to new directory
|
||||||
|
model.save_pretrained("path/to/save")
|
||||||
|
|
||||||
|
|
||||||
#### 1. 自回归
|
|
||||||
|
|
||||||
假设我们有一个句子被拆分成如下tokens列表:
|
|
||||||
|
|
||||||
```
|
|
||||||
["你好", "," "今天", "天气"]
|
|
||||||
```
|
```
|
||||||
|
|
||||||
接下来,模型会基于这个序列预测下一个可能出现的token。这通常以概率分布的形式给出,比如:
|
The Transformer model is registered via `@AutoModel.register('transformer')` decorator, allowing easy extension for new model types.
|
||||||
|
|
||||||
```
|
```mermaid
|
||||||
-> {"token": "不错", "probability": 0.4}
|
flowchart TB
|
||||||
-> {"token": "晴朗", "probability": 0.2}
|
subgraph Layers["Transformer Layers"]
|
||||||
-> ......
|
direction TB
|
||||||
|
A[Input Embedding] --> B[Transformer Block\nLayer 1]
|
||||||
|
B --> C[Transformer Block\nLayer ...]
|
||||||
|
C --> D[Transformer Block\nLayer 32]
|
||||||
|
D --> E[RMSNorm]
|
||||||
|
E --> F[Linear]
|
||||||
|
F --> G[SoftMax]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph TransformerBlock["Transformer Block"]
|
||||||
|
direction TB
|
||||||
|
H[x] --> I[RMSNorm]
|
||||||
|
I --> J[Linear → Q/K/V]
|
||||||
|
J --> K[Q]
|
||||||
|
J --> L[K]
|
||||||
|
J --> M[V]
|
||||||
|
K --> N[RoPE]
|
||||||
|
L --> O[RoPE]
|
||||||
|
N --> P["Q @ K^T / sqrt(d)"]
|
||||||
|
O --> P
|
||||||
|
P --> Q[Masked SoftMax]
|
||||||
|
Q --> R[S @ V]
|
||||||
|
M --> R
|
||||||
|
R --> S[Linear]
|
||||||
|
S --> T[+]
|
||||||
|
H --> T
|
||||||
|
T --> U[RMSNorm]
|
||||||
|
U --> V["Linear (gate)"]
|
||||||
|
U --> W["Linear (up)"]
|
||||||
|
V --> X[SiLU]
|
||||||
|
X --> Y[×]
|
||||||
|
W --> Y
|
||||||
|
Y --> Z["Linear (down)"]
|
||||||
|
Z --> AA[+]
|
||||||
|
T --> AA
|
||||||
|
AA --> BB[x']
|
||||||
|
end
|
||||||
|
|
||||||
|
classDef main fill:#e6f3ff,stroke:#0066cc;
|
||||||
|
classDef block fill:#fff2e6,stroke:#cc6600;
|
||||||
|
class Layers main;
|
||||||
|
class TransformerBlock block;
|
||||||
```
|
```
|
||||||
|
|
||||||
这里,“不错”和“晴朗”是两个可能跟随在“天气”之后的tokens,并且给出了每个token成为下一个token的可能性大小。
|
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).
|
||||||
|
|
||||||
之后,我们通过采样(通过top_k, top_p, temperature参数调整采样后的结果)得到下一个token并且将下一个token加入序列作为输入
|
#### 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.
|
||||||
["你好", "," "今天", "天气", "不错"]
|
|
||||||
```
|
|
||||||
|
|
||||||
之后都是在重复这个流程, 直到遇到控制流程结束的token(<|end_of_seqence|>)模型停止处理(一般模型都会设置控制token, 不然模型会一直输出到显存爆炸)。
|
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:
|
||||||
|
|
||||||
#### 2. 因果掩码
|
|
||||||
|
|
||||||
transformer 中采用注意力机制,输入的形状一般为[bsz, seq_len], 输出为[bsz, seq_len,n_dim], 为了实现预测下一个token, 模型的输入和输出必须错开来一个位置。模型预测的target必须错开一个位置, 在训练的时候我们也采用错开一个位置的方法
|
|
||||||
|
|
||||||
```
|
```
|
||||||
sequence : [[1, 2, 3, 4, 5, 6]]
|
sequence : [[1, 2, 3, 4, 5, 6]]
|
||||||
@@ -52,18 +85,14 @@ input_ids: [[1, 2, 3, 4, 5]]
|
|||||||
target_ids: [[2, 3, 4, 5, 6]]
|
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} = softmax(\frac{q_i^Tk_j}{\sqrt{d_k}}) $$
|
||||||
$$ s_{ij} := s_{ij} + mask_{ij} $$
|
$$ s_{ij} := s_{ij} + mask_{ij} $$
|
||||||
|
|
||||||
|
Here, the attention score represents the degree to which the model attends to the similarity between two tokens.
|
||||||
|
|
||||||
其中注意力得分代表了模型对两个token之间相似程度的关注程度
|
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:
|
||||||
|
|
||||||
对于decoder only结构的模型, 为了防止模型从未来的位置偷到信息, 在注意力的计算过程中需要增加掩码,我们需要在注意力得分计算之前应用一个掩码。这个掩码通常是一个下三角矩阵,对于长度为n的序列,它的形状是[n, n]。下面以一个长度为5的序列为例,展示如何创建这样的因果掩码矩阵:
|
|
||||||
|
|
||||||
```
|
```
|
||||||
[[0, -inf, -inf, -inf, -inf],
|
[[0, -inf, -inf, -inf, -inf],
|
||||||
@@ -73,17 +102,233 @@ $$ s_{ij} := s_{ij} + mask_{ij} $$
|
|||||||
[0, 0, 0, 0, 0]]
|
[0, 0, 0, 0, 0]]
|
||||||
```
|
```
|
||||||
|
|
||||||
在这个矩阵中,0表示可以注意到的位置,而-inf表示应该被掩盖(即不应注意到)的位置。因为这个句子保证了注意力得分中 $j > i$ 的部分通过softmax 之后由`inf` 变成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.
|
||||||
#### 3. 旋转位置编码
|
|
||||||
|
|
||||||
旋转位置编码(Rotary Position Embedding, RoPE)是一种为了解决Transformer模型中缺乏对序列位置信息直接建模的问题而设计的位置编码方法。与传统的位置编码(如正弦和余弦函数的位置编码)不同,RoPE通过将位置信息直接嵌入到查询(Query, Q)和键(Key, K)向量中来实现,使得模型能够更自然地处理序列中的相对位置关系。
|
|
||||||
|
|
||||||
|
|
||||||
$$ q_i = R_i W_q x_i $$
|
$$ q_i = R_i W_q x_i $$
|
||||||
$$ k_j = R_j W_k x_j $$
|
$$ 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 $$
|
$$ q_i^T k_j = (R_i W_q x_i)^T( R_j W_k x_j) = x_i^T W_q^T R_{i-j} W_k x_j $$
|
||||||
|
|
||||||
其中的 $R_{i-j}$ 控制了模型的不同token 在不同相对距离上注意力的衰减,在 $i - j$ 绝对值越大的时候, 衰减的程度越强, 通过这种方式能让模型学习到相对位置关系, 从而使得模型可以扩展和适应长序列
|
The $R_{i-j}$ controls the attenuation of attention for different tokens at different relative distances. When the absolute value of $i - j$ is larger, the degree of attenuation is stronger. This approach allows the model to learn relative position relationships, enabling the model to scale and adapt to longer sequences.
|
||||||
|
|
||||||
|
## KV Cache Implementation
|
||||||
|
|
||||||
|
According to the attention calculation formula:
|
||||||
|
|
||||||
|
$$
|
||||||
|
\begin{align*}
|
||||||
|
o_i &= \sum_j s_{ij} v_{j} \newline
|
||||||
|
s_{ij} &= \text{softmax}\left( \frac{q_{i} k_{j}}{\sqrt{d_k}} \right)
|
||||||
|
\end{align*}
|
||||||
|
$$
|
||||||
|
|
||||||
|
Since the model is an autoregressive model, we only need to calculate for the last part of the sequence, meaning the index $i$ is fixed as the last element of the sequence, and we compute $o_{n}$:
|
||||||
|
|
||||||
|
$$
|
||||||
|
\begin{align*}
|
||||||
|
o_n &= \sum_j s_{j}v_{j} \newline
|
||||||
|
s_j &= \text{softmax}\left(\frac{q_n k_{j}}{\sqrt{d_k}} \right)
|
||||||
|
\end{align*}
|
||||||
|
$$
|
||||||
|
|
||||||
|
If we expand the expression:
|
||||||
|
|
||||||
|
$$
|
||||||
|
o_n = \sum_j \text{softmax}\left(\frac{q_n k_{j}}{\sqrt{d_k}}\right)v_{j}
|
||||||
|
$$
|
||||||
|
|
||||||
|
In the above expression, only k and v have length indices, while $q$ does not. Therefore, during the calculation process, the input of $q$ is fixed as the last token from the previous input, while $k$ and $v$ need to be cached for parts of different lengths. Also, when caching, note that position encoding calculation should be performed before KV cache computation, otherwise there will be position encoding calculation errors.
|
||||||
|
|
||||||
|
### 4. AutoModel Loading
|
||||||
|
|
||||||
|
The project now uses the **AutoModel** base class for flexible model loading and saving:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from astrai.model import AutoModel
|
||||||
|
|
||||||
|
# Load model from checkpoint
|
||||||
|
model = AutoModel.from_pretrained("path/to/model")
|
||||||
|
|
||||||
|
# Save model to new directory
|
||||||
|
model.save_pretrained("path/to/save")
|
||||||
|
```
|
||||||
|
|
||||||
|
The Transformer model is registered via `@AutoModel.register('transformer')` decorator, allowing easy extension for new model types. The `from_pretrained` method automatically loads the `config.json` to determine the model type and uses safetensors format for weights.
|
||||||
|
|
||||||
|
### 5. Continuous Batching Inference
|
||||||
|
|
||||||
|
The inference engine supports **continuous batching** for efficient batch processing:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from astrai.inference import InferenceEngine, GenerationRequest
|
||||||
|
|
||||||
|
# Create inference engine with continuous batching
|
||||||
|
engine = InferenceEngine(
|
||||||
|
model=model,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Use GenerationRequest with messages format
|
||||||
|
request = GenerationRequest(
|
||||||
|
messages=[
|
||||||
|
{"role": "system", "content": "You are a helpful assistant."},
|
||||||
|
{"role": "user", "content": "Hello"},
|
||||||
|
],
|
||||||
|
temperature=0.8,
|
||||||
|
top_p=0.95,
|
||||||
|
top_k=50,
|
||||||
|
max_len=1024,
|
||||||
|
stream=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Generate with streaming
|
||||||
|
for token in engine.generate_with_request(request):
|
||||||
|
print(token, end="", flush=True)
|
||||||
|
```
|
||||||
|
|
||||||
|
The continuous batching feature allows dynamic batch composition where new requests can join at any time and completed requests are released immediately.
|
||||||
|
|
||||||
|
## HTTP API Usage
|
||||||
|
|
||||||
|
The inference server provides HTTP endpoints for remote inference. Start the server first:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python -m scripts.tools.server --port 8000
|
||||||
|
```
|
||||||
|
|
||||||
|
### OpenAI-Compatible Endpoint
|
||||||
|
|
||||||
|
The server provides an OpenAI-compatible chat completion endpoint at `/v1/chat/completions`:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"messages": [
|
||||||
|
{"role": "system", "content": "You are a helpful assistant."},
|
||||||
|
{"role": "user", "content": "Hello, how are you?"}
|
||||||
|
],
|
||||||
|
"temperature": 0.8,
|
||||||
|
"max_tokens": 2048,
|
||||||
|
"stream": false
|
||||||
|
}'
|
||||||
|
```
|
||||||
|
|
||||||
|
**Request Parameters:**
|
||||||
|
| Parameter | Type | Default | Description |
|
||||||
|
|-----------|------|---------|-------------|
|
||||||
|
| `messages` | List[dict] | Required | Chat messages with role and content |
|
||||||
|
| `temperature` | float | 1.0 | Sampling temperature (0.0-2.0) |
|
||||||
|
| `top_p` | float | 1.0 | Nucleus sampling threshold |
|
||||||
|
| `top_k` | int | 50 | Top-k sampling parameter |
|
||||||
|
| `max_tokens` | int | 1024 | Maximum tokens to generate |
|
||||||
|
| `stream` | bool | false | Enable streaming response |
|
||||||
|
|
||||||
|
**Response (non-streaming):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"id": "chatcmpl-1234567890",
|
||||||
|
"object": "chat.completion",
|
||||||
|
"created": 1234567890,
|
||||||
|
"model": "astrai",
|
||||||
|
"choices": [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"message": {"role": "assistant", "content": "Hello! I'm doing well..."},
|
||||||
|
"finish_reason": "stop"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"usage": {
|
||||||
|
"prompt_tokens": 20,
|
||||||
|
"completion_tokens": 15,
|
||||||
|
"total_tokens": 35
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Streaming Response
|
||||||
|
|
||||||
|
Enable streaming for real-time token-by-token output:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"messages": [{"role": "user", "content": "Write a story"}],
|
||||||
|
"stream": true,
|
||||||
|
"max_tokens": 500
|
||||||
|
}'
|
||||||
|
```
|
||||||
|
|
||||||
|
The server uses Server-Sent Events (SSE) with content type `text/event-stream`.
|
||||||
|
|
||||||
|
### Anthropic-Compatible Endpoint
|
||||||
|
|
||||||
|
The server also provides an Anthropic-compatible endpoint at `/v1/messages`:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"model": "astrai",
|
||||||
|
"system": "You are a helpful assistant.",
|
||||||
|
"messages": [{"role": "user", "content": "Hello, how are you?"}],
|
||||||
|
"max_tokens": 2048
|
||||||
|
}'
|
||||||
|
```
|
||||||
|
|
||||||
|
Response:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"id": "msg_abc123...",
|
||||||
|
"type": "message",
|
||||||
|
"role": "assistant",
|
||||||
|
"model": "astrai",
|
||||||
|
"content": [{"type": "text", "text": "Hello! I am doing well..."}],
|
||||||
|
"stop_reason": "end_turn",
|
||||||
|
"stop_sequence": null,
|
||||||
|
"usage": {"input_tokens": 20, "output_tokens": 15}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Streaming:
|
||||||
|
```bash
|
||||||
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"model": "astrai",
|
||||||
|
"system": "You are a helpful assistant.",
|
||||||
|
"messages": [{"role": "user", "content": "Write a short poem"}],
|
||||||
|
"max_tokens": 500,
|
||||||
|
"stream": true
|
||||||
|
}'
|
||||||
|
```
|
||||||
|
|
||||||
|
Supports `stop_sequences` for early termination:
|
||||||
|
```bash
|
||||||
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"model": "astrai",
|
||||||
|
"messages": [{"role": "user", "content": "Write a story"}],
|
||||||
|
"max_tokens": 500,
|
||||||
|
"stop_sequences": ["The end", "THE END"]
|
||||||
|
}'
|
||||||
|
```
|
||||||
|
|
||||||
|
### Health Check
|
||||||
|
|
||||||
|
Monitor server and model status:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl http://localhost:8000/health
|
||||||
|
# {"status": "ok", "model_loaded": true}
|
||||||
|
|
||||||
|
curl http://localhost:8000/stats
|
||||||
|
# {"total_tasks": 10, "total_tokens": 5000, "active_tasks": 1, "waiting_queue": 0}
|
||||||
|
```
|
||||||
|
|
||||||
|
> Document Update Time: 2026-04-09
|
||||||
@@ -1,27 +0,0 @@
|
|||||||
## kv_cache 实现
|
|
||||||
|
|
||||||
根据注意力的计算公式
|
|
||||||
|
|
||||||
$$
|
|
||||||
\begin{align*}
|
|
||||||
o_i &= \sum_j s_{ij} v_{j} \\
|
|
||||||
s_{ij} &= \text{softmax}\left( \sum_n \frac{q_{i,n} k_{j,n}}{\sqrt{d_k}} \right)
|
|
||||||
\end{align*}
|
|
||||||
$$
|
|
||||||
|
|
||||||
由于模型是自回归模型, 我们只用求序列最后一个部分,也就是说 $ i $ 的下标是确定的, 是序列最后一个元素, 我们求的是 $o_{n} $
|
|
||||||
|
|
||||||
$$
|
|
||||||
\begin{align*}
|
|
||||||
o_n &= \sum_j s_{j}v_{j,n} \\
|
|
||||||
s_j &= \text{softmax}\left(\sum_n\frac{q_n k_{j,n}}{\sqrt{d_k}} \right)
|
|
||||||
\end{align*}
|
|
||||||
$$
|
|
||||||
|
|
||||||
如果我们把式子展开
|
|
||||||
|
|
||||||
$$
|
|
||||||
o_n = \sum_j \sum_n \text{softmax}\left(\frac{q_n k_{j,n}}{\sqrt{d_k}}\right)v_{j,n}
|
|
||||||
$$
|
|
||||||
|
|
||||||
以上表达式只有k和v存在长度下标, 而 $q$ 没有, 所以计算过程中 $q$ 的输入是确定的上次输入的最后一个token, 而 $k, v$ 是需要对不同长度的部分进行缓存的,同时缓存的时候应该注意位置编码的计算应该在kvcache的计算之前进行,否则会存在位置编码的计算错误
|
|
||||||
@@ -0,0 +1,158 @@
|
|||||||
|
# Parameter Documentation
|
||||||
|
|
||||||
|
## Training Parameters
|
||||||
|
|
||||||
|
### Basic Parameters
|
||||||
|
|
||||||
|
| Parameter | Description | Default |
|
||||||
|
|-----------|-------------|---------|
|
||||||
|
| `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`) | required |
|
||||||
|
| `--data_root_path` | Dataset root directory | required |
|
||||||
|
| `--param_path` | Model parameters or checkpoint path | required |
|
||||||
|
| `--n_epoch` | Total training epochs | 1 |
|
||||||
|
| `--batch_size` | Batch size | 1 |
|
||||||
|
| `--accumulation_steps` | Gradient accumulation steps between optimizer steps | 1 |
|
||||||
|
|
||||||
|
### Learning Rate Scheduling
|
||||||
|
|
||||||
|
| Parameter | Description | Default |
|
||||||
|
|-----------|-------------|---------|
|
||||||
|
| `--warmup_steps` | Warmup steps | 1000 |
|
||||||
|
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
||||||
|
| `--max_grad_norm` | Maximum gradient norm for clipping | 1.0 |
|
||||||
|
|
||||||
|
### Optimizer (AdamW)
|
||||||
|
|
||||||
|
| Parameter | Description | Default |
|
||||||
|
|-----------|-------------|---------|
|
||||||
|
| `--adamw_beta1` | AdamW beta1 | 0.9 |
|
||||||
|
| `--adamw_beta2` | AdamW beta2 | 0.95 |
|
||||||
|
| `--adamw_weight_decay` | AdamW weight decay | 0.01 |
|
||||||
|
|
||||||
|
### Data Loading
|
||||||
|
|
||||||
|
| Parameter | Description | Default |
|
||||||
|
|-----------|-------------|---------|
|
||||||
|
| `--window_size` | Max input sequence length | model config `max_len` |
|
||||||
|
| `--stride` | Stride for sliding window over sequences | None |
|
||||||
|
| `--random_seed` | Random seed for reproducibility | 3407 |
|
||||||
|
| `--num_workers` | DataLoader worker processes | 4 |
|
||||||
|
| `--no_pin_memory` | Disable pin_memory (enabled by default) | (flag) |
|
||||||
|
|
||||||
|
### Checkpoint & Resume
|
||||||
|
|
||||||
|
| Parameter | Description | Default |
|
||||||
|
|-----------|-------------|---------|
|
||||||
|
| `--ckpt_interval` | Iterations between checkpoints | 5000 |
|
||||||
|
| `--ckpt_dir` | Checkpoint save directory | checkpoint |
|
||||||
|
| `--start_epoch` | Resume from epoch (0 = from scratch) | 0 |
|
||||||
|
| `--start_batch` | Resume from batch iteration | 0 |
|
||||||
|
|
||||||
|
### Distributed Training
|
||||||
|
|
||||||
|
| Parameter | Description | Default |
|
||||||
|
|-----------|-------------|---------|
|
||||||
|
| `--nprocs` | Number of GPUs / processes | 1 |
|
||||||
|
| `--device_type` | Device type | cuda |
|
||||||
|
|
||||||
|
### Strategy-specific
|
||||||
|
|
||||||
|
| Parameter | Description | Default | Used by |
|
||||||
|
|-----------|-------------|---------|---------|
|
||||||
|
| `--dpo_beta` | DPO beta value | 0.1 | `dpo` |
|
||||||
|
| `--label_smoothing` | Label smoothing for cross-entropy loss | 0.1 | `seq`, `sft` |
|
||||||
|
| `--group_size` | GRPO group size | 4 | `grpo` |
|
||||||
|
| `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo` |
|
||||||
|
| `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo` |
|
||||||
|
| `--grpo_sync_interval` | GRPO ref_model sync interval (steps) | 200 | `grpo` |
|
||||||
|
|
||||||
|
### Usage Example
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/tools/train.py \
|
||||||
|
--train_type seq \
|
||||||
|
--data_root_path /path/to/dataset \
|
||||||
|
--param_path /path/to/model \
|
||||||
|
--n_epoch 3 \
|
||||||
|
--batch_size 4 \
|
||||||
|
--accumulation_steps 8 \
|
||||||
|
--max_lr 3e-4 \
|
||||||
|
--warmup_steps 2000 \
|
||||||
|
--max_grad_norm 1.0 \
|
||||||
|
--ckpt_interval 5000 \
|
||||||
|
--ckpt_dir ./checkpoints \
|
||||||
|
--num_workers 4 \
|
||||||
|
--nprocs 1 \
|
||||||
|
--device_type cuda
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Generation Parameters
|
||||||
|
|
||||||
|
### GenerationRequest Parameters
|
||||||
|
|
||||||
|
| Parameter | Description | Default Value |
|
||||||
|
|-----------|-------------|---------------|
|
||||||
|
| `messages` | List of message dictionaries (role, content) | required |
|
||||||
|
| `temperature` | Sampling temperature (higher = more random) | 1.0 |
|
||||||
|
| `top_p` | Nucleus sampling threshold | 1.0 |
|
||||||
|
| `top_k` | Top-k sampling count | 50 |
|
||||||
|
| `max_len` | Maximum generation length | 1024 |
|
||||||
|
| `stream` | Whether to stream output | False |
|
||||||
|
|
||||||
|
### Usage Example
|
||||||
|
|
||||||
|
```python
|
||||||
|
import torch
|
||||||
|
from astrai.model import AutoModel
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
from astrai.inference import InferenceEngine, GenerationRequest
|
||||||
|
|
||||||
|
# Load model using AutoModel
|
||||||
|
model = AutoModel.from_pretrained("your_model_dir")
|
||||||
|
|
||||||
|
# Load tokenizer
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained("your_model_dir")
|
||||||
|
|
||||||
|
# Create engine with separate model and tokenizer
|
||||||
|
engine = InferenceEngine(
|
||||||
|
model=model,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Build request with messages format
|
||||||
|
request = GenerationRequest(
|
||||||
|
messages=[
|
||||||
|
{"role": "system", "content": "You are a helpful assistant."},
|
||||||
|
{"role": "user", "content": "Hello"},
|
||||||
|
],
|
||||||
|
temperature=0.8,
|
||||||
|
top_p=0.95,
|
||||||
|
top_k=50,
|
||||||
|
max_len=1024,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Generate (streaming)
|
||||||
|
for token in engine.generate_with_request(request):
|
||||||
|
print(token, end="", flush=True)
|
||||||
|
|
||||||
|
# Or use simple generate interface
|
||||||
|
result = engine.generate(
|
||||||
|
prompt="Hello",
|
||||||
|
stream=False,
|
||||||
|
max_tokens=1024,
|
||||||
|
temperature=0.8,
|
||||||
|
top_p=0.95,
|
||||||
|
top_k=50,
|
||||||
|
)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Generation Modes
|
||||||
|
|
||||||
|
| Mode | Description |
|
||||||
|
|------|-------------|
|
||||||
|
| `stream=True` | Streaming output, yields token by token |
|
||||||
|
| `stream=False` | Non-streaming output, returns complete result |
|
||||||
|
|
||||||
|
> Document Update Time: 2026-04-09
|
||||||
Binary file not shown.
|
After Width: | Height: | Size: 281 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 21 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 11 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 590 KiB |
@@ -0,0 +1,32 @@
|
|||||||
|
__version__ = "1.3.4"
|
||||||
|
__author__ = "ViperEkura"
|
||||||
|
|
||||||
|
from astrai.config import (
|
||||||
|
ModelConfig,
|
||||||
|
TrainConfig,
|
||||||
|
)
|
||||||
|
from astrai.dataset import DatasetFactory
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.inference import (
|
||||||
|
GenerationRequest,
|
||||||
|
InferenceEngine,
|
||||||
|
)
|
||||||
|
from astrai.model import AutoModel, Transformer
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
from astrai.trainer import CallbackFactory, SchedulerFactory, StrategyFactory, Trainer
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"Transformer",
|
||||||
|
"ModelConfig",
|
||||||
|
"TrainConfig",
|
||||||
|
"DatasetFactory",
|
||||||
|
"AutoTokenizer",
|
||||||
|
"GenerationRequest",
|
||||||
|
"InferenceEngine",
|
||||||
|
"Trainer",
|
||||||
|
"CallbackFactory",
|
||||||
|
"StrategyFactory",
|
||||||
|
"SchedulerFactory",
|
||||||
|
"BaseFactory",
|
||||||
|
"AutoModel",
|
||||||
|
]
|
||||||
@@ -0,0 +1,8 @@
|
|||||||
|
from astrai.config.model_config import ModelConfig
|
||||||
|
from astrai.config.train_config import TrainConfig
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
# Model configuration
|
||||||
|
"ModelConfig",
|
||||||
|
"TrainConfig",
|
||||||
|
]
|
||||||
@@ -0,0 +1,42 @@
|
|||||||
|
import json
|
||||||
|
from dataclasses import asdict, dataclass
|
||||||
|
from typing import Optional, Self
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ModelConfig:
|
||||||
|
# basic config
|
||||||
|
model_type: Optional[str] = None
|
||||||
|
vocab_size: Optional[int] = None
|
||||||
|
dim: Optional[int] = None
|
||||||
|
|
||||||
|
n_layers: Optional[int] = None
|
||||||
|
norm_eps: Optional[float] = None
|
||||||
|
dim_ffn: Optional[int] = None
|
||||||
|
tie_weight: Optional[bool] = None
|
||||||
|
|
||||||
|
# RoPE
|
||||||
|
max_len: Optional[int] = None
|
||||||
|
rope_theta: Optional[float] = None
|
||||||
|
|
||||||
|
# GQA
|
||||||
|
n_heads: Optional[int] = None
|
||||||
|
n_kv_heads: Optional[int] = None
|
||||||
|
use_qk_norm: Optional[bool] = None
|
||||||
|
use_gated_attention: Optional[bool] = None
|
||||||
|
|
||||||
|
def load(self, config_path: str) -> Self:
|
||||||
|
config = {}
|
||||||
|
with open(config_path, "r") as f:
|
||||||
|
config.update(json.load(f))
|
||||||
|
|
||||||
|
for key, value in config.items():
|
||||||
|
if hasattr(self, key):
|
||||||
|
setattr(self, key, value)
|
||||||
|
|
||||||
|
return self
|
||||||
|
|
||||||
|
def save(self, config_path: str):
|
||||||
|
config_dict = {k: v for k, v in asdict(self).items() if v is not None}
|
||||||
|
with open(config_path, "w") as f:
|
||||||
|
json.dump(config_dict, f, indent=4)
|
||||||
@@ -0,0 +1,98 @@
|
|||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Callable, Optional
|
||||||
|
|
||||||
|
import torch.nn as nn
|
||||||
|
from torch.optim import Optimizer
|
||||||
|
from torch.optim.lr_scheduler import LRScheduler
|
||||||
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class TrainConfig:
|
||||||
|
# basic setting
|
||||||
|
model: nn.Module = field(default=None, metadata={"help": "Model for training."})
|
||||||
|
strategy: str = field(default=None, metadata={"help": "Training strategy."})
|
||||||
|
dataset: Dataset = field(default=None, metadata={"help": "Dataset for training."})
|
||||||
|
optimizer_fn: Callable[[nn.Module], Optimizer] = field(
|
||||||
|
default=None, metadata={"help": "Optimizer factory for training."}
|
||||||
|
)
|
||||||
|
scheduler_fn: Callable[[Optimizer], LRScheduler] = field(
|
||||||
|
default=None, metadata={"help": "Scheduler factory for training."}
|
||||||
|
)
|
||||||
|
n_epoch: int = field(default=1, metadata={"help": "Number of epochs for training."})
|
||||||
|
batch_size: int = field(default=4, metadata={"help": "Batch size for training."})
|
||||||
|
accumulation_steps: int = field(
|
||||||
|
default=1, metadata={"help": "Number of iterations between steps."}
|
||||||
|
)
|
||||||
|
max_grad_norm: float = field(
|
||||||
|
default=1.0, metadata={"help": "Maximum gradient norm."}
|
||||||
|
)
|
||||||
|
|
||||||
|
# checkpoint setting
|
||||||
|
start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."})
|
||||||
|
start_batch: int = field(
|
||||||
|
default=0, metadata={"help": "Start batch iteration for training."}
|
||||||
|
)
|
||||||
|
ckpt_dir: str = field(
|
||||||
|
default="./checkpoint", metadata={"help": "Checkpoint directory."}
|
||||||
|
)
|
||||||
|
ckpt_interval: int = field(
|
||||||
|
default=5000, metadata={"help": "Number of iterations between checkpoints."}
|
||||||
|
)
|
||||||
|
|
||||||
|
# dataloader setting
|
||||||
|
random_seed: int = field(default=3407, metadata={"help": "Random seed."})
|
||||||
|
num_workers: int = field(
|
||||||
|
default=0, metadata={"help": "Number of workers for dataloader."}
|
||||||
|
)
|
||||||
|
prefetch_factor: Optional[int] = field(
|
||||||
|
default=None, metadata={"help": "Prefetch factor for dataloader."}
|
||||||
|
)
|
||||||
|
pin_memory: bool = field(
|
||||||
|
default=False, metadata={"help": "Pin memory for dataloader."}
|
||||||
|
)
|
||||||
|
|
||||||
|
# distributed training
|
||||||
|
nprocs: int = field(
|
||||||
|
default=1, metadata={"help": "Number of processes for distributed training."}
|
||||||
|
)
|
||||||
|
backend: str = field(
|
||||||
|
default="nccl", metadata={"help": "Distributed training backend."}
|
||||||
|
)
|
||||||
|
master_addr: str = field(
|
||||||
|
default="localhost",
|
||||||
|
metadata={"help": "Master address for distributed training."},
|
||||||
|
)
|
||||||
|
master_port: str = field(
|
||||||
|
default="29500", metadata={"help": "Master port for distributed training."}
|
||||||
|
)
|
||||||
|
parallel_wrapper: Optional[Callable] = field(
|
||||||
|
default=None, metadata={"help": "Parallel function for training."}
|
||||||
|
)
|
||||||
|
state_dict_fn: Optional[Callable] = field(
|
||||||
|
default=None, metadata={"help": "Parallel function for state dict saving."}
|
||||||
|
)
|
||||||
|
|
||||||
|
# others
|
||||||
|
device_type: str = field(
|
||||||
|
default="cuda", metadata={"help": "Device type for distributed training."}
|
||||||
|
)
|
||||||
|
extra_kwargs: dict = field(
|
||||||
|
default_factory=dict, metadata={"help": "Other arguments."}
|
||||||
|
)
|
||||||
|
|
||||||
|
def __post_init__(self):
|
||||||
|
self.validate()
|
||||||
|
|
||||||
|
def validate(self):
|
||||||
|
required_fields = [
|
||||||
|
"model",
|
||||||
|
"strategy",
|
||||||
|
"dataset",
|
||||||
|
"optimizer_fn",
|
||||||
|
"scheduler_fn",
|
||||||
|
]
|
||||||
|
|
||||||
|
for field_name in required_fields:
|
||||||
|
if getattr(self, field_name) is None:
|
||||||
|
raise ValueError(f"{field_name} is required.")
|
||||||
@@ -0,0 +1,19 @@
|
|||||||
|
from astrai.dataset.dataset import (
|
||||||
|
BaseDataset,
|
||||||
|
BaseSegmentFetcher,
|
||||||
|
DatasetFactory,
|
||||||
|
MultiSegmentFetcher,
|
||||||
|
)
|
||||||
|
from astrai.dataset.sampler import ResumableDistributedSampler
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
# Base classes
|
||||||
|
"BaseDataset",
|
||||||
|
# Factory
|
||||||
|
"DatasetFactory",
|
||||||
|
# Fetchers
|
||||||
|
"BaseSegmentFetcher",
|
||||||
|
"MultiSegmentFetcher",
|
||||||
|
# Sampler
|
||||||
|
"ResumableDistributedSampler",
|
||||||
|
]
|
||||||
@@ -0,0 +1,338 @@
|
|||||||
|
"""Dataset implementations with factory pattern for training."""
|
||||||
|
|
||||||
|
import bisect
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from typing import Dict, List, Optional, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.serialization import load_h5
|
||||||
|
|
||||||
|
|
||||||
|
class BaseSegmentFetcher:
|
||||||
|
"""Fetches data segments across multiple tensor segments.
|
||||||
|
|
||||||
|
Maintains cumulative lengths for efficient range queries across
|
||||||
|
multiple discontinuous segments.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, segments: List[Tensor]):
|
||||||
|
self.segments = segments
|
||||||
|
self.cum_lengths = []
|
||||||
|
|
||||||
|
total = 0
|
||||||
|
for seg in segments:
|
||||||
|
total += torch.numel(seg)
|
||||||
|
self.cum_lengths.append(total)
|
||||||
|
|
||||||
|
self.total_length = total
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
return self.total_length
|
||||||
|
|
||||||
|
def fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
||||||
|
"""Fetch data in the range [begin_idx, end_idx).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
begin_idx: Starting index (inclusive)
|
||||||
|
end_idx: Ending index (exclusive)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Concatenated tensor of data in the specified range
|
||||||
|
"""
|
||||||
|
if not (
|
||||||
|
0 <= begin_idx < self.total_length and 0 <= end_idx <= self.total_length
|
||||||
|
):
|
||||||
|
raise ValueError("begin_idx or end_idx out of bounds")
|
||||||
|
if begin_idx >= end_idx:
|
||||||
|
return torch.tensor([], dtype=torch.long)
|
||||||
|
|
||||||
|
# Find segment boundaries for the range
|
||||||
|
seg_start_idx = bisect.bisect_right(self.cum_lengths, begin_idx)
|
||||||
|
seg_end_idx = bisect.bisect_left(self.cum_lengths, end_idx)
|
||||||
|
|
||||||
|
result_segments = []
|
||||||
|
|
||||||
|
for i in range(seg_start_idx, seg_end_idx + 1):
|
||||||
|
prev_cum = self.cum_lengths[i - 1] if i > 0 else 0
|
||||||
|
start = max(begin_idx - prev_cum, 0)
|
||||||
|
end = min(end_idx - prev_cum, len(self.segments[i]))
|
||||||
|
data = self.segments[i][start:end]
|
||||||
|
result_segments.append(data)
|
||||||
|
|
||||||
|
return torch.cat(result_segments, dim=0)
|
||||||
|
|
||||||
|
|
||||||
|
class MultiSegmentFetcher:
|
||||||
|
"""Manages multiple segment fetchers for different data keys.
|
||||||
|
|
||||||
|
Each key corresponds to a different type of data (e.g., "sequence", "mask").
|
||||||
|
"""
|
||||||
|
|
||||||
|
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."""
|
||||||
|
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.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
begin_idx: Starting index
|
||||||
|
end_idx: Ending index
|
||||||
|
keys: Single key or list of keys to fetch
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary of tensors if multiple keys, single tensor if one key
|
||||||
|
"""
|
||||||
|
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 BaseDataset(Dataset, ABC):
|
||||||
|
"""Abstract base class for all dataset types.
|
||||||
|
|
||||||
|
Implements common functionality for window-based data fetching.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, window_size: int, stride: int):
|
||||||
|
super().__init__()
|
||||||
|
self.segments = {}
|
||||||
|
self.window_size = window_size
|
||||||
|
self.stride = stride
|
||||||
|
self.total_samples = None
|
||||||
|
self.fetcher: Optional[MultiSegmentFetcher] = None
|
||||||
|
|
||||||
|
def load(self, load_path: str):
|
||||||
|
"""Load dataset from HDF5 file.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
load_path: Path to the HDF5 data file
|
||||||
|
"""
|
||||||
|
self.segments = load_h5(load_path)
|
||||||
|
self.fetcher = MultiSegmentFetcher(self.segments)
|
||||||
|
self.total_samples = len(self.fetcher)
|
||||||
|
|
||||||
|
def get_index(self, index: int) -> tuple:
|
||||||
|
"""Calculate begin and end indices for a sample.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
index: Sample index
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (begin_idx, end_idx)
|
||||||
|
"""
|
||||||
|
assert self.total_samples > self.window_size
|
||||||
|
|
||||||
|
begin_idx = min(index * self.stride, self.total_samples - 1 - self.window_size)
|
||||||
|
end_idx = min(begin_idx + self.window_size, self.total_samples - 1)
|
||||||
|
|
||||||
|
return begin_idx, end_idx
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
|
"""Get a single sample by index.
|
||||||
|
|
||||||
|
Must be implemented by subclasses.
|
||||||
|
"""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
assert self.total_samples is not None
|
||||||
|
if self.total_samples <= self.window_size:
|
||||||
|
return 0
|
||||||
|
return (self.total_samples - 1 - self.window_size) // self.stride + 1
|
||||||
|
|
||||||
|
|
||||||
|
class DatasetFactory(BaseFactory["BaseDataset"]):
|
||||||
|
"""Factory class for creating dataset instances.
|
||||||
|
|
||||||
|
Supports decorator-based registration for extensible dataset types.
|
||||||
|
All default dataset types (seq, sft, dpo, grpo) are registered automatically
|
||||||
|
when their classes are defined with the decorator.
|
||||||
|
|
||||||
|
Example usage:
|
||||||
|
@DatasetFactory.register("custom")
|
||||||
|
class CustomDataset(BaseDataset):
|
||||||
|
...
|
||||||
|
|
||||||
|
dataset = DatasetFactory.create("custom", window_size, stride)
|
||||||
|
"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _validate_component(cls, dataset_cls: type) -> None:
|
||||||
|
"""Validate that the dataset class inherits from BaseDataset."""
|
||||||
|
if not issubclass(dataset_cls, BaseDataset):
|
||||||
|
raise TypeError(f"{dataset_cls.__name__} must inherit from BaseDataset")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(cls, train_type: str, window_size: int, stride: int) -> "BaseDataset":
|
||||||
|
"""Create a dataset instance.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
train_type: Type of training ("seq", "sft", "dpo", "grpo")
|
||||||
|
window_size: Window size for data sampling
|
||||||
|
stride: Stride between consecutive samples
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dataset instance
|
||||||
|
"""
|
||||||
|
return super().create(train_type, window_size, stride)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def load(
|
||||||
|
cls,
|
||||||
|
train_type: str,
|
||||||
|
load_path: str,
|
||||||
|
window_size: int,
|
||||||
|
stride: Optional[int] = None,
|
||||||
|
) -> "BaseDataset":
|
||||||
|
"""Create and load a dataset in one step.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
train_type: Type of training dataset
|
||||||
|
load_path: Path to the data file
|
||||||
|
window_size: Window size for data sampling
|
||||||
|
stride: Stride between consecutive samples (default: same as window_size)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Loaded dataset instance
|
||||||
|
"""
|
||||||
|
if stride is None:
|
||||||
|
stride = window_size
|
||||||
|
|
||||||
|
dataset = cls.create(train_type, window_size, stride)
|
||||||
|
dataset.load(load_path)
|
||||||
|
|
||||||
|
return dataset
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def available_types(cls) -> list:
|
||||||
|
"""Return list of registered dataset type names."""
|
||||||
|
return cls.list_registered()
|
||||||
|
|
||||||
|
|
||||||
|
# ============== Dataset Classes ==============
|
||||||
|
# All dataset classes are registered at class definition time using the decorator
|
||||||
|
|
||||||
|
|
||||||
|
@DatasetFactory.register("seq")
|
||||||
|
class SEQDataset(BaseDataset):
|
||||||
|
"""Dataset for sequential next-token prediction training."""
|
||||||
|
|
||||||
|
def __init__(self, window_size: int, stride: int):
|
||||||
|
super().__init__(window_size, stride)
|
||||||
|
|
||||||
|
def _fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
||||||
|
return self.fetcher.key_fetch(begin_idx, end_idx, "sequence")
|
||||||
|
|
||||||
|
def __getitem__(self, index):
|
||||||
|
begin_idx, end_idx = self.get_index(index)
|
||||||
|
|
||||||
|
x = self._fetch_data(begin_idx, end_idx).to(dtype=torch.long)
|
||||||
|
y = self._fetch_data(begin_idx + 1, end_idx + 1).to(dtype=torch.long)
|
||||||
|
|
||||||
|
return {"input_ids": x, "target_ids": y}
|
||||||
|
|
||||||
|
|
||||||
|
@DatasetFactory.register("sft")
|
||||||
|
class SFTDataset(BaseDataset):
|
||||||
|
"""Dataset for supervised fine-tuning with loss masking."""
|
||||||
|
|
||||||
|
def __init__(self, window_size: int, stride: int):
|
||||||
|
super().__init__(window_size, stride)
|
||||||
|
|
||||||
|
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||||
|
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
||||||
|
|
||||||
|
def __getitem__(self, index):
|
||||||
|
begin_idx, end_idx = self.get_index(index)
|
||||||
|
|
||||||
|
x = self._fetch_data(begin_idx, end_idx, "sequence").to(dtype=torch.long)
|
||||||
|
y = self._fetch_data(begin_idx + 1, end_idx + 1, "sequence").to(
|
||||||
|
dtype=torch.long
|
||||||
|
)
|
||||||
|
loss_mask = self._fetch_data(begin_idx + 1, end_idx + 1, "loss_mask").to(
|
||||||
|
dtype=torch.bool
|
||||||
|
)
|
||||||
|
|
||||||
|
return {"input_ids": x, "target_ids": y, "loss_mask": loss_mask}
|
||||||
|
|
||||||
|
|
||||||
|
@DatasetFactory.register("dpo")
|
||||||
|
class DPODataset(BaseDataset):
|
||||||
|
"""Dataset for Direct Preference Optimization training."""
|
||||||
|
|
||||||
|
def __init__(self, window_size: int, stride: int):
|
||||||
|
super().__init__(window_size, stride)
|
||||||
|
|
||||||
|
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||||
|
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
||||||
|
|
||||||
|
def __getitem__(self, index: int):
|
||||||
|
begin_idx, end_idx = self.get_index(index)
|
||||||
|
|
||||||
|
chosen = self._fetch_data(begin_idx, end_idx, "chosen").to(dtype=torch.long)
|
||||||
|
rejected = self._fetch_data(begin_idx, end_idx, "rejected").to(dtype=torch.long)
|
||||||
|
chosen_mask = self._fetch_data(begin_idx, end_idx, "chosen_mask").to(
|
||||||
|
dtype=torch.bool
|
||||||
|
)
|
||||||
|
rejected_mask = self._fetch_data(begin_idx, end_idx, "rejected_mask").to(
|
||||||
|
dtype=torch.bool
|
||||||
|
)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"chosen": chosen,
|
||||||
|
"rejected": rejected,
|
||||||
|
"chosen_mask": chosen_mask,
|
||||||
|
"rejected_mask": rejected_mask,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@DatasetFactory.register("grpo")
|
||||||
|
class GRPODataset(BaseDataset):
|
||||||
|
"""Dataset for Group Relative Policy Optimization training."""
|
||||||
|
|
||||||
|
def __init__(self, window_size: int, stride: int):
|
||||||
|
super().__init__(window_size, stride)
|
||||||
|
|
||||||
|
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||||
|
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
||||||
|
|
||||||
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
|
begin_idx, end_idx = self.get_index(index)
|
||||||
|
|
||||||
|
prompts = self._fetch_data(begin_idx, end_idx, "prompts")
|
||||||
|
responses = self._fetch_data(begin_idx, end_idx, "responses")
|
||||||
|
masks = self._fetch_data(begin_idx, end_idx, "masks")
|
||||||
|
rewards = self._fetch_data(begin_idx, end_idx, "rewards")
|
||||||
|
|
||||||
|
return {
|
||||||
|
"prompts": prompts,
|
||||||
|
"responses": responses,
|
||||||
|
"masks": masks,
|
||||||
|
"rewards": rewards,
|
||||||
|
}
|
||||||
@@ -0,0 +1,78 @@
|
|||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
|
from torch.utils.data import Dataset, Sampler
|
||||||
|
|
||||||
|
|
||||||
|
class ResumableDistributedSampler(Sampler[int]):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
data_source: Dataset,
|
||||||
|
start_epoch: int = 0,
|
||||||
|
start_iter: int = 0,
|
||||||
|
seed: int = 42,
|
||||||
|
drop_last: bool = False,
|
||||||
|
shuffle: bool = True,
|
||||||
|
process_group: Optional[dist.ProcessGroup] = None,
|
||||||
|
):
|
||||||
|
self.epoch = start_epoch
|
||||||
|
self.iter = start_iter
|
||||||
|
self.seed = seed
|
||||||
|
self.num_samples = len(data_source)
|
||||||
|
|
||||||
|
if process_group is not None:
|
||||||
|
# input process group
|
||||||
|
self.rank = dist.get_rank(process_group)
|
||||||
|
self.num_replicas = dist.get_world_size(process_group)
|
||||||
|
|
||||||
|
elif dist.is_available() and dist.is_initialized():
|
||||||
|
# use default process group
|
||||||
|
process_group = dist.group.WORLD
|
||||||
|
self.rank = dist.get_rank()
|
||||||
|
self.num_replicas = dist.get_world_size()
|
||||||
|
|
||||||
|
else:
|
||||||
|
# single process
|
||||||
|
self.rank = 0
|
||||||
|
self.num_replicas = 1
|
||||||
|
|
||||||
|
self.drop_last = drop_last
|
||||||
|
self.shuffle = shuffle
|
||||||
|
|
||||||
|
offset = 0 if drop_last else self.num_replicas - 1
|
||||||
|
self.num_samples_per_replica = (self.num_samples + offset) // self.num_replicas
|
||||||
|
self.total_size = self.num_samples_per_replica * self.num_replicas
|
||||||
|
|
||||||
|
self._indices = None
|
||||||
|
|
||||||
|
def _get_indices(self):
|
||||||
|
if self.shuffle:
|
||||||
|
generator = torch.Generator()
|
||||||
|
generator.manual_seed(self.seed + self.epoch)
|
||||||
|
indices = torch.randperm(self.num_samples, generator=generator).tolist()
|
||||||
|
else:
|
||||||
|
indices = torch.arange(self.num_samples).tolist()
|
||||||
|
|
||||||
|
if not self.drop_last and self.num_samples < self.total_size:
|
||||||
|
padding_size = self.total_size - len(indices)
|
||||||
|
indices += indices[:padding_size]
|
||||||
|
|
||||||
|
local_indices = indices[self.rank : self.total_size : self.num_replicas]
|
||||||
|
|
||||||
|
self.iter = self.iter % self.num_samples_per_replica
|
||||||
|
self._indices = local_indices[self.iter :]
|
||||||
|
|
||||||
|
def __iter__(self):
|
||||||
|
if self._indices is None:
|
||||||
|
self._get_indices()
|
||||||
|
|
||||||
|
for i in self._indices:
|
||||||
|
self.iter += 1
|
||||||
|
yield i
|
||||||
|
|
||||||
|
self.epoch += 1
|
||||||
|
self._indices = None
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return self.num_samples_per_replica
|
||||||
@@ -0,0 +1,190 @@
|
|||||||
|
"""Base factory class for extensible component registration."""
|
||||||
|
|
||||||
|
from abc import ABC
|
||||||
|
from typing import Callable, Dict, Generic, List, Optional, Tuple, Type, TypeVar
|
||||||
|
|
||||||
|
T = TypeVar("T")
|
||||||
|
|
||||||
|
|
||||||
|
class Registry:
|
||||||
|
"""Flexible registry for component classes with category and priority support.
|
||||||
|
|
||||||
|
This registry stores component classes with optional metadata (category, priority).
|
||||||
|
It provides methods for registration, retrieval, and listing with filtering.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._entries = {} # name -> (component_cls, category, priority)
|
||||||
|
|
||||||
|
def register(
|
||||||
|
self,
|
||||||
|
name: str,
|
||||||
|
component_cls: Type,
|
||||||
|
category: Optional[str] = None,
|
||||||
|
priority: int = 0,
|
||||||
|
) -> None:
|
||||||
|
"""Register a component class with optional category and priority."""
|
||||||
|
if name in self._entries:
|
||||||
|
raise ValueError(f"Component '{name}' is already registered")
|
||||||
|
self._entries[name] = (component_cls, category, priority)
|
||||||
|
|
||||||
|
def get(self, name: str) -> Type:
|
||||||
|
"""Get component class by name."""
|
||||||
|
if name not in self._entries:
|
||||||
|
raise KeyError(f"Component '{name}' not found in registry")
|
||||||
|
return self._entries[name][0]
|
||||||
|
|
||||||
|
def get_with_metadata(self, name: str) -> Tuple[Type, Optional[str], int]:
|
||||||
|
"""Get component class with its metadata."""
|
||||||
|
entry = self._entries.get(name)
|
||||||
|
if entry is None:
|
||||||
|
raise KeyError(f"Component '{name}' not found in registry")
|
||||||
|
return entry
|
||||||
|
|
||||||
|
def contains(self, name: str) -> bool:
|
||||||
|
"""Check if a name is registered."""
|
||||||
|
return name in self._entries
|
||||||
|
|
||||||
|
def list_names(self) -> List[str]:
|
||||||
|
"""Return list of registered component names."""
|
||||||
|
return sorted(self._entries.keys())
|
||||||
|
|
||||||
|
def list_by_category(self, category: str) -> List[str]:
|
||||||
|
"""Return names of components belonging to a specific category."""
|
||||||
|
return sorted(
|
||||||
|
name for name, (_, cat, _) in self._entries.items() if cat == category
|
||||||
|
)
|
||||||
|
|
||||||
|
def list_by_priority(self, reverse: bool = False) -> List[str]:
|
||||||
|
"""Return names sorted by priority (default ascending)."""
|
||||||
|
return sorted(
|
||||||
|
self._entries.keys(),
|
||||||
|
key=lambda name: self._entries[name][2],
|
||||||
|
reverse=reverse,
|
||||||
|
)
|
||||||
|
|
||||||
|
def entries(self) -> Dict[str, Tuple[Type, Optional[str], int]]:
|
||||||
|
"""Return raw entries dictionary."""
|
||||||
|
return self._entries.copy()
|
||||||
|
|
||||||
|
|
||||||
|
class BaseFactory(ABC, Generic[T]):
|
||||||
|
"""Generic factory class for component registration and creation.
|
||||||
|
|
||||||
|
This base class provides a decorator-based registration pattern
|
||||||
|
for creating extensible component factories.
|
||||||
|
|
||||||
|
Example usage:
|
||||||
|
class MyFactory(BaseFactory[MyBaseClass]):
|
||||||
|
pass
|
||||||
|
|
||||||
|
@MyFactory.register("custom")
|
||||||
|
class CustomComponent(MyBaseClass):
|
||||||
|
...
|
||||||
|
|
||||||
|
component = MyFactory.create("custom", *args, **kwargs)
|
||||||
|
"""
|
||||||
|
|
||||||
|
_registry: Registry
|
||||||
|
|
||||||
|
def __init_subclass__(cls, **kwargs):
|
||||||
|
super().__init_subclass__(**kwargs)
|
||||||
|
cls._registry = Registry()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def register(
|
||||||
|
cls, name: str, category: Optional[str] = None, priority: int = 0
|
||||||
|
) -> Callable[[Type[T]], Type[T]]:
|
||||||
|
"""Decorator to register a component class with optional category and priority.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
name: Registration name for the component
|
||||||
|
category: Optional category for grouping components
|
||||||
|
priority: Priority for ordering (default 0)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Decorator function that registers the component class
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
TypeError: If the decorated class doesn't inherit from the base type
|
||||||
|
"""
|
||||||
|
|
||||||
|
def decorator(component_cls: Type[T]) -> Type[T]:
|
||||||
|
cls._validate_component(component_cls)
|
||||||
|
cls._registry.register(
|
||||||
|
name, component_cls, category=category, priority=priority
|
||||||
|
)
|
||||||
|
return component_cls
|
||||||
|
|
||||||
|
return decorator
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(cls, name: str, *args, **kwargs) -> T:
|
||||||
|
"""Create a component instance by name.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
name: Registered name of the component
|
||||||
|
*args: Positional arguments passed to component constructor
|
||||||
|
**kwargs: Keyword arguments passed to component constructor
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Component instance
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If the component name is not registered
|
||||||
|
"""
|
||||||
|
if not cls._registry.contains(name):
|
||||||
|
raise ValueError(
|
||||||
|
f"Unknown component: '{name}'. "
|
||||||
|
f"Supported types: {sorted(cls._registry.list_names())}"
|
||||||
|
)
|
||||||
|
component_cls = cls._registry.get(name)
|
||||||
|
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
|
||||||
|
def list_registered(cls) -> list:
|
||||||
|
"""List all registered component names.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List of registered component names
|
||||||
|
"""
|
||||||
|
return cls._registry.list_names()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def is_registered(cls, name: str) -> bool:
|
||||||
|
"""Check if a component name is registered.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
name: Component name to check
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
True if registered, False otherwise
|
||||||
|
"""
|
||||||
|
return cls._registry.contains(name)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def list_by_category(cls, category: str) -> List[str]:
|
||||||
|
"""List registered component names in a category."""
|
||||||
|
return cls._registry.list_by_category(category)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def list_by_priority(cls, reverse: bool = False) -> List[str]:
|
||||||
|
"""List registered component names sorted by priority."""
|
||||||
|
return cls._registry.list_by_priority(reverse)
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["Registry", "BaseFactory"]
|
||||||
@@ -0,0 +1,46 @@
|
|||||||
|
"""Inference module for continuous batching.
|
||||||
|
|
||||||
|
Layers:
|
||||||
|
- engine.py: Facade (InferenceEngine), Value Object (GenerationParams, GenerationRequest)
|
||||||
|
- scheduler.py: Continuous-batching loop, Task state machine, TaskStatus enum
|
||||||
|
- cache.py: PagedCache (page-table-indirected KV cache with alloc/free)
|
||||||
|
- sampling.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy)
|
||||||
|
- server.py: FastAPI HTTP server (OpenAI-compatible endpoints)
|
||||||
|
"""
|
||||||
|
|
||||||
|
from astrai.inference.engine import (
|
||||||
|
GenerationParams,
|
||||||
|
GenerationRequest,
|
||||||
|
InferenceEngine,
|
||||||
|
)
|
||||||
|
from astrai.inference.sampling import (
|
||||||
|
BaseSamplingStrategy,
|
||||||
|
SamplingPipeline,
|
||||||
|
TemperatureStrategy,
|
||||||
|
TopKStrategy,
|
||||||
|
TopPStrategy,
|
||||||
|
sample,
|
||||||
|
)
|
||||||
|
from astrai.inference.scheduler import (
|
||||||
|
InferenceScheduler,
|
||||||
|
Task,
|
||||||
|
TaskStatus,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
# Engine / Requests
|
||||||
|
"InferenceEngine",
|
||||||
|
"GenerationRequest",
|
||||||
|
"GenerationParams",
|
||||||
|
# Scheduler
|
||||||
|
"InferenceScheduler",
|
||||||
|
"Task",
|
||||||
|
"TaskStatus",
|
||||||
|
# Sampling (Strategy pattern)
|
||||||
|
"sample",
|
||||||
|
"BaseSamplingStrategy",
|
||||||
|
"TemperatureStrategy",
|
||||||
|
"TopKStrategy",
|
||||||
|
"TopPStrategy",
|
||||||
|
"SamplingPipeline",
|
||||||
|
]
|
||||||
@@ -0,0 +1,174 @@
|
|||||||
|
"""Page-based KV cache with page-table-indirected read/write.
|
||||||
|
|
||||||
|
Provides:
|
||||||
|
- PagedCache: paged KV cache combining page pool and tensor storage.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Dict, List, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
STOP = object()
|
||||||
|
|
||||||
|
|
||||||
|
def page_hash(token_ids: List[int], page_idx: int, page_size: int) -> int:
|
||||||
|
start = page_idx * page_size
|
||||||
|
end = min(start + page_size, len(token_ids))
|
||||||
|
h = 0
|
||||||
|
for i in range(start, end):
|
||||||
|
h = (h * 31 + token_ids[i]) & 0xFFFFFFFFFFFFFFFF
|
||||||
|
return h
|
||||||
|
|
||||||
|
|
||||||
|
class PagedCache:
|
||||||
|
"""Paged KV cache with page-table-indirected read/write.
|
||||||
|
|
||||||
|
Combines:
|
||||||
|
- Page pool (ref-counted alloc/free via bitmask)
|
||||||
|
- KV tensor storage (k_cache, v_cache)
|
||||||
|
- Prefix-cache hash lookup (page_content_hash -> physical_page_idx)
|
||||||
|
|
||||||
|
Call :meth:`bind` to obtain a batch view for the attention layers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
n_layers: int,
|
||||||
|
n_pages: int,
|
||||||
|
page_size: int,
|
||||||
|
n_kv_heads: int,
|
||||||
|
head_dim: int,
|
||||||
|
device: torch.device,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
):
|
||||||
|
self.page_size = page_size
|
||||||
|
self._free_mask = (1 << n_pages) - 1
|
||||||
|
self._refs: List[int] = [0] * n_pages
|
||||||
|
self.k_cache = torch.empty(
|
||||||
|
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
|
||||||
|
device=device,
|
||||||
|
dtype=dtype,
|
||||||
|
)
|
||||||
|
self.v_cache = torch.empty(
|
||||||
|
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
|
||||||
|
device=device,
|
||||||
|
dtype=dtype,
|
||||||
|
)
|
||||||
|
self._page_to_hash: Dict[int, int] = {}
|
||||||
|
self._hash_to_page: Dict[int, int] = {}
|
||||||
|
|
||||||
|
def record_page(
|
||||||
|
self, page_idx: int, token_ids: List[int], logical_page_idx: int
|
||||||
|
) -> None:
|
||||||
|
h = page_hash(token_ids, logical_page_idx, self.page_size)
|
||||||
|
old_h = self._page_to_hash.pop(page_idx, None)
|
||||||
|
if old_h is not None:
|
||||||
|
self._hash_to_page.pop(old_h, None)
|
||||||
|
self._page_to_hash[page_idx] = h
|
||||||
|
self._hash_to_page[h] = page_idx
|
||||||
|
|
||||||
|
def lookup_prefix(self, token_ids: List[int]) -> List[int]:
|
||||||
|
full_pages = len(token_ids) // self.page_size
|
||||||
|
hits: List[int] = []
|
||||||
|
for i in range(full_pages):
|
||||||
|
h = page_hash(token_ids, i, self.page_size)
|
||||||
|
p = self._hash_to_page.get(h)
|
||||||
|
if p is None:
|
||||||
|
break
|
||||||
|
hits.append(p)
|
||||||
|
return hits
|
||||||
|
|
||||||
|
def inc_ref(self, idx: int) -> None:
|
||||||
|
self._refs[idx] += 1
|
||||||
|
|
||||||
|
def alloc(self) -> int:
|
||||||
|
lsb = self._free_mask & -self._free_mask
|
||||||
|
if lsb == 0:
|
||||||
|
return -1
|
||||||
|
idx = lsb.bit_length() - 1
|
||||||
|
self._free_mask ^= lsb
|
||||||
|
self._refs[idx] = 1
|
||||||
|
return idx
|
||||||
|
|
||||||
|
def alloc_n(self, n: int) -> List[int]:
|
||||||
|
pages = [self.alloc() for _ in range(n)]
|
||||||
|
if any(p < 0 for p in pages):
|
||||||
|
for p in pages:
|
||||||
|
if p >= 0:
|
||||||
|
self.free(p)
|
||||||
|
return []
|
||||||
|
return pages
|
||||||
|
|
||||||
|
def free(self, idx: int) -> None:
|
||||||
|
self._refs[idx] -= 1
|
||||||
|
if self._refs[idx] == 0:
|
||||||
|
self._free_mask |= 1 << idx
|
||||||
|
h = self._page_to_hash.pop(idx, None)
|
||||||
|
if h is not None:
|
||||||
|
self._hash_to_page.pop(h, None)
|
||||||
|
|
||||||
|
def bind(self, page_table: Tensor, total_len: int = 0) -> "CacheView":
|
||||||
|
return CacheView(self, page_table, total_len)
|
||||||
|
|
||||||
|
def write(
|
||||||
|
self, layer_id: int, page_table: Tensor, start_pos: int, k: Tensor, v: Tensor
|
||||||
|
) -> None:
|
||||||
|
seq_len = k.size(1)
|
||||||
|
if seq_len == 0:
|
||||||
|
return
|
||||||
|
page_size = self.page_size
|
||||||
|
written = 0
|
||||||
|
first_page = start_pos // page_size
|
||||||
|
last_page = (start_pos + seq_len - 1) // page_size
|
||||||
|
for pi in range(first_page, last_page + 1):
|
||||||
|
phys_pages = page_table[:, pi]
|
||||||
|
page_start = pi * page_size
|
||||||
|
write_start = max(page_start, start_pos)
|
||||||
|
write_end = min(page_start + page_size, start_pos + seq_len)
|
||||||
|
offset = write_start - page_start
|
||||||
|
chunk = write_end - write_start
|
||||||
|
self.k_cache[layer_id, phys_pages, offset : offset + chunk] = k[
|
||||||
|
:, written : written + chunk
|
||||||
|
]
|
||||||
|
self.v_cache[layer_id, phys_pages, offset : offset + chunk] = v[
|
||||||
|
:, written : written + chunk
|
||||||
|
]
|
||||||
|
written += chunk
|
||||||
|
|
||||||
|
def gather(self, layer_id: int, page_table: Tensor) -> Tuple[Tensor, Tensor]:
|
||||||
|
k_parts, v_parts = [], []
|
||||||
|
for pi in range(page_table.size(1)):
|
||||||
|
phys_pages = page_table[:, pi]
|
||||||
|
if not (phys_pages >= 0).any():
|
||||||
|
break
|
||||||
|
k_parts.append(self.k_cache[layer_id, phys_pages])
|
||||||
|
v_parts.append(self.v_cache[layer_id, phys_pages])
|
||||||
|
k = torch.cat(k_parts, dim=1)
|
||||||
|
v = torch.cat(v_parts, dim=1)
|
||||||
|
return k, v
|
||||||
|
|
||||||
|
|
||||||
|
class CacheView:
|
||||||
|
"""Per-batch view that bundles PagedCache + page_table + total_len.
|
||||||
|
|
||||||
|
Attention layers receive this as ``paged_cache`` and only see
|
||||||
|
``write()`` / ``gather()``, never raw page tables or length params.
|
||||||
|
"""
|
||||||
|
|
||||||
|
__slots__ = ("_cache", "_page_table", "_total_len")
|
||||||
|
|
||||||
|
def __init__(self, cache: PagedCache, page_table: Tensor, total_len: int = 0):
|
||||||
|
self._cache = cache
|
||||||
|
self._page_table = page_table
|
||||||
|
self._total_len = total_len
|
||||||
|
|
||||||
|
def write(self, layer_id: int, start_pos: int, k: Tensor, v: Tensor) -> None:
|
||||||
|
self._cache.write(layer_id, self._page_table, start_pos, k, v)
|
||||||
|
|
||||||
|
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
||||||
|
k, v = self._cache.gather(layer_id, self._page_table)
|
||||||
|
if self._total_len:
|
||||||
|
k = k[:, : self._total_len]
|
||||||
|
v = v[:, : self._total_len]
|
||||||
|
return k, v
|
||||||
@@ -0,0 +1,460 @@
|
|||||||
|
"""Unified inference engine for continuous batching.
|
||||||
|
|
||||||
|
Layers:
|
||||||
|
- GenerationParams: Immutable value object for sampling parameters.
|
||||||
|
- GenerationRequest: User-facing request DTO with validation.
|
||||||
|
- _Result: Thread-safe token accumulator (Observer pattern).
|
||||||
|
- InferenceEngine: Facade over InferenceScheduler + async wrapper.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import gc
|
||||||
|
import threading
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
from astrai.inference.cache import STOP
|
||||||
|
from astrai.inference.scheduler import InferenceScheduler
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class GenerationParams:
|
||||||
|
"""Immutable value object for sampling hyperparameters."""
|
||||||
|
|
||||||
|
top_k: int = 50
|
||||||
|
top_p: float = 1.0
|
||||||
|
temperature: float = 1.0
|
||||||
|
max_tokens: int = 1024
|
||||||
|
|
||||||
|
|
||||||
|
class GenerationRequest:
|
||||||
|
"""Request parameters for text generation.
|
||||||
|
|
||||||
|
Encapsulates messages, sampling parameters (via GenerationParams),
|
||||||
|
and streaming preference for a single generation request.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
messages: List[Dict[str, str]],
|
||||||
|
top_k: int = 50,
|
||||||
|
top_p: float = 1.0,
|
||||||
|
temperature: float = 1.0,
|
||||||
|
max_len: int = 1024,
|
||||||
|
stream: bool = False,
|
||||||
|
):
|
||||||
|
"""Initializes a generation request.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
messages: Conversation history as list of {"role": ..., "content": ...}.
|
||||||
|
top_k: Top-k sampling count (0 disables).
|
||||||
|
top_p: Nucleus sampling probability threshold.
|
||||||
|
temperature: Sampling temperature.
|
||||||
|
max_len: Maximum tokens to generate.
|
||||||
|
stream: Whether to return output as a token stream.
|
||||||
|
"""
|
||||||
|
self.messages = messages
|
||||||
|
self.params = GenerationParams(
|
||||||
|
top_k=top_k,
|
||||||
|
top_p=top_p,
|
||||||
|
temperature=temperature,
|
||||||
|
max_tokens=max_len,
|
||||||
|
)
|
||||||
|
self.stream = stream
|
||||||
|
self._validate()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def top_k(self) -> int:
|
||||||
|
return self.params.top_k
|
||||||
|
|
||||||
|
@property
|
||||||
|
def top_p(self) -> float:
|
||||||
|
return self.params.top_p
|
||||||
|
|
||||||
|
@property
|
||||||
|
def temperature(self) -> float:
|
||||||
|
return self.params.temperature
|
||||||
|
|
||||||
|
@property
|
||||||
|
def max_len(self) -> int:
|
||||||
|
return self.params.max_tokens
|
||||||
|
|
||||||
|
def _validate(self):
|
||||||
|
"""Validates sampling parameter ranges."""
|
||||||
|
if not (isinstance(self.top_k, int) and self.top_k >= 0):
|
||||||
|
raise ValueError("top_k must be a non-negative integer")
|
||||||
|
if not (0.0 <= self.top_p <= 1.0):
|
||||||
|
raise ValueError("top_p must be a float between 0.0 and 1.0")
|
||||||
|
if not (isinstance(self.temperature, (int, float)) and self.temperature >= 0):
|
||||||
|
raise ValueError("temperature must be a non-negative number")
|
||||||
|
|
||||||
|
|
||||||
|
class _Result:
|
||||||
|
"""Thread-safe token accumulator for streaming and non-streaming modes.
|
||||||
|
|
||||||
|
Supports multiple concurrent generation tasks with per-index result tracking.
|
||||||
|
Uses a threading.Condition for efficient completion notification
|
||||||
|
and a threading.Event for streaming wakeup.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, count: int = 1):
|
||||||
|
"""Initializes the accumulator.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
count: Number of concurrent generation tasks to track.
|
||||||
|
"""
|
||||||
|
self._cond = threading.Condition()
|
||||||
|
self._event = threading.Event()
|
||||||
|
self.tokens: List[str] = []
|
||||||
|
self.results: List[str] = [""] * count
|
||||||
|
self._done: List[bool] = [False] * count
|
||||||
|
self._completed = 0
|
||||||
|
self._total = count
|
||||||
|
|
||||||
|
def append(self, token: str, idx: int = 0):
|
||||||
|
"""Appends a token to the result buffer.
|
||||||
|
|
||||||
|
In non-streaming mode, tokens are concatenated into results[idx].
|
||||||
|
The sentinel STOP marks a task as complete.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
token: The decoded token string, or STOP sentinel.
|
||||||
|
idx: Index of the generation task this token belongs to.
|
||||||
|
"""
|
||||||
|
with self._cond:
|
||||||
|
self.tokens.append(token)
|
||||||
|
if token is not STOP:
|
||||||
|
self.results[idx] += token
|
||||||
|
else:
|
||||||
|
if not self._done[idx]:
|
||||||
|
self._done[idx] = True
|
||||||
|
self._completed += 1
|
||||||
|
self._cond.notify_all()
|
||||||
|
self._event.set()
|
||||||
|
|
||||||
|
def pop_all(self) -> List[str]:
|
||||||
|
"""Returns and clears all accumulated tokens.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List of token strings since the last call.
|
||||||
|
"""
|
||||||
|
with self._cond:
|
||||||
|
out = self.tokens.copy()
|
||||||
|
self.tokens.clear()
|
||||||
|
if not out:
|
||||||
|
self._event.clear()
|
||||||
|
return out
|
||||||
|
|
||||||
|
def wait(self, timeout: Optional[float] = None) -> bool:
|
||||||
|
"""Blocks until new tokens arrive or the timeout expires.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
timeout: Maximum wait time in seconds (None = infinite).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
True if the event was set (new data available), False on timeout.
|
||||||
|
"""
|
||||||
|
return self._event.wait(timeout=timeout)
|
||||||
|
|
||||||
|
def wait_completion(self) -> None:
|
||||||
|
"""Blocks until all tasks complete (non-streaming).
|
||||||
|
|
||||||
|
Uses a Condition to sleep efficiently instead of busy-waiting.
|
||||||
|
The calling thread is parked until a STOP signal arrives.
|
||||||
|
"""
|
||||||
|
with self._cond:
|
||||||
|
self._cond.wait_for(lambda: self._completed >= self._total)
|
||||||
|
|
||||||
|
def get_results(self) -> List[str]:
|
||||||
|
"""Returns all accumulated results for non-streaming mode.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List of complete generated strings, one per task index.
|
||||||
|
"""
|
||||||
|
with self._cond:
|
||||||
|
return self.results.copy()
|
||||||
|
|
||||||
|
|
||||||
|
class InferenceEngine:
|
||||||
|
"""Unified inference engine backed by continuous-batching scheduler.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
with InferenceEngine(model, tokenizer) as engine:
|
||||||
|
for token in engine.generate("hello", stream=True):
|
||||||
|
print(token, end="")
|
||||||
|
|
||||||
|
text = engine.generate("hello")
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: nn.Module,
|
||||||
|
tokenizer: AutoTokenizer,
|
||||||
|
max_batch_size: int = 1,
|
||||||
|
max_seq_len: Optional[int] = None,
|
||||||
|
max_prompt_len: int = 2048,
|
||||||
|
page_size: int = 128,
|
||||||
|
):
|
||||||
|
"""Initializes the inference engine.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model: The model instance.
|
||||||
|
tokenizer: The tokenizer instance.
|
||||||
|
max_batch_size: Maximum number of concurrent tasks.
|
||||||
|
max_seq_len: Maximum sequence length.
|
||||||
|
max_prompt_len: Maximum prompt tokens.
|
||||||
|
compile: Whether to compile the model with torch.compile.
|
||||||
|
page_size: Number of tokens per KV cache page.
|
||||||
|
"""
|
||||||
|
self.model = model
|
||||||
|
self.tokenizer = tokenizer
|
||||||
|
self.scheduler = InferenceScheduler(
|
||||||
|
model=self.model,
|
||||||
|
tokenizer=self.tokenizer,
|
||||||
|
max_batch_size=max_batch_size,
|
||||||
|
max_seq_len=max_seq_len,
|
||||||
|
max_prompt_len=max_prompt_len,
|
||||||
|
page_size=page_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.scheduler.start()
|
||||||
|
|
||||||
|
def __enter__(self):
|
||||||
|
return self
|
||||||
|
|
||||||
|
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||||
|
self.shutdown()
|
||||||
|
return False
|
||||||
|
|
||||||
|
def generate(
|
||||||
|
self,
|
||||||
|
prompt: Union[str, List[str]],
|
||||||
|
stream: bool = False,
|
||||||
|
max_tokens: int = 1024,
|
||||||
|
temperature: float = 1.0,
|
||||||
|
top_p: float = 1.0,
|
||||||
|
top_k: int = 50,
|
||||||
|
) -> Union[Generator[str, None, None], str, List[str]]:
|
||||||
|
"""Generates text from a prompt.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prompt: Single string or list of strings for batch generation.
|
||||||
|
stream: If True, returns a generator yielding tokens one by one.
|
||||||
|
max_tokens: Maximum number of tokens to generate.
|
||||||
|
temperature: Sampling temperature.
|
||||||
|
top_p: Nucleus sampling probability threshold.
|
||||||
|
top_k: Top-k sampling count (0 disables).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Generator (stream=True), single string (non-stream, single prompt),
|
||||||
|
or list of strings (non-stream, batch prompts).
|
||||||
|
"""
|
||||||
|
is_batch = isinstance(prompt, list)
|
||||||
|
prompts = prompt if is_batch else [prompt]
|
||||||
|
|
||||||
|
if stream:
|
||||||
|
return self._generate_streaming(
|
||||||
|
prompts, is_batch, max_tokens, temperature, top_p, top_k
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
return self._generate_non_streaming(
|
||||||
|
prompts, is_batch, max_tokens, temperature, top_p, top_k
|
||||||
|
)
|
||||||
|
|
||||||
|
def generate_async(
|
||||||
|
self,
|
||||||
|
prompt: str,
|
||||||
|
max_tokens: int = 1024,
|
||||||
|
temperature: float = 1.0,
|
||||||
|
top_p: float = 1.0,
|
||||||
|
top_k: int = 50,
|
||||||
|
) -> AsyncGenerator[str, None]:
|
||||||
|
"""Async streaming generator that does not block the event loop.
|
||||||
|
|
||||||
|
Runs the synchronous generator in a background thread pool executor,
|
||||||
|
yielding tokens to the async consumer as they arrive.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prompt: Input text to generate from.
|
||||||
|
max_tokens: Maximum tokens to generate.
|
||||||
|
temperature: Sampling temperature.
|
||||||
|
top_p: Nucleus sampling threshold.
|
||||||
|
top_k: Top-k sampling count.
|
||||||
|
|
||||||
|
Yields:
|
||||||
|
Decoded token strings as they are generated.
|
||||||
|
"""
|
||||||
|
sync_gen = self._generate_streaming(
|
||||||
|
[prompt], False, max_tokens, temperature, top_p, top_k
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _agen():
|
||||||
|
loop = asyncio.get_event_loop()
|
||||||
|
while True:
|
||||||
|
token = await loop.run_in_executor(None, self._next_token, sync_gen)
|
||||||
|
if token is None:
|
||||||
|
break
|
||||||
|
yield token
|
||||||
|
|
||||||
|
return _agen()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _next_token(gen: Generator) -> Optional[str]:
|
||||||
|
"""Retrieves the next token from a synchronous generator.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
gen: A synchronous generator yielding token strings.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The next token, or None if the generator is exhausted.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
return next(gen)
|
||||||
|
except StopIteration:
|
||||||
|
return None
|
||||||
|
|
||||||
|
def generate_with_request(
|
||||||
|
self, request: GenerationRequest
|
||||||
|
) -> Union[Generator[str, None, None], str, List[str]]:
|
||||||
|
"""Generates text from a structured GenerationRequest.
|
||||||
|
|
||||||
|
Applies the chat template to the request's messages before generation.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
request: A GenerationRequest with messages and parameters.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Generator, string, or list of strings (see generate()).
|
||||||
|
"""
|
||||||
|
prompt = self.tokenizer.apply_chat_template(request.messages, tokenize=False)
|
||||||
|
return self.generate(
|
||||||
|
prompt=prompt,
|
||||||
|
stream=request.stream,
|
||||||
|
max_tokens=request.params.max_tokens,
|
||||||
|
temperature=request.params.temperature,
|
||||||
|
top_p=request.params.top_p,
|
||||||
|
top_k=request.params.top_k,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _generate_streaming(
|
||||||
|
self,
|
||||||
|
prompts: List[str],
|
||||||
|
is_batch: bool,
|
||||||
|
max_tokens: int,
|
||||||
|
temperature: float,
|
||||||
|
top_p: float,
|
||||||
|
top_k: int,
|
||||||
|
) -> Generator[str, None, None]:
|
||||||
|
"""Internal streaming generator.
|
||||||
|
|
||||||
|
Polls the _Result accumulator in a loop, yielding tokens as they arrive.
|
||||||
|
Cleans up the scheduler task on GeneratorExit.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prompts: List of prompts (only first is used; batch not yet supported).
|
||||||
|
is_batch: If True, raises NotImplementedError.
|
||||||
|
max_tokens: Maximum tokens to generate.
|
||||||
|
temperature: Sampling temperature.
|
||||||
|
top_p: Nucleus sampling threshold.
|
||||||
|
top_k: Top-k sampling count.
|
||||||
|
|
||||||
|
Yields:
|
||||||
|
Decoded token strings.
|
||||||
|
"""
|
||||||
|
if is_batch:
|
||||||
|
raise NotImplementedError("Batch streaming not yet supported")
|
||||||
|
|
||||||
|
result = _Result()
|
||||||
|
|
||||||
|
task_id = self.scheduler.add_task(
|
||||||
|
prompt=prompts[0],
|
||||||
|
max_tokens=max_tokens,
|
||||||
|
temperature=temperature,
|
||||||
|
top_p=top_p,
|
||||||
|
top_k=top_k,
|
||||||
|
stream_callback=lambda tok: result.append(tok, 0),
|
||||||
|
)
|
||||||
|
|
||||||
|
def gen():
|
||||||
|
try:
|
||||||
|
while True:
|
||||||
|
tokens = result.pop_all()
|
||||||
|
for token in tokens:
|
||||||
|
if token is STOP:
|
||||||
|
return
|
||||||
|
yield token
|
||||||
|
if not result.wait(timeout=0.05):
|
||||||
|
pass
|
||||||
|
finally:
|
||||||
|
self.scheduler.remove_task(task_id)
|
||||||
|
|
||||||
|
return gen()
|
||||||
|
|
||||||
|
def _generate_non_streaming(
|
||||||
|
self,
|
||||||
|
prompts: List[str],
|
||||||
|
is_batch: bool,
|
||||||
|
max_tokens: int,
|
||||||
|
temperature: float,
|
||||||
|
top_p: float,
|
||||||
|
top_k: int,
|
||||||
|
) -> Union[str, List[str]]:
|
||||||
|
"""Internal non-streaming generator.
|
||||||
|
|
||||||
|
Submits all prompts to the scheduler and waits for all to complete.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prompts: List of prompt strings.
|
||||||
|
is_batch: Whether multiple prompts were provided.
|
||||||
|
max_tokens: Maximum tokens to generate.
|
||||||
|
temperature: Sampling temperature.
|
||||||
|
top_p: Nucleus sampling threshold.
|
||||||
|
top_k: Top-k sampling count.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Single string for one prompt, list of strings for batch.
|
||||||
|
"""
|
||||||
|
result = _Result(count=len(prompts))
|
||||||
|
task_ids = []
|
||||||
|
|
||||||
|
for i, p in enumerate(prompts):
|
||||||
|
|
||||||
|
def make_cb(idx):
|
||||||
|
return lambda tok: result.append(tok, idx)
|
||||||
|
|
||||||
|
task_id = self.scheduler.add_task(
|
||||||
|
prompt=p,
|
||||||
|
max_tokens=max_tokens,
|
||||||
|
temperature=temperature,
|
||||||
|
top_p=top_p,
|
||||||
|
top_k=top_k,
|
||||||
|
stream_callback=make_cb(i),
|
||||||
|
)
|
||||||
|
task_ids.append(task_id)
|
||||||
|
|
||||||
|
result.wait_completion()
|
||||||
|
|
||||||
|
for task_id in task_ids:
|
||||||
|
self.scheduler.remove_task(task_id)
|
||||||
|
|
||||||
|
res = result.get_results()
|
||||||
|
return res if is_batch else res[0]
|
||||||
|
|
||||||
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
|
"""Returns current engine statistics.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict with total_tasks, total_tokens, active_tasks, waiting_queue.
|
||||||
|
"""
|
||||||
|
return self.scheduler.get_stats()
|
||||||
|
|
||||||
|
def shutdown(self) -> None:
|
||||||
|
"""Shuts down the engine, stops the scheduler, and frees GPU memory."""
|
||||||
|
self.scheduler.stop()
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
gc.collect()
|
||||||
@@ -0,0 +1,178 @@
|
|||||||
|
"""Composable sampling strategies for logit transformation.
|
||||||
|
|
||||||
|
Implements the Strategy pattern: each sampling technique
|
||||||
|
(temperature, top-k, top-p) is a pluggable strategy that
|
||||||
|
can be composed into a pipeline.
|
||||||
|
|
||||||
|
All strategies accept both scalar and per-sample tensor
|
||||||
|
parameters, so a single pipeline works for any batch size.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from typing import List, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
class BaseSamplingStrategy(ABC):
|
||||||
|
"""Abstract base for a logit transformation strategy."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
||||||
|
"""Applies the strategy to logits.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
logits: Raw logits tensor (batch, vocab_size).
|
||||||
|
filter_value: Value assigned to filtered-out positions.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Transformed logits tensor.
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
class TemperatureStrategy(BaseSamplingStrategy):
|
||||||
|
"""Divides logits by temperature to control randomness.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
temperature: Scalar or ``[batch]`` tensor.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, temperature: Union[float, Tensor] = 1.0):
|
||||||
|
self.temperature = temperature
|
||||||
|
|
||||||
|
def apply(self, logits, filter_value=-float("inf")):
|
||||||
|
t = self.temperature
|
||||||
|
if isinstance(t, Tensor):
|
||||||
|
if (t != 1.0).any():
|
||||||
|
logits = logits / t.to(logits.device, non_blocking=True).view(-1, 1)
|
||||||
|
elif t != 1.0:
|
||||||
|
logits = logits / t
|
||||||
|
return logits
|
||||||
|
|
||||||
|
|
||||||
|
class TopKStrategy(BaseSamplingStrategy):
|
||||||
|
"""Keeps only the top-k logits, setting the rest to filter_value.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
top_k: Scalar or ``[batch]`` tensor (0 disables).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, top_k: Union[int, Tensor] = 0):
|
||||||
|
self.top_k = top_k
|
||||||
|
|
||||||
|
def apply(self, logits, filter_value=-float("inf")):
|
||||||
|
tk = self.top_k
|
||||||
|
if isinstance(tk, Tensor):
|
||||||
|
max_k = int(tk.max().item())
|
||||||
|
if max_k <= 0:
|
||||||
|
return logits
|
||||||
|
k = min(max_k, logits.size(-1))
|
||||||
|
elif tk > 0:
|
||||||
|
k = min(tk, logits.size(-1))
|
||||||
|
else:
|
||||||
|
return logits
|
||||||
|
thresholds = torch.topk(logits, k, dim=-1)[0][..., -1:]
|
||||||
|
logits[logits < thresholds] = filter_value
|
||||||
|
return logits
|
||||||
|
|
||||||
|
|
||||||
|
class TopPStrategy(BaseSamplingStrategy):
|
||||||
|
"""Nucleus (top-p) filtering: keeps the smallest set of tokens whose
|
||||||
|
cumulative probability exceeds top_p.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
top_p: Scalar or ``[batch]`` tensor (1.0 disables).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, top_p: Union[float, Tensor] = 1.0):
|
||||||
|
self.top_p = top_p
|
||||||
|
|
||||||
|
def _apply(self, logits, top_p, filter_value):
|
||||||
|
sorted_logits, sorted_indices = torch.sort(logits, descending=True, dim=-1)
|
||||||
|
cum_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
|
||||||
|
remove = cum_probs > top_p
|
||||||
|
remove[..., 1:] = remove[..., :-1].clone()
|
||||||
|
remove[..., 0] = False
|
||||||
|
mask = torch.zeros_like(logits, dtype=torch.bool)
|
||||||
|
mask.scatter_(1, sorted_indices, remove)
|
||||||
|
logits[mask] = filter_value
|
||||||
|
return logits
|
||||||
|
|
||||||
|
def apply(self, logits, filter_value=-float("inf")):
|
||||||
|
tp = self.top_p
|
||||||
|
if isinstance(tp, Tensor):
|
||||||
|
tp = tp.to(logits.device, non_blocking=True)
|
||||||
|
if (tp < 1.0).any():
|
||||||
|
logits = self._apply(logits, tp.view(-1, 1), filter_value)
|
||||||
|
elif tp < 1.0:
|
||||||
|
logits = self._apply(logits, tp, filter_value)
|
||||||
|
return logits
|
||||||
|
|
||||||
|
|
||||||
|
class SamplingPipeline(BaseSamplingStrategy):
|
||||||
|
"""Composes multiple sampling strategies into a single transformation.
|
||||||
|
|
||||||
|
Strategies are applied sequentially in the order they are provided,
|
||||||
|
matching the original temperature -> top-k -> top-p ordering.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
pipeline = SamplingPipeline([
|
||||||
|
TemperatureStrategy(0.8),
|
||||||
|
TopKStrategy(50),
|
||||||
|
TopPStrategy(0.95),
|
||||||
|
])
|
||||||
|
logits = pipeline.apply(logits)
|
||||||
|
token = pipeline.sample(logits) # softmax + multinomial
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, strategies: List[BaseSamplingStrategy]):
|
||||||
|
self.strategies = strategies
|
||||||
|
|
||||||
|
def apply(self, logits, filter_value=-float("inf")):
|
||||||
|
for strategy in self.strategies:
|
||||||
|
logits = strategy.apply(logits, filter_value)
|
||||||
|
return logits
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def sample(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
||||||
|
"""Apply strategies then sample (softmax + multinomial).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
logits: Raw logits ``[batch, vocab_size]``.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Sampled token IDs ``[batch]``.
|
||||||
|
"""
|
||||||
|
return torch.multinomial(
|
||||||
|
torch.softmax(self.apply(logits, filter_value), dim=-1),
|
||||||
|
num_samples=1,
|
||||||
|
).squeeze(-1)
|
||||||
|
|
||||||
|
|
||||||
|
@torch.inference_mode()
|
||||||
|
def sample(
|
||||||
|
logits: Tensor,
|
||||||
|
temperature: Union[float, Tensor] = 1.0,
|
||||||
|
top_k: Union[int, Tensor] = 0,
|
||||||
|
top_p: Union[float, Tensor] = 1.0,
|
||||||
|
filter_value: float = -float("inf"),
|
||||||
|
) -> Tensor:
|
||||||
|
"""Apply sampling strategies then sample (softmax + multinomial).
|
||||||
|
|
||||||
|
Shortcut for ``SamplingPipeline(...).sample(logits)``.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
logits: Raw logits ``[batch, vocab_size]``.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Sampled token IDs ``[batch]``.
|
||||||
|
"""
|
||||||
|
return SamplingPipeline(
|
||||||
|
[
|
||||||
|
TemperatureStrategy(temperature),
|
||||||
|
TopKStrategy(top_k),
|
||||||
|
TopPStrategy(top_p),
|
||||||
|
]
|
||||||
|
).sample(logits, filter_value)
|
||||||
@@ -0,0 +1,411 @@
|
|||||||
|
"""Inference scheduler for single-GPU continuous batching with paged KV cache."""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
from enum import Enum
|
||||||
|
from typing import Any, Callable, Dict, List, Optional, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.inference.cache import STOP, PagedCache
|
||||||
|
from astrai.inference.sampling import sample
|
||||||
|
from astrai.model.automodel import AutoModel
|
||||||
|
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class TaskStatus(Enum):
|
||||||
|
"""Task states in the continuous batching lifecycle."""
|
||||||
|
|
||||||
|
PENDING = "pending"
|
||||||
|
RUNNING = "running"
|
||||||
|
FINISHED = "finished"
|
||||||
|
ABORTED = "aborted"
|
||||||
|
|
||||||
|
|
||||||
|
class Task:
|
||||||
|
"""Represents a single generation request with paged KV cache tracking."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
prompt_ids: List[int],
|
||||||
|
max_tokens: int = 1024,
|
||||||
|
temperature: float = 1.0,
|
||||||
|
top_p: float = 1.0,
|
||||||
|
top_k: int = 50,
|
||||||
|
stream_callback: Optional[Callable[[str], None]] = None,
|
||||||
|
):
|
||||||
|
self.task_id = task_id
|
||||||
|
self.prompt_ids = prompt_ids
|
||||||
|
self.max_tokens = max_tokens
|
||||||
|
self.temperature = temperature
|
||||||
|
self.top_p = top_p
|
||||||
|
self.top_k = top_k
|
||||||
|
|
||||||
|
self.status = TaskStatus.PENDING
|
||||||
|
self.output_ids: List[int] = []
|
||||||
|
self.input_tokens: int = 0
|
||||||
|
self.output_tokens: int = 0
|
||||||
|
self.page_table: List[int] = []
|
||||||
|
self.n_pages: int = 0
|
||||||
|
self._prefix_cached_tokens: int = 0
|
||||||
|
self.arrival_time = time.time()
|
||||||
|
self.finish_time: Optional[float] = None
|
||||||
|
self.stream_callback = stream_callback
|
||||||
|
self._pages_freed: bool = False
|
||||||
|
|
||||||
|
@property
|
||||||
|
def next_pos(self) -> int:
|
||||||
|
return self.input_tokens + len(self.output_ids)
|
||||||
|
|
||||||
|
def is_finished(self, stop_ids: List[int]) -> bool:
|
||||||
|
if self.output_tokens >= self.max_tokens:
|
||||||
|
return True
|
||||||
|
if self.output_ids and self.output_ids[-1] in stop_ids:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
class InferenceScheduler:
|
||||||
|
"""Continuous batching scheduler with paged KV cache.
|
||||||
|
|
||||||
|
Runs a background generation loop with four phases per iteration:
|
||||||
|
1. Cleanup finished tasks and release resources.
|
||||||
|
2. Refill active batch from the waiting queue.
|
||||||
|
3. Prefill newly activated tasks.
|
||||||
|
4. Decode the largest same-position group of active tasks.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: AutoModel,
|
||||||
|
tokenizer: AutoTokenizer,
|
||||||
|
max_batch_size: int = 16,
|
||||||
|
max_seq_len: Optional[int] = None,
|
||||||
|
max_prompt_len: int = 512,
|
||||||
|
page_size: int = 64,
|
||||||
|
device: Optional[str] = None,
|
||||||
|
dtype: Optional[torch.dtype] = None,
|
||||||
|
):
|
||||||
|
config = model.config
|
||||||
|
|
||||||
|
self.model = model
|
||||||
|
self.tokenizer = tokenizer
|
||||||
|
self.max_batch_size = max_batch_size
|
||||||
|
self.max_seq_len = max_seq_len or config.max_len
|
||||||
|
self.max_prompt_len = max_prompt_len
|
||||||
|
self.page_size = page_size
|
||||||
|
self.device = device or next(model.parameters()).device
|
||||||
|
self.dtype = dtype or next(model.parameters()).dtype
|
||||||
|
|
||||||
|
n_kv_heads = config.n_kv_heads
|
||||||
|
head_dim = config.dim // config.n_heads
|
||||||
|
n_layers = config.n_layers
|
||||||
|
n_pages = (
|
||||||
|
max_batch_size * (self.max_seq_len + page_size) + page_size - 1
|
||||||
|
) // page_size
|
||||||
|
|
||||||
|
self.page_cache = PagedCache(
|
||||||
|
n_layers,
|
||||||
|
n_pages,
|
||||||
|
page_size,
|
||||||
|
n_kv_heads,
|
||||||
|
head_dim,
|
||||||
|
self.device,
|
||||||
|
self.dtype,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.waiting_queue: List[Task] = []
|
||||||
|
self.active_tasks: List[Task] = []
|
||||||
|
|
||||||
|
self._running = False
|
||||||
|
self._task_event = threading.Event()
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
self._total_tasks = 0
|
||||||
|
self._total_tokens = 0
|
||||||
|
|
||||||
|
def _n_pages_for(self, n_tokens: int) -> int:
|
||||||
|
return (n_tokens + self.page_size - 1) // self.page_size
|
||||||
|
|
||||||
|
def add_task(
|
||||||
|
self,
|
||||||
|
prompt: str,
|
||||||
|
max_tokens: int = 1024,
|
||||||
|
temperature: float = 1.0,
|
||||||
|
top_p: float = 1.0,
|
||||||
|
top_k: int = 50,
|
||||||
|
stream_callback: Optional[Callable[[str], None]] = None,
|
||||||
|
) -> str:
|
||||||
|
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
|
||||||
|
prompt_ids = self.tokenizer.encode(prompt)
|
||||||
|
if len(prompt_ids) > self.max_prompt_len:
|
||||||
|
prompt_ids = prompt_ids[-self.max_prompt_len :]
|
||||||
|
|
||||||
|
task = Task(
|
||||||
|
task_id=task_id,
|
||||||
|
prompt_ids=prompt_ids,
|
||||||
|
max_tokens=max_tokens,
|
||||||
|
temperature=temperature,
|
||||||
|
top_p=top_p,
|
||||||
|
top_k=top_k,
|
||||||
|
stream_callback=stream_callback,
|
||||||
|
)
|
||||||
|
|
||||||
|
with self._lock:
|
||||||
|
self.waiting_queue.append(task)
|
||||||
|
self._total_tasks += 1
|
||||||
|
|
||||||
|
self._task_event.set()
|
||||||
|
return task_id
|
||||||
|
|
||||||
|
def remove_task(self, task_id: str) -> None:
|
||||||
|
with self._lock:
|
||||||
|
removed_active = [t for t in self.active_tasks if t.task_id == task_id]
|
||||||
|
self.waiting_queue = [t for t in self.waiting_queue if t.task_id != task_id]
|
||||||
|
self.active_tasks = [t for t in self.active_tasks if t.task_id != task_id]
|
||||||
|
|
||||||
|
for task in removed_active:
|
||||||
|
if not task._pages_freed:
|
||||||
|
self._free_pages(task.page_table)
|
||||||
|
task.page_table.clear()
|
||||||
|
task.n_pages = 0
|
||||||
|
task._pages_freed = True
|
||||||
|
|
||||||
|
def _free_pages(self, indices: List[int]) -> None:
|
||||||
|
for idx in indices:
|
||||||
|
self.page_cache.free(idx)
|
||||||
|
|
||||||
|
def _record_page_hashes(self, task: Task, start_logical_page: int = 0) -> None:
|
||||||
|
full_pages = len(task.prompt_ids) // self.page_size
|
||||||
|
for i in range(start_logical_page, full_pages):
|
||||||
|
self.page_cache.record_page(task.page_table[i], task.prompt_ids, i)
|
||||||
|
|
||||||
|
def _remove_finished_tasks(self) -> None:
|
||||||
|
finished = []
|
||||||
|
for task in self.active_tasks:
|
||||||
|
if task.is_finished(self.tokenizer.stop_ids):
|
||||||
|
task.status = TaskStatus.FINISHED
|
||||||
|
task.finish_time = time.time()
|
||||||
|
finished.append(task)
|
||||||
|
self._total_tokens += task.output_tokens
|
||||||
|
|
||||||
|
for task in finished:
|
||||||
|
if not task._pages_freed:
|
||||||
|
self._free_pages(task.page_table)
|
||||||
|
task.page_table.clear()
|
||||||
|
task.n_pages = 0
|
||||||
|
task._pages_freed = True
|
||||||
|
|
||||||
|
self.active_tasks = [
|
||||||
|
t for t in self.active_tasks if t.status != TaskStatus.FINISHED
|
||||||
|
]
|
||||||
|
|
||||||
|
def _refill_active_batch(self) -> None:
|
||||||
|
available = self.max_batch_size - len(self.active_tasks)
|
||||||
|
if available <= 0:
|
||||||
|
return
|
||||||
|
|
||||||
|
to_add: List[Task] = []
|
||||||
|
with self._lock:
|
||||||
|
n = min(available, len(self.waiting_queue))
|
||||||
|
for _ in range(n):
|
||||||
|
to_add.append(self.waiting_queue.pop(0))
|
||||||
|
|
||||||
|
failed: List[Task] = []
|
||||||
|
for task in to_add:
|
||||||
|
prompt_len = len(task.prompt_ids)
|
||||||
|
|
||||||
|
hit_pages = self.page_cache.lookup_prefix(task.prompt_ids)
|
||||||
|
cached_tokens = len(hit_pages) * self.page_size
|
||||||
|
for p in hit_pages:
|
||||||
|
self.page_cache.inc_ref(p)
|
||||||
|
|
||||||
|
remaining = prompt_len - cached_tokens
|
||||||
|
n_new = self._n_pages_for(remaining) if remaining > 0 else 0
|
||||||
|
new_pages = self.page_cache.alloc_n(n_new) if n_new > 0 else []
|
||||||
|
|
||||||
|
if remaining > 0 and not new_pages:
|
||||||
|
for p in hit_pages:
|
||||||
|
self.page_cache.free(p)
|
||||||
|
failed.append(task)
|
||||||
|
continue
|
||||||
|
|
||||||
|
task.page_table = hit_pages + new_pages
|
||||||
|
task.n_pages = len(task.page_table)
|
||||||
|
task._prefix_cached_tokens = cached_tokens
|
||||||
|
task.status = TaskStatus.RUNNING
|
||||||
|
self.active_tasks.append(task)
|
||||||
|
|
||||||
|
if failed:
|
||||||
|
with self._lock:
|
||||||
|
self.waiting_queue[:0] = failed
|
||||||
|
|
||||||
|
def _execute_prefill(
|
||||||
|
self, tasks: List[Task], prompt_len: int, start_pos: int = 0
|
||||||
|
) -> None:
|
||||||
|
tasks = sorted(tasks, key=lambda t: t.task_id)
|
||||||
|
batch_sz = len(tasks)
|
||||||
|
|
||||||
|
seq_len = prompt_len - start_pos
|
||||||
|
input_ids = torch.empty(batch_sz, seq_len, dtype=torch.long, device=self.device)
|
||||||
|
input_mask = torch.ones(batch_sz, seq_len, dtype=torch.bool, device=self.device)
|
||||||
|
|
||||||
|
for i, t in enumerate(tasks):
|
||||||
|
input_ids[i] = torch.tensor(
|
||||||
|
t.prompt_ids[start_pos:prompt_len], device=self.device
|
||||||
|
)
|
||||||
|
|
||||||
|
page_tables = self._make_page_table_tensor(tasks)
|
||||||
|
|
||||||
|
with torch.inference_mode():
|
||||||
|
self.model(
|
||||||
|
input_ids,
|
||||||
|
input_mask=input_mask,
|
||||||
|
start_pos=start_pos,
|
||||||
|
paged_cache=self.page_cache.bind(page_tables, total_len=prompt_len),
|
||||||
|
)
|
||||||
|
|
||||||
|
start_logical_page = start_pos // self.page_size
|
||||||
|
for t in tasks:
|
||||||
|
self._record_page_hashes(t, start_logical_page=start_logical_page)
|
||||||
|
|
||||||
|
def _execute_decode(self, tasks: List[Task], start_pos: int) -> None:
|
||||||
|
if not tasks:
|
||||||
|
return
|
||||||
|
|
||||||
|
tasks = sorted(tasks, key=lambda t: t.task_id)
|
||||||
|
batch_sz = len(tasks)
|
||||||
|
|
||||||
|
for t in tasks:
|
||||||
|
self._maybe_alloc_page(t, start_pos)
|
||||||
|
|
||||||
|
input_ids = torch.tensor(
|
||||||
|
[t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks],
|
||||||
|
dtype=torch.long,
|
||||||
|
device=self.device,
|
||||||
|
)
|
||||||
|
|
||||||
|
active_mask = torch.ones((batch_sz, 1), dtype=torch.bool, device=self.device)
|
||||||
|
|
||||||
|
page_tables = self._make_page_table_tensor(tasks)
|
||||||
|
total_len = start_pos + 1
|
||||||
|
|
||||||
|
temperatures = torch.tensor([t.temperature for t in tasks], device=self.device)
|
||||||
|
top_ks = torch.tensor([t.top_k for t in tasks], device=self.device)
|
||||||
|
top_ps = torch.tensor([t.top_p for t in tasks], device=self.device)
|
||||||
|
|
||||||
|
with torch.inference_mode():
|
||||||
|
outputs = self.model(
|
||||||
|
input_ids.unsqueeze(1),
|
||||||
|
input_mask=active_mask,
|
||||||
|
paged_cache=self.page_cache.bind(page_tables, total_len=total_len),
|
||||||
|
start_pos=start_pos,
|
||||||
|
)
|
||||||
|
logits = outputs["logits"][:, -1, :]
|
||||||
|
|
||||||
|
next_tokens = sample(
|
||||||
|
logits,
|
||||||
|
temperature=temperatures,
|
||||||
|
top_k=top_ks,
|
||||||
|
top_p=top_ps,
|
||||||
|
).tolist()
|
||||||
|
|
||||||
|
for t, ntok in zip(tasks, next_tokens):
|
||||||
|
t.output_ids.append(ntok)
|
||||||
|
t.output_tokens += 1
|
||||||
|
pos = t.input_tokens + t.output_tokens
|
||||||
|
self._maybe_alloc_page(t, pos)
|
||||||
|
if t.stream_callback:
|
||||||
|
t.stream_callback(self.tokenizer.decode([ntok]))
|
||||||
|
|
||||||
|
for t in tasks:
|
||||||
|
if t.is_finished(self.tokenizer.stop_ids):
|
||||||
|
if t.stream_callback:
|
||||||
|
t.stream_callback(STOP)
|
||||||
|
|
||||||
|
def _make_page_table_tensor(self, tasks: List[Task]) -> Tensor:
|
||||||
|
max_pages = max(t.n_pages for t in tasks)
|
||||||
|
rows = [t.page_table + [-1] * (max_pages - t.n_pages) for t in tasks]
|
||||||
|
return torch.tensor(rows, dtype=torch.long, device=self.device)
|
||||||
|
|
||||||
|
def _maybe_alloc_page(self, task: Task, pos: int) -> None:
|
||||||
|
needed = self._n_pages_for(pos + 1)
|
||||||
|
while task.n_pages < needed:
|
||||||
|
p = self.page_cache.alloc()
|
||||||
|
if p < 0:
|
||||||
|
break
|
||||||
|
task.page_table.append(p)
|
||||||
|
task.n_pages += 1
|
||||||
|
|
||||||
|
def _run_generation_loop(self) -> None:
|
||||||
|
try:
|
||||||
|
while self._running:
|
||||||
|
self._remove_finished_tasks()
|
||||||
|
self._refill_active_batch()
|
||||||
|
|
||||||
|
if not self.active_tasks and not self.waiting_queue:
|
||||||
|
self._task_event.clear()
|
||||||
|
self._task_event.wait(timeout=1.0)
|
||||||
|
continue
|
||||||
|
|
||||||
|
to_prefill = [t for t in self.active_tasks if t.output_tokens == 0]
|
||||||
|
if to_prefill:
|
||||||
|
for t in to_prefill:
|
||||||
|
t.input_tokens = len(t.prompt_ids)
|
||||||
|
|
||||||
|
groups: Dict[Tuple[int, int], List[Task]] = {}
|
||||||
|
for t in to_prefill:
|
||||||
|
key = (len(t.prompt_ids), t._prefix_cached_tokens)
|
||||||
|
groups.setdefault(key, []).append(t)
|
||||||
|
|
||||||
|
for (prompt_len, start_pos), group in groups.items():
|
||||||
|
if start_pos < prompt_len:
|
||||||
|
self._execute_prefill(group, prompt_len, start_pos)
|
||||||
|
|
||||||
|
pos_groups: Dict[int, List[Task]] = {}
|
||||||
|
for t in self.active_tasks:
|
||||||
|
pos_groups.setdefault(t.next_pos, []).append(t)
|
||||||
|
|
||||||
|
if pos_groups:
|
||||||
|
best_pos = max(pos_groups, key=lambda p: len(pos_groups[p]))
|
||||||
|
self._execute_decode(pos_groups[best_pos], best_pos)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
||||||
|
for task in self.active_tasks:
|
||||||
|
if task.stream_callback:
|
||||||
|
task.stream_callback(STOP)
|
||||||
|
for task in self.waiting_queue:
|
||||||
|
if task.stream_callback:
|
||||||
|
task.stream_callback(STOP)
|
||||||
|
raise
|
||||||
|
|
||||||
|
def start(self) -> None:
|
||||||
|
if not self._running:
|
||||||
|
self._running = True
|
||||||
|
t = threading.Thread(target=self._run_generation_loop, daemon=True)
|
||||||
|
t.start()
|
||||||
|
self._loop_thread = t
|
||||||
|
|
||||||
|
def stop(self) -> None:
|
||||||
|
self._running = False
|
||||||
|
self._task_event.set()
|
||||||
|
if hasattr(self, "_loop_thread"):
|
||||||
|
self._loop_thread.join(timeout=2.0)
|
||||||
|
self.waiting_queue.clear()
|
||||||
|
self.active_tasks.clear()
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"total_tasks": self._total_tasks,
|
||||||
|
"total_tokens": self._total_tokens,
|
||||||
|
"active_tasks": len(self.active_tasks),
|
||||||
|
"waiting_queue": len(self.waiting_queue),
|
||||||
|
}
|
||||||
@@ -0,0 +1,486 @@
|
|||||||
|
"""
|
||||||
|
OpenAI / Anthropic-compatible chat completion server backed by continuous-batching inference.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
from contextlib import asynccontextmanager
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import uvicorn
|
||||||
|
from fastapi import FastAPI, HTTPException
|
||||||
|
from fastapi.responses import StreamingResponse
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
from astrai.inference.engine import InferenceEngine
|
||||||
|
from astrai.model import AutoModel
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_project_root = Path(__file__).parent.parent.parent
|
||||||
|
|
||||||
|
|
||||||
|
class ServerState:
|
||||||
|
def __init__(self):
|
||||||
|
self.engine: Optional[InferenceEngine] = None
|
||||||
|
self.config: Dict[str, Any] = {
|
||||||
|
"device": "cuda",
|
||||||
|
"dtype": torch.bfloat16,
|
||||||
|
"param_path": None,
|
||||||
|
"max_batch_size": 16,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
_state = ServerState()
|
||||||
|
|
||||||
|
|
||||||
|
class ChatMessage(BaseModel):
|
||||||
|
role: str
|
||||||
|
content: str
|
||||||
|
|
||||||
|
|
||||||
|
class ChatCompletionRequest(BaseModel):
|
||||||
|
"""OpenAI Chat Completion API request body."""
|
||||||
|
|
||||||
|
model: str = "astrai"
|
||||||
|
messages: List[ChatMessage]
|
||||||
|
temperature: Optional[float] = Field(default=1.0, ge=0.0, le=2.0)
|
||||||
|
top_p: Optional[float] = Field(default=1.0, ge=0.0, le=1.0)
|
||||||
|
top_k: Optional[int] = Field(default=50, ge=1)
|
||||||
|
stream: Optional[bool] = False
|
||||||
|
stop: Optional[Union[str, List[str]]] = None
|
||||||
|
max_tokens: Optional[int] = Field(default=2048, ge=1)
|
||||||
|
n: Optional[int] = Field(default=1, ge=1)
|
||||||
|
presence_penalty: Optional[float] = Field(default=0.0, ge=-2.0, le=2.0)
|
||||||
|
frequency_penalty: Optional[float] = Field(default=0.0, ge=-2.0, le=2.0)
|
||||||
|
logit_bias: Optional[Dict[int, float]] = None
|
||||||
|
user: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
|
class AnthropicMessage(BaseModel):
|
||||||
|
role: str
|
||||||
|
content: Union[str, List[Dict[str, Any]]]
|
||||||
|
|
||||||
|
|
||||||
|
class MessagesRequest(BaseModel):
|
||||||
|
"""Anthropic Messages API request body."""
|
||||||
|
|
||||||
|
model: str = "astrai"
|
||||||
|
max_tokens: int = Field(default=1024, ge=1)
|
||||||
|
messages: List[AnthropicMessage]
|
||||||
|
system: Optional[str] = None
|
||||||
|
temperature: Optional[float] = Field(default=1.0, ge=0.0, le=2.0)
|
||||||
|
top_p: Optional[float] = Field(default=1.0, ge=0.0, le=1.0)
|
||||||
|
top_k: Optional[int] = Field(default=50, ge=1)
|
||||||
|
stream: Optional[bool] = False
|
||||||
|
stop_sequences: Optional[List[str]] = None
|
||||||
|
|
||||||
|
|
||||||
|
def configure_server(
|
||||||
|
device: str = "cuda",
|
||||||
|
dtype: torch.dtype = torch.bfloat16,
|
||||||
|
param_path: Optional[Path] = None,
|
||||||
|
max_batch_size: int = 16,
|
||||||
|
):
|
||||||
|
_state.config.update(
|
||||||
|
device=device,
|
||||||
|
dtype=dtype,
|
||||||
|
param_path=param_path,
|
||||||
|
max_batch_size=max_batch_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
|
async def lifespan(app: FastAPI):
|
||||||
|
try:
|
||||||
|
load_model(
|
||||||
|
param_path=_state.config["param_path"],
|
||||||
|
device=_state.config["device"],
|
||||||
|
dtype=_state.config["dtype"],
|
||||||
|
max_batch_size=_state.config["max_batch_size"],
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Failed to load model: {e}")
|
||||||
|
raise
|
||||||
|
yield
|
||||||
|
if _state.engine:
|
||||||
|
_state.engine.shutdown()
|
||||||
|
logger.info("Inference engine shutdown complete")
|
||||||
|
|
||||||
|
|
||||||
|
app = FastAPI(title="AstrAI Inference Server", version="0.2.0", lifespan=lifespan)
|
||||||
|
|
||||||
|
|
||||||
|
def load_model(
|
||||||
|
param_path: Optional[Path] = None,
|
||||||
|
device: str = "cuda",
|
||||||
|
dtype: torch.dtype = torch.bfloat16,
|
||||||
|
max_batch_size: int = 16,
|
||||||
|
):
|
||||||
|
if param_path is None:
|
||||||
|
param_path = _project_root / "params"
|
||||||
|
if not param_path.exists():
|
||||||
|
raise FileNotFoundError(f"Parameter directory not found: {param_path}")
|
||||||
|
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||||
|
model = AutoModel.from_pretrained(param_path)
|
||||||
|
model.to(device=device, dtype=dtype)
|
||||||
|
logger.info(f"Model loaded on {device} with dtype {dtype}")
|
||||||
|
|
||||||
|
_state.engine = InferenceEngine(
|
||||||
|
model=model,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
max_batch_size=max_batch_size,
|
||||||
|
)
|
||||||
|
logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}")
|
||||||
|
|
||||||
|
|
||||||
|
def _get_engine() -> InferenceEngine:
|
||||||
|
if _state.engine is None:
|
||||||
|
raise HTTPException(status_code=503, detail="Engine not initialized")
|
||||||
|
return _state.engine
|
||||||
|
|
||||||
|
|
||||||
|
def _make_chunk(
|
||||||
|
delta: Dict[str, str],
|
||||||
|
finish_reason: Optional[str] = None,
|
||||||
|
*,
|
||||||
|
resp_id: str,
|
||||||
|
created: int,
|
||||||
|
model: str,
|
||||||
|
index: int = 0,
|
||||||
|
) -> str:
|
||||||
|
"""Build a single SSE ``data:`` chunk matching OpenAI streaming format."""
|
||||||
|
data = {
|
||||||
|
"id": resp_id,
|
||||||
|
"object": "chat.completion.chunk",
|
||||||
|
"created": created,
|
||||||
|
"model": model,
|
||||||
|
"choices": [
|
||||||
|
{
|
||||||
|
"index": index,
|
||||||
|
"delta": delta,
|
||||||
|
"finish_reason": finish_reason,
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
return f"data: {json.dumps(data, ensure_ascii=False)}\n\n"
|
||||||
|
|
||||||
|
|
||||||
|
@app.get("/health")
|
||||||
|
async def health():
|
||||||
|
return {
|
||||||
|
"status": "ok",
|
||||||
|
"model_loaded": _state.engine is not None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@app.get("/stats")
|
||||||
|
async def get_stats():
|
||||||
|
return _get_engine().get_stats()
|
||||||
|
|
||||||
|
|
||||||
|
@app.post("/v1/chat/completions")
|
||||||
|
async def chat_completion(request: ChatCompletionRequest):
|
||||||
|
"""OpenAI-compatible chat completion endpoint (streaming + non-streaming)."""
|
||||||
|
engine = _get_engine()
|
||||||
|
resp_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
||||||
|
created = int(time.time())
|
||||||
|
model = request.model
|
||||||
|
|
||||||
|
prompt = engine.tokenizer.apply_chat_template(
|
||||||
|
[{"role": m.role, "content": m.content} for m in request.messages],
|
||||||
|
tokenize=False,
|
||||||
|
)
|
||||||
|
prompt_tokens = len(engine.tokenizer.encode(prompt))
|
||||||
|
|
||||||
|
if request.stream:
|
||||||
|
agen = engine.generate_async(
|
||||||
|
prompt=prompt,
|
||||||
|
max_tokens=request.max_tokens,
|
||||||
|
temperature=request.temperature,
|
||||||
|
top_p=request.top_p,
|
||||||
|
top_k=request.top_k,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def event_stream():
|
||||||
|
yield _make_chunk(
|
||||||
|
{"role": "assistant"},
|
||||||
|
finish_reason=None,
|
||||||
|
resp_id=resp_id,
|
||||||
|
created=created,
|
||||||
|
model=model,
|
||||||
|
)
|
||||||
|
|
||||||
|
completion_tokens = 0
|
||||||
|
async for token in agen:
|
||||||
|
yield _make_chunk(
|
||||||
|
{"content": token},
|
||||||
|
finish_reason=None,
|
||||||
|
resp_id=resp_id,
|
||||||
|
created=created,
|
||||||
|
model=model,
|
||||||
|
)
|
||||||
|
completion_tokens += 1
|
||||||
|
|
||||||
|
yield _make_chunk(
|
||||||
|
{},
|
||||||
|
finish_reason="stop",
|
||||||
|
resp_id=resp_id,
|
||||||
|
created=created,
|
||||||
|
model=model,
|
||||||
|
)
|
||||||
|
|
||||||
|
usage = {
|
||||||
|
"prompt_tokens": prompt_tokens,
|
||||||
|
"completion_tokens": completion_tokens,
|
||||||
|
"total_tokens": prompt_tokens + completion_tokens,
|
||||||
|
}
|
||||||
|
yield f"data: {json.dumps(usage, ensure_ascii=False)}\n\n"
|
||||||
|
yield "data: [DONE]\n\n"
|
||||||
|
|
||||||
|
return StreamingResponse(
|
||||||
|
event_stream(),
|
||||||
|
media_type="text/event-stream",
|
||||||
|
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
||||||
|
)
|
||||||
|
|
||||||
|
completion_tokens = 0
|
||||||
|
chunks: List[str] = []
|
||||||
|
agen = engine.generate_async(
|
||||||
|
prompt=prompt,
|
||||||
|
max_tokens=request.max_tokens,
|
||||||
|
temperature=request.temperature,
|
||||||
|
top_p=request.top_p,
|
||||||
|
top_k=request.top_k,
|
||||||
|
)
|
||||||
|
async for token in agen:
|
||||||
|
chunks.append(token)
|
||||||
|
completion_tokens += 1
|
||||||
|
content = "".join(chunks)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"id": resp_id,
|
||||||
|
"object": "chat.completion",
|
||||||
|
"created": created,
|
||||||
|
"model": model,
|
||||||
|
"choices": [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"message": {"role": "assistant", "content": content},
|
||||||
|
"finish_reason": "stop",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"usage": {
|
||||||
|
"prompt_tokens": prompt_tokens,
|
||||||
|
"completion_tokens": completion_tokens,
|
||||||
|
"total_tokens": prompt_tokens + completion_tokens,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _make_anthropic_sse(event: str, data: Dict[str, Any]) -> str:
|
||||||
|
return f"event: {event}\ndata: {json.dumps(data, ensure_ascii=False)}\n\n"
|
||||||
|
|
||||||
|
|
||||||
|
def _check_stop_sequence(text: str, stop_sequences: List[str]) -> Optional[str]:
|
||||||
|
for seq in stop_sequences:
|
||||||
|
if seq and seq in text:
|
||||||
|
return seq
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_text_content(content: Union[str, List[Dict[str, Any]]]) -> str:
|
||||||
|
if isinstance(content, str):
|
||||||
|
return content
|
||||||
|
if isinstance(content, list):
|
||||||
|
for block in content:
|
||||||
|
if isinstance(block, dict) and block.get("type") == "text":
|
||||||
|
return block.get("text", "")
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
def _build_anthropic_messages(
|
||||||
|
messages: List[AnthropicMessage], system: Optional[str]
|
||||||
|
) -> List[Dict[str, str]]:
|
||||||
|
result: List[Dict[str, str]] = []
|
||||||
|
if system:
|
||||||
|
result.append({"role": "system", "content": system})
|
||||||
|
for m in messages:
|
||||||
|
content = _extract_text_content(m.content)
|
||||||
|
if content:
|
||||||
|
result.append({"role": m.role, "content": content})
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
@app.post("/v1/messages")
|
||||||
|
async def create_message(request: MessagesRequest):
|
||||||
|
"""Anthropic-compatible Messages API endpoint (streaming + non-streaming)."""
|
||||||
|
engine = _get_engine()
|
||||||
|
resp_id = f"msg_{uuid.uuid4().hex[:24]}"
|
||||||
|
model = request.model
|
||||||
|
|
||||||
|
chat_messages = _build_anthropic_messages(request.messages, request.system)
|
||||||
|
prompt = engine.tokenizer.apply_chat_template(chat_messages, tokenize=False)
|
||||||
|
prompt_tokens = len(engine.tokenizer.encode(prompt))
|
||||||
|
|
||||||
|
stop_sequences = request.stop_sequences or []
|
||||||
|
|
||||||
|
if request.stream:
|
||||||
|
agen = engine.generate_async(
|
||||||
|
prompt=prompt,
|
||||||
|
max_tokens=request.max_tokens,
|
||||||
|
temperature=request.temperature,
|
||||||
|
top_p=request.top_p,
|
||||||
|
top_k=request.top_k,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def event_stream():
|
||||||
|
yield _make_anthropic_sse(
|
||||||
|
"message_start",
|
||||||
|
{
|
||||||
|
"type": "message_start",
|
||||||
|
"message": {
|
||||||
|
"id": resp_id,
|
||||||
|
"type": "message",
|
||||||
|
"role": "assistant",
|
||||||
|
"model": model,
|
||||||
|
"content": [],
|
||||||
|
"usage": {"input_tokens": prompt_tokens},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
yield _make_anthropic_sse(
|
||||||
|
"content_block_start",
|
||||||
|
{
|
||||||
|
"type": "content_block_start",
|
||||||
|
"index": 0,
|
||||||
|
"content_block": {"type": "text", "text": ""},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
completion_tokens = 0
|
||||||
|
accumulated = ""
|
||||||
|
stopped_seq: Optional[str] = None
|
||||||
|
async for token in agen:
|
||||||
|
accumulated += token
|
||||||
|
completion_tokens += 1
|
||||||
|
|
||||||
|
matched = _check_stop_sequence(accumulated, stop_sequences)
|
||||||
|
if matched:
|
||||||
|
text = accumulated[: accumulated.rfind(matched)]
|
||||||
|
stopped_seq = matched
|
||||||
|
if text:
|
||||||
|
yield _make_anthropic_sse(
|
||||||
|
"content_block_delta",
|
||||||
|
{
|
||||||
|
"type": "content_block_delta",
|
||||||
|
"index": 0,
|
||||||
|
"delta": {"type": "text_delta", "text": text},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
break
|
||||||
|
|
||||||
|
yield _make_anthropic_sse(
|
||||||
|
"content_block_delta",
|
||||||
|
{
|
||||||
|
"type": "content_block_delta",
|
||||||
|
"index": 0,
|
||||||
|
"delta": {"type": "text_delta", "text": token},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
yield _make_anthropic_sse(
|
||||||
|
"content_block_stop",
|
||||||
|
{"type": "content_block_stop", "index": 0},
|
||||||
|
)
|
||||||
|
|
||||||
|
stop_reason = "stop_sequence" if stopped_seq else "end_turn"
|
||||||
|
yield _make_anthropic_sse(
|
||||||
|
"message_delta",
|
||||||
|
{
|
||||||
|
"type": "message_delta",
|
||||||
|
"delta": {"stop_reason": stop_reason, "stop_sequence": stopped_seq},
|
||||||
|
"usage": {"output_tokens": completion_tokens},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
yield _make_anthropic_sse(
|
||||||
|
"message_stop",
|
||||||
|
{"type": "message_stop"},
|
||||||
|
)
|
||||||
|
|
||||||
|
return StreamingResponse(
|
||||||
|
event_stream(),
|
||||||
|
media_type="text/event-stream",
|
||||||
|
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
||||||
|
)
|
||||||
|
|
||||||
|
completion_tokens = 0
|
||||||
|
chunks: List[str] = []
|
||||||
|
agen = engine.generate_async(
|
||||||
|
prompt=prompt,
|
||||||
|
max_tokens=request.max_tokens,
|
||||||
|
temperature=request.temperature,
|
||||||
|
top_p=request.top_p,
|
||||||
|
top_k=request.top_k,
|
||||||
|
)
|
||||||
|
stopped_seq: Optional[str] = None
|
||||||
|
accumulated = ""
|
||||||
|
async for token in agen:
|
||||||
|
chunks.append(token)
|
||||||
|
completion_tokens += 1
|
||||||
|
accumulated += token
|
||||||
|
matched = _check_stop_sequence(accumulated, stop_sequences)
|
||||||
|
if matched:
|
||||||
|
stopped_seq = matched
|
||||||
|
break
|
||||||
|
|
||||||
|
content = "".join(chunks)
|
||||||
|
if stopped_seq:
|
||||||
|
idx = content.rfind(stopped_seq)
|
||||||
|
if idx != -1:
|
||||||
|
content = content[:idx]
|
||||||
|
|
||||||
|
return {
|
||||||
|
"id": resp_id,
|
||||||
|
"type": "message",
|
||||||
|
"role": "assistant",
|
||||||
|
"model": model,
|
||||||
|
"content": [{"type": "text", "text": content}],
|
||||||
|
"stop_reason": "stop_sequence" if stopped_seq else "end_turn",
|
||||||
|
"stop_sequence": stopped_seq,
|
||||||
|
"usage": {
|
||||||
|
"input_tokens": prompt_tokens,
|
||||||
|
"output_tokens": completion_tokens,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def run_server(
|
||||||
|
host: str = "0.0.0.0",
|
||||||
|
port: int = 8000,
|
||||||
|
reload: bool = False,
|
||||||
|
device: str = "cuda",
|
||||||
|
dtype: torch.dtype = torch.bfloat16,
|
||||||
|
param_path: Optional[Path] = None,
|
||||||
|
max_batch_size: int = 16,
|
||||||
|
):
|
||||||
|
configure_server(
|
||||||
|
device=device,
|
||||||
|
dtype=dtype,
|
||||||
|
param_path=param_path,
|
||||||
|
max_batch_size=max_batch_size,
|
||||||
|
)
|
||||||
|
uvicorn.run(
|
||||||
|
"astrai.inference.server:app",
|
||||||
|
host=host,
|
||||||
|
port=port,
|
||||||
|
reload=reload,
|
||||||
|
)
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
from astrai.model.automodel import AutoModel
|
||||||
|
from astrai.model.module import (
|
||||||
|
GQA,
|
||||||
|
MLP,
|
||||||
|
DecoderBlock,
|
||||||
|
Linear,
|
||||||
|
RMSNorm,
|
||||||
|
)
|
||||||
|
from astrai.model.transformer import Transformer
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
# Modules
|
||||||
|
"Linear",
|
||||||
|
"RMSNorm",
|
||||||
|
"MLP",
|
||||||
|
"GQA",
|
||||||
|
"DecoderBlock",
|
||||||
|
# Models
|
||||||
|
"Transformer",
|
||||||
|
"AutoModel",
|
||||||
|
]
|
||||||
@@ -0,0 +1,129 @@
|
|||||||
|
"""
|
||||||
|
AutoModel base class for model loading and saving.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Self, Type, Union
|
||||||
|
|
||||||
|
import safetensors.torch as st
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
from astrai.config import ModelConfig
|
||||||
|
from astrai.factory import Registry
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def _disable_random_init(enable: bool = True):
|
||||||
|
init_functions = [
|
||||||
|
"xavier_normal_",
|
||||||
|
"xavier_uniform_",
|
||||||
|
"kaiming_normal_",
|
||||||
|
"kaiming_uniform_",
|
||||||
|
"zeros_",
|
||||||
|
"ones_",
|
||||||
|
"constant_",
|
||||||
|
"normal_",
|
||||||
|
"uniform_",
|
||||||
|
]
|
||||||
|
original_funcs = {}
|
||||||
|
for name in init_functions:
|
||||||
|
if enable and hasattr(nn.init, name):
|
||||||
|
original_funcs[name] = getattr(nn.init, name)
|
||||||
|
setattr(nn.init, name, lambda *args, **kwargs: None)
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
if enable:
|
||||||
|
for name, orig_func in original_funcs.items():
|
||||||
|
setattr(nn.init, name, orig_func)
|
||||||
|
|
||||||
|
|
||||||
|
class AutoModel(nn.Module):
|
||||||
|
"""
|
||||||
|
Autoregressive language model base class.
|
||||||
|
Provides model loading/saving and generation capabilities.
|
||||||
|
"""
|
||||||
|
|
||||||
|
_registry = Registry()
|
||||||
|
|
||||||
|
def __init__(self, config: ModelConfig):
|
||||||
|
super().__init__()
|
||||||
|
self.config = config
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def register(cls, model_type: str):
|
||||||
|
"""
|
||||||
|
Class method decorator to register model type.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
@AutoModel.register('transformer')
|
||||||
|
class Transformer(AutoModel):
|
||||||
|
...
|
||||||
|
"""
|
||||||
|
|
||||||
|
def decorator(sub_cls: Type["AutoModel"]) -> Type["AutoModel"]:
|
||||||
|
cls._registry.register(model_type.lower(), sub_cls)
|
||||||
|
return sub_cls
|
||||||
|
|
||||||
|
return decorator
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_model_class(cls, model_type: str) -> Type["AutoModel"]:
|
||||||
|
"""Get model class by model_type string."""
|
||||||
|
model_type = model_type.lower()
|
||||||
|
if not cls._registry.contains(model_type):
|
||||||
|
available = cls._registry.list_names()
|
||||||
|
raise ValueError(
|
||||||
|
f"Unknown model_type: {model_type}. Available: {available}"
|
||||||
|
)
|
||||||
|
return cls._registry.get(model_type)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_pretrained(
|
||||||
|
cls,
|
||||||
|
path: Union[str, Path],
|
||||||
|
disable_random_init: bool = True,
|
||||||
|
strict: bool = True,
|
||||||
|
) -> nn.Module:
|
||||||
|
|
||||||
|
model_path = Path(path)
|
||||||
|
|
||||||
|
# Load config
|
||||||
|
config = ModelConfig()
|
||||||
|
config_path = model_path / "config.json"
|
||||||
|
if config_path.exists():
|
||||||
|
config.load(str(config_path))
|
||||||
|
else:
|
||||||
|
raise FileNotFoundError(f"Config file not found: {config_path}")
|
||||||
|
|
||||||
|
model_type = config.model_type or "transformer"
|
||||||
|
actual_cls = cls.get_model_class(model_type)
|
||||||
|
|
||||||
|
with _disable_random_init(enable=disable_random_init):
|
||||||
|
model = actual_cls(config)
|
||||||
|
|
||||||
|
# Load weights
|
||||||
|
weights_path = model_path / "model.safetensors"
|
||||||
|
if weights_path.exists():
|
||||||
|
state_dict = st.load_file(str(weights_path))
|
||||||
|
model.load_state_dict(state_dict, strict=strict)
|
||||||
|
|
||||||
|
return model
|
||||||
|
|
||||||
|
def save_pretrained(
|
||||||
|
self,
|
||||||
|
save_directory: Union[str, Path],
|
||||||
|
) -> None:
|
||||||
|
save_path = Path(save_directory)
|
||||||
|
save_path.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
# Save config
|
||||||
|
self.config.save(str(save_path / "config.json"))
|
||||||
|
|
||||||
|
# Save weights
|
||||||
|
st.save_file(self.state_dict(), str(save_path / "model.safetensors"))
|
||||||
|
|
||||||
|
def to(self, *args, **kwargs) -> Self:
|
||||||
|
"""Move model to device/dtype."""
|
||||||
|
return super().to(*args, **kwargs)
|
||||||
@@ -0,0 +1,337 @@
|
|||||||
|
from typing import Optional, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.inference.cache import CacheView
|
||||||
|
|
||||||
|
|
||||||
|
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
||||||
|
"""Repeat KV heads n_rep times for GQA."""
|
||||||
|
bs, slen, n_heads, head_dim = x.shape
|
||||||
|
if n_rep == 1:
|
||||||
|
return x
|
||||||
|
return (
|
||||||
|
x[:, :, :, None, :]
|
||||||
|
.expand(bs, slen, n_heads, n_rep, head_dim)
|
||||||
|
.reshape(bs, slen, n_heads * n_rep, head_dim)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def get_rotary_emb(
|
||||||
|
dim: int,
|
||||||
|
max_len: int,
|
||||||
|
base: float = 10000,
|
||||||
|
device: Optional[torch.device] = None,
|
||||||
|
) -> Tuple[Tensor, Tensor]:
|
||||||
|
"""Precompute cos/sin for RoPE."""
|
||||||
|
theta = base ** (-torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim)
|
||||||
|
t = torch.arange(0, max_len, dtype=torch.float64, device=device)
|
||||||
|
freqs = torch.outer(t, theta)
|
||||||
|
return torch.cos(freqs).float(), torch.sin(freqs).float()
|
||||||
|
|
||||||
|
|
||||||
|
def apply_rotary_emb(x: torch.Tensor, rotary_emb: Tuple[Tensor, Tensor]) -> Tensor:
|
||||||
|
"""Apply rotary embedding via cos/sin (shape-preserving)."""
|
||||||
|
dtype = x.dtype
|
||||||
|
cos, sin = rotary_emb
|
||||||
|
cos = cos.unsqueeze(0).unsqueeze(2)
|
||||||
|
sin = sin.unsqueeze(0).unsqueeze(2)
|
||||||
|
x_real = x[..., 0::2]
|
||||||
|
x_imag = x[..., 1::2]
|
||||||
|
x_real_rot = x_real * cos - x_imag * sin
|
||||||
|
x_imag_rot = x_real * sin + x_imag * cos
|
||||||
|
x_out = torch.stack([x_real_rot, x_imag_rot], dim=-1)
|
||||||
|
x_out = x_out.view(*x_out.shape[:-2], -1)
|
||||||
|
return x_out.to(dtype)
|
||||||
|
|
||||||
|
|
||||||
|
class RotaryEmbedding(nn.Module):
|
||||||
|
def __init__(self, dim: int, max_len: int, base: int = 10000):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
self.max_len = max_len
|
||||||
|
self.base = base
|
||||||
|
self.max_len_cached = None
|
||||||
|
self._set_rotary_buffer(self.max_len, None)
|
||||||
|
|
||||||
|
def _set_rotary_buffer(self, max_len: int, device: Optional[torch.device] = None):
|
||||||
|
cos_cached, sin_cached = get_rotary_emb(self.dim, max_len, self.base, device)
|
||||||
|
self.register_buffer("cos_cached", cos_cached, persistent=False)
|
||||||
|
self.register_buffer("sin_cached", sin_cached, persistent=False)
|
||||||
|
self.max_len_cached = max_len
|
||||||
|
|
||||||
|
def forward(self, x: Tensor, start_pos: int = 0) -> Tuple[Tensor, Tensor]:
|
||||||
|
seq_len = x.size(1)
|
||||||
|
if self.max_len_cached < seq_len + start_pos:
|
||||||
|
self._set_rotary_buffer(self.max_len_cached * 2, x.device)
|
||||||
|
cos = self.cos_cached[start_pos : start_pos + seq_len]
|
||||||
|
sin = self.sin_cached[start_pos : start_pos + seq_len]
|
||||||
|
return (cos, sin)
|
||||||
|
|
||||||
|
|
||||||
|
class Linear(nn.Module):
|
||||||
|
def __init__(self, in_dim: int, out_dim: int, bias: bool = False):
|
||||||
|
super().__init__()
|
||||||
|
self.weight = nn.Parameter(torch.empty((out_dim, in_dim)))
|
||||||
|
self.bias = nn.Parameter(torch.zeros(out_dim)) if bias else None
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
return F.linear(x, self.weight, self.bias)
|
||||||
|
|
||||||
|
|
||||||
|
class RMSNorm(nn.Module):
|
||||||
|
def __init__(self, dim, norm_eps):
|
||||||
|
super().__init__()
|
||||||
|
self.weight = nn.Parameter(torch.ones(dim))
|
||||||
|
self.normalized_shape = (dim,)
|
||||||
|
self.norm_eps = norm_eps
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
return F.rms_norm(x, self.normalized_shape, self.weight, self.norm_eps)
|
||||||
|
|
||||||
|
|
||||||
|
class MLP(nn.Module):
|
||||||
|
def __init__(self, dim: int, dim_feed_forward: int):
|
||||||
|
super().__init__()
|
||||||
|
self.up = Linear(dim, dim_feed_forward)
|
||||||
|
self.gate = Linear(dim, dim_feed_forward)
|
||||||
|
self.down = Linear(dim_feed_forward, dim)
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
gated = self.up(x) * F.silu(self.gate(x))
|
||||||
|
out = self.down(gated)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
class GQA(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim: int,
|
||||||
|
n_heads: int,
|
||||||
|
n_kv_heads: int,
|
||||||
|
use_qk_norm: bool,
|
||||||
|
norm_eps: float,
|
||||||
|
use_gated_attention: bool,
|
||||||
|
layer_id: int,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
assert dim % n_heads == 0
|
||||||
|
assert n_heads % n_kv_heads == 0
|
||||||
|
|
||||||
|
self.head_dim = dim // n_heads
|
||||||
|
self.layer_id = layer_id
|
||||||
|
self.dim = dim
|
||||||
|
self.n_heads = n_heads
|
||||||
|
self.n_kv_heads = n_kv_heads
|
||||||
|
self.n_rep = n_heads // n_kv_heads
|
||||||
|
self.use_qk_norm = use_qk_norm
|
||||||
|
self.use_gated_attention = use_gated_attention
|
||||||
|
|
||||||
|
self.q_proj = Linear(dim, n_heads * self.head_dim)
|
||||||
|
self.k_proj = Linear(dim, n_kv_heads * self.head_dim)
|
||||||
|
self.v_proj = Linear(dim, n_kv_heads * self.head_dim)
|
||||||
|
self.o_proj = Linear(dim, dim)
|
||||||
|
|
||||||
|
if self.use_qk_norm:
|
||||||
|
self.q_norm = RMSNorm(self.head_dim, norm_eps)
|
||||||
|
self.k_norm = RMSNorm(self.head_dim, norm_eps)
|
||||||
|
|
||||||
|
if self.use_gated_attention:
|
||||||
|
self.gate = Linear(dim, dim)
|
||||||
|
|
||||||
|
def _split_heads(self, x: Tensor, n_heads) -> Tensor:
|
||||||
|
batch_size, seq_len, _ = x.shape
|
||||||
|
x = x.reshape(batch_size, seq_len, n_heads, self.head_dim)
|
||||||
|
return x
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
x: Tensor,
|
||||||
|
rotary_emb: Tuple[Tensor, Tensor],
|
||||||
|
mask: Tensor = None,
|
||||||
|
paged_cache: Optional[CacheView] = None,
|
||||||
|
start_pos: int = 0,
|
||||||
|
) -> Tensor:
|
||||||
|
bsz, seq_len, _ = x.size()
|
||||||
|
is_causal = mask is None
|
||||||
|
|
||||||
|
# (bsz, seq_len, dim) -> (bsz, seq_len, n_heads, head_dim)
|
||||||
|
q = self._split_heads(self.q_proj(x), self.n_heads)
|
||||||
|
k = self._split_heads(self.k_proj(x), self.n_kv_heads)
|
||||||
|
v = self._split_heads(self.v_proj(x), self.n_kv_heads)
|
||||||
|
q, k = apply_rotary_emb(q, rotary_emb), apply_rotary_emb(k, rotary_emb)
|
||||||
|
|
||||||
|
if self.use_qk_norm:
|
||||||
|
q, k = self.q_norm(q), self.k_norm(k)
|
||||||
|
|
||||||
|
if paged_cache is not None:
|
||||||
|
paged_cache.write(self.layer_id, start_pos, k, v)
|
||||||
|
k, v = paged_cache.gather(self.layer_id)
|
||||||
|
|
||||||
|
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)
|
||||||
|
|
||||||
|
# (bsz, seq_len, n_heads, head_dim) -> (bsz, n_heads, seq_len, head_dim)
|
||||||
|
q, k, v = q.permute(0, 2, 1, 3), k.permute(0, 2, 1, 3), v.permute(0, 2, 1, 3)
|
||||||
|
sdqa_out = (
|
||||||
|
F.scaled_dot_product_attention(q, k, v, mask, is_causal=is_causal)
|
||||||
|
.permute(0, 2, 1, 3)
|
||||||
|
.contiguous()
|
||||||
|
.flatten(2)
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.use_gated_attention:
|
||||||
|
sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
|
||||||
|
|
||||||
|
out = self.o_proj(sdqa_out)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
class MLA(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim: int,
|
||||||
|
n_heads: int,
|
||||||
|
n_kv_heads: int,
|
||||||
|
kv_lora_rank: int,
|
||||||
|
qk_nope_head_dim: int,
|
||||||
|
qk_rope_head_dim: int,
|
||||||
|
norm_eps: float,
|
||||||
|
use_gated_attention: bool,
|
||||||
|
layer_id: int,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
self.n_heads = n_heads
|
||||||
|
self.n_kv_heads = n_kv_heads
|
||||||
|
self.kv_lora_rank = kv_lora_rank
|
||||||
|
self.qk_nope_head_dim = qk_nope_head_dim
|
||||||
|
self.qk_rope_head_dim = qk_rope_head_dim
|
||||||
|
self.head_dim = qk_nope_head_dim + qk_rope_head_dim
|
||||||
|
self.layer_id = layer_id
|
||||||
|
self.n_rep = n_heads // n_kv_heads
|
||||||
|
self.use_gated_attention = use_gated_attention
|
||||||
|
|
||||||
|
self.q_proj = Linear(dim, n_heads * self.head_dim, bias=False)
|
||||||
|
self.kv_a_proj = Linear(dim, kv_lora_rank, bias=False)
|
||||||
|
self.kv_norm = RMSNorm(kv_lora_rank, norm_eps)
|
||||||
|
|
||||||
|
# fused KV: (k_nope, k_rope, v)
|
||||||
|
self.kv_b_proj = Linear(
|
||||||
|
kv_lora_rank,
|
||||||
|
n_kv_heads * (self.head_dim + qk_rope_head_dim + self.head_dim),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.o_proj = Linear(dim, dim, bias=False)
|
||||||
|
|
||||||
|
if use_gated_attention:
|
||||||
|
self.gate = Linear(dim, dim, bias=False)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
x: Tensor,
|
||||||
|
rotary_emb: Tuple[Tensor, Tensor],
|
||||||
|
mask: Tensor = None,
|
||||||
|
paged_cache: Optional[CacheView] = None,
|
||||||
|
start_pos: int = 0,
|
||||||
|
) -> Tensor:
|
||||||
|
bsz, seq_len, _ = x.size()
|
||||||
|
is_causal = mask is None
|
||||||
|
|
||||||
|
q = self.q_proj(x)
|
||||||
|
q = q.view(bsz, seq_len, self.n_heads, self.head_dim)
|
||||||
|
|
||||||
|
kv_compressed = self.kv_a_proj(x)
|
||||||
|
kv_compressed = self.kv_norm(kv_compressed)
|
||||||
|
|
||||||
|
kv = self.kv_b_proj(kv_compressed)
|
||||||
|
kv = kv.view(bsz, seq_len, self.n_kv_heads, -1)
|
||||||
|
|
||||||
|
k_nope, k_rope, v = torch.split(
|
||||||
|
kv, [self.qk_nope_head_dim, self.qk_rope_head_dim, self.head_dim], dim=-1
|
||||||
|
)
|
||||||
|
|
||||||
|
q_nope, q_rope = (
|
||||||
|
q[..., : self.qk_nope_head_dim],
|
||||||
|
q[..., self.qk_rope_head_dim :],
|
||||||
|
)
|
||||||
|
q_rope = apply_rotary_emb(q_rope, rotary_emb)
|
||||||
|
k_rope = apply_rotary_emb(k_rope, rotary_emb)
|
||||||
|
|
||||||
|
q = torch.cat([q_nope, q_rope], dim=-1)
|
||||||
|
k = torch.cat([k_nope, k_rope], dim=-1)
|
||||||
|
|
||||||
|
if paged_cache is not None:
|
||||||
|
paged_cache.write(self.layer_id, start_pos, k, v)
|
||||||
|
k, v = paged_cache.gather(self.layer_id)
|
||||||
|
|
||||||
|
q = q.permute(0, 2, 1, 3)
|
||||||
|
k = k.permute(0, 2, 1, 3)
|
||||||
|
v = v.permute(0, 2, 1, 3)
|
||||||
|
|
||||||
|
attn_out = F.scaled_dot_product_attention(q, k, v, mask, is_causal=is_causal)
|
||||||
|
attn_out = attn_out.permute(0, 2, 1, 3).contiguous().flatten(2)
|
||||||
|
|
||||||
|
if self.use_gated_attention:
|
||||||
|
attn_out = attn_out * F.sigmoid(self.gate(x))
|
||||||
|
|
||||||
|
out = self.o_proj(attn_out)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
class DecoderBlock(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim: int,
|
||||||
|
n_heads: int,
|
||||||
|
dim_ffn: int,
|
||||||
|
n_kv_heads: int,
|
||||||
|
norm_eps: int,
|
||||||
|
use_qk_norm: bool,
|
||||||
|
use_gated_attention: bool,
|
||||||
|
layer_id: int,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.attention = GQA(
|
||||||
|
dim,
|
||||||
|
n_heads,
|
||||||
|
n_kv_heads,
|
||||||
|
use_qk_norm,
|
||||||
|
norm_eps,
|
||||||
|
use_gated_attention,
|
||||||
|
layer_id,
|
||||||
|
)
|
||||||
|
self.input_norm = RMSNorm(dim, norm_eps)
|
||||||
|
self.mlp = MLP(dim, dim_ffn)
|
||||||
|
self.post_attention_norm = RMSNorm(dim, norm_eps)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
x: Tensor,
|
||||||
|
rotary_emb: Tuple[Tensor, Tensor],
|
||||||
|
attention_mask: Optional[Tensor] = None,
|
||||||
|
paged_cache: Optional[CacheView] = None,
|
||||||
|
start_pos: int = 0,
|
||||||
|
) -> Tensor:
|
||||||
|
attn_output = self.attention(
|
||||||
|
self.input_norm(x),
|
||||||
|
rotary_emb,
|
||||||
|
attention_mask,
|
||||||
|
paged_cache,
|
||||||
|
start_pos,
|
||||||
|
)
|
||||||
|
x = attn_output + x
|
||||||
|
|
||||||
|
x = self.mlp(self.post_attention_norm(x)) + x
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
class Embedding(nn.Module):
|
||||||
|
def __init__(self, vocab_size: int, embedding_dim: int):
|
||||||
|
super().__init__()
|
||||||
|
self.weight = nn.Parameter(torch.empty((vocab_size, embedding_dim)))
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
return F.embedding(x, self.weight)
|
||||||
@@ -0,0 +1,146 @@
|
|||||||
|
from typing import Any, Mapping, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.config.model_config import ModelConfig
|
||||||
|
from astrai.inference.cache import CacheView
|
||||||
|
from astrai.model.automodel import AutoModel
|
||||||
|
from astrai.model.module import (
|
||||||
|
DecoderBlock,
|
||||||
|
Embedding,
|
||||||
|
Linear,
|
||||||
|
RMSNorm,
|
||||||
|
RotaryEmbedding,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def process_attention_mask(
|
||||||
|
seq_mask: Tensor,
|
||||||
|
input_tensor: Tensor,
|
||||||
|
start_pos: int = 0,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Build 4D attention mask from 2D seq_mask, with optional causal masking."""
|
||||||
|
device = input_tensor.device
|
||||||
|
dtype = input_tensor.dtype
|
||||||
|
seq_len = input_tensor.size(1)
|
||||||
|
|
||||||
|
if seq_mask is None:
|
||||||
|
if start_pos != 0:
|
||||||
|
seq_mask = torch.ones((1, seq_len), dtype=torch.bool, device=device)
|
||||||
|
else:
|
||||||
|
return None
|
||||||
|
|
||||||
|
if seq_mask.dim() > 2:
|
||||||
|
return seq_mask
|
||||||
|
|
||||||
|
batch_size = seq_mask.size(0)
|
||||||
|
seq_mask = seq_mask[:, : start_pos + seq_len].to(device=device, dtype=torch.bool)
|
||||||
|
expanded_mask = seq_mask.unsqueeze(1).expand(
|
||||||
|
batch_size, seq_len, start_pos + seq_len
|
||||||
|
)
|
||||||
|
|
||||||
|
if is_causal:
|
||||||
|
expanded_mask = torch.tril(expanded_mask, diagonal=start_pos)
|
||||||
|
|
||||||
|
attention_mask = torch.zeros_like(expanded_mask, dtype=dtype, device=device)
|
||||||
|
attention_mask = attention_mask.masked_fill_(
|
||||||
|
~expanded_mask, -torch.finfo(dtype).max / 2
|
||||||
|
).unsqueeze(1)
|
||||||
|
|
||||||
|
return attention_mask
|
||||||
|
|
||||||
|
|
||||||
|
@AutoModel.register("transformer")
|
||||||
|
class Transformer(AutoModel):
|
||||||
|
"""Transformer language model with paged KV cache."""
|
||||||
|
|
||||||
|
def __init__(self, config: ModelConfig):
|
||||||
|
super().__init__(config)
|
||||||
|
self.config = config
|
||||||
|
self.rotary_embedding = RotaryEmbedding(
|
||||||
|
config.dim // config.n_heads, config.max_len
|
||||||
|
)
|
||||||
|
self.embed_tokens = Embedding(config.vocab_size, config.dim)
|
||||||
|
|
||||||
|
self.layers = nn.ModuleList(
|
||||||
|
[
|
||||||
|
DecoderBlock(
|
||||||
|
config.dim,
|
||||||
|
config.n_heads,
|
||||||
|
config.dim_ffn,
|
||||||
|
config.n_kv_heads,
|
||||||
|
config.norm_eps,
|
||||||
|
config.use_qk_norm,
|
||||||
|
config.use_gated_attention,
|
||||||
|
layer_id,
|
||||||
|
)
|
||||||
|
for layer_id in range(config.n_layers)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
self.norm = RMSNorm(config.dim, config.norm_eps)
|
||||||
|
self.lm_head = Linear(config.dim, config.vocab_size)
|
||||||
|
|
||||||
|
if self.config.tie_weight:
|
||||||
|
self.lm_head.weight = self.embed_tokens.weight
|
||||||
|
|
||||||
|
self._init_weights()
|
||||||
|
|
||||||
|
def _init_weights(self):
|
||||||
|
for param in self.parameters():
|
||||||
|
if param.dim() > 1:
|
||||||
|
nn.init.normal_(param, mean=0.0, std=0.006)
|
||||||
|
|
||||||
|
def load_state_dict(self, state_dict: Mapping[str, Any], strict=True, assign=False):
|
||||||
|
lm_head_key = "lm_head.weight"
|
||||||
|
embed_key = "embed_tokens.weight"
|
||||||
|
|
||||||
|
state_dict = dict(state_dict)
|
||||||
|
|
||||||
|
if self.config.tie_weight:
|
||||||
|
# same tensor for embed and lm_head
|
||||||
|
if embed_key in state_dict:
|
||||||
|
state_dict[lm_head_key] = state_dict[embed_key]
|
||||||
|
else:
|
||||||
|
if lm_head_key not in state_dict and embed_key in state_dict:
|
||||||
|
# clone to avoid sharing gradients
|
||||||
|
state_dict[lm_head_key] = torch.clone(state_dict[embed_key])
|
||||||
|
|
||||||
|
return super().load_state_dict(state_dict, strict, assign)
|
||||||
|
|
||||||
|
def state_dict(self, destination=None, prefix="", keep_vars=False):
|
||||||
|
state_dict = super().state_dict(
|
||||||
|
destination=destination, prefix=prefix, keep_vars=keep_vars
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.config.tie_weight:
|
||||||
|
lm_head_key = prefix + "lm_head.weight"
|
||||||
|
if lm_head_key in state_dict:
|
||||||
|
del state_dict[lm_head_key]
|
||||||
|
|
||||||
|
return state_dict
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
input_ids: Tensor,
|
||||||
|
input_mask: Optional[Tensor] = None,
|
||||||
|
paged_cache: Optional[CacheView] = None,
|
||||||
|
start_pos: int = 0,
|
||||||
|
) -> Tensor:
|
||||||
|
assert input_ids.ndim == 2
|
||||||
|
|
||||||
|
x = self.embed_tokens(input_ids)
|
||||||
|
rotary_emb = self.rotary_embedding(x, start_pos)
|
||||||
|
|
||||||
|
attn_mask = process_attention_mask(input_mask, x, start_pos, is_causal=True)
|
||||||
|
|
||||||
|
for layer in self.layers:
|
||||||
|
x = layer(x, rotary_emb, attn_mask, paged_cache, start_pos)
|
||||||
|
|
||||||
|
hidden_states = self.norm(x)
|
||||||
|
logits = self.lm_head(hidden_states)
|
||||||
|
|
||||||
|
return {"logits": logits, "hidden_states": hidden_states}
|
||||||
@@ -0,0 +1,20 @@
|
|||||||
|
from astrai.parallel.module import ColumnParallelLinear, RowParallelLinear
|
||||||
|
from astrai.parallel.setup import (
|
||||||
|
get_current_device,
|
||||||
|
get_rank,
|
||||||
|
get_world_size,
|
||||||
|
only_on_rank,
|
||||||
|
setup_parallel,
|
||||||
|
spawn_parallel_fn,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"get_world_size",
|
||||||
|
"get_rank",
|
||||||
|
"get_current_device",
|
||||||
|
"only_on_rank",
|
||||||
|
"setup_parallel",
|
||||||
|
"spawn_parallel_fn",
|
||||||
|
"RowParallelLinear",
|
||||||
|
"ColumnParallelLinear",
|
||||||
|
]
|
||||||
@@ -0,0 +1,115 @@
|
|||||||
|
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)
|
||||||
@@ -0,0 +1,161 @@
|
|||||||
|
import os
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from functools import wraps
|
||||||
|
from typing import Callable
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
|
import torch.multiprocessing as mp
|
||||||
|
|
||||||
|
|
||||||
|
def get_current_device():
|
||||||
|
return os.environ["LOCAL_DEVICE"]
|
||||||
|
|
||||||
|
|
||||||
|
def get_world_size() -> int:
|
||||||
|
if dist.is_available() and dist.is_initialized():
|
||||||
|
return dist.get_world_size()
|
||||||
|
else:
|
||||||
|
return 1
|
||||||
|
|
||||||
|
|
||||||
|
def get_rank() -> int:
|
||||||
|
if dist.is_available() and dist.is_initialized():
|
||||||
|
return dist.get_rank()
|
||||||
|
else:
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def setup_parallel(
|
||||||
|
rank: int,
|
||||||
|
world_size: int,
|
||||||
|
backend: str = "nccl",
|
||||||
|
master_addr: str = "localhost",
|
||||||
|
master_port: str = "29500",
|
||||||
|
device_type: str = "cuda",
|
||||||
|
):
|
||||||
|
|
||||||
|
if dist.is_available() and dist.is_initialized():
|
||||||
|
yield dist.group.WORLD
|
||||||
|
return
|
||||||
|
|
||||||
|
if world_size <= 1:
|
||||||
|
yield None
|
||||||
|
return
|
||||||
|
|
||||||
|
device_id = torch.device(device_type, rank)
|
||||||
|
|
||||||
|
os.environ["MASTER_ADDR"] = master_addr
|
||||||
|
os.environ["MASTER_PORT"] = master_port
|
||||||
|
os.environ["LOCAL_RANK"] = str(rank)
|
||||||
|
os.environ["WORLD_SIZE"] = str(world_size)
|
||||||
|
os.environ["LOCAL_DEVICE"] = str(device_id)
|
||||||
|
|
||||||
|
dist.init_process_group(
|
||||||
|
rank=rank, world_size=world_size, backend=backend, device_id=device_id
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
if backend == "nccl" and torch.cuda.is_available():
|
||||||
|
torch.cuda.set_device(device_id)
|
||||||
|
elif backend == "ccl" and hasattr(torch, "xpu") and torch.xpu.is_available():
|
||||||
|
torch.xpu.set_device(device_id)
|
||||||
|
|
||||||
|
yield dist.group.WORLD
|
||||||
|
finally:
|
||||||
|
if dist.is_initialized():
|
||||||
|
dist.destroy_process_group()
|
||||||
|
|
||||||
|
|
||||||
|
def only_on_rank(rank, sync=False):
|
||||||
|
"""
|
||||||
|
decorator to run a function only on a specific rank.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def decorator(func):
|
||||||
|
@wraps(func)
|
||||||
|
def wrapper(*args, **kwargs):
|
||||||
|
ret_args = None
|
||||||
|
if get_rank() == rank:
|
||||||
|
ret_args = func(*args, **kwargs)
|
||||||
|
|
||||||
|
if sync and dist.is_available() and dist.is_initialized():
|
||||||
|
dist.barrier()
|
||||||
|
|
||||||
|
return ret_args
|
||||||
|
|
||||||
|
return wrapper
|
||||||
|
|
||||||
|
return decorator
|
||||||
|
|
||||||
|
|
||||||
|
def wrapper_spawn_func(
|
||||||
|
rank: int,
|
||||||
|
world_size: int,
|
||||||
|
backend: str,
|
||||||
|
master_addr: str,
|
||||||
|
master_port: str,
|
||||||
|
device_type: str,
|
||||||
|
func: Callable,
|
||||||
|
kwargs: dict,
|
||||||
|
):
|
||||||
|
try:
|
||||||
|
with setup_parallel(
|
||||||
|
rank=rank,
|
||||||
|
world_size=world_size,
|
||||||
|
backend=backend,
|
||||||
|
master_addr=master_addr,
|
||||||
|
master_port=master_port,
|
||||||
|
device_type=device_type,
|
||||||
|
):
|
||||||
|
func(**kwargs)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Error in rank {rank}: {e}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
|
def spawn_parallel_fn(
|
||||||
|
func: Callable,
|
||||||
|
world_size: int,
|
||||||
|
backend: str = "nccl",
|
||||||
|
master_addr: str = "localhost",
|
||||||
|
master_port: str = "29500",
|
||||||
|
device_type: str = "cuda",
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
# clear environment variables
|
||||||
|
for key in [
|
||||||
|
"MASTER_ADDR",
|
||||||
|
"MASTER_PORT",
|
||||||
|
"RANK",
|
||||||
|
"WORLD_SIZE",
|
||||||
|
"LOCAL_RANK",
|
||||||
|
"LOCAL_DEVICE",
|
||||||
|
]:
|
||||||
|
if key in os.environ:
|
||||||
|
del os.environ[key]
|
||||||
|
|
||||||
|
if world_size == 1:
|
||||||
|
device_id = torch.device(device_type, 0)
|
||||||
|
os.environ["LOCAL_RANK"] = "0"
|
||||||
|
os.environ["WORLD_SIZE"] = "1"
|
||||||
|
os.environ["LOCAL_DEVICE"] = str(device_id)
|
||||||
|
|
||||||
|
func(**kwargs)
|
||||||
|
return
|
||||||
|
|
||||||
|
wrapper_spawn_func_args = (
|
||||||
|
world_size,
|
||||||
|
backend,
|
||||||
|
master_addr,
|
||||||
|
master_port,
|
||||||
|
device_type,
|
||||||
|
func,
|
||||||
|
kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
mp.spawn(
|
||||||
|
wrapper_spawn_func, nprocs=world_size, args=wrapper_spawn_func_args, join=True
|
||||||
|
)
|
||||||
@@ -0,0 +1,116 @@
|
|||||||
|
import json
|
||||||
|
import os
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
import h5py
|
||||||
|
import safetensors.torch as st
|
||||||
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.parallel.setup import get_rank
|
||||||
|
|
||||||
|
|
||||||
|
def save_h5(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
|
||||||
|
os.makedirs(file_path, exist_ok=True)
|
||||||
|
full_file_path = os.path.join(file_path, f"{file_name}.h5")
|
||||||
|
with h5py.File(full_file_path, "w") as f:
|
||||||
|
for key, tensors in tensor_group.items():
|
||||||
|
grp = f.create_group(key)
|
||||||
|
for idx, tensor in enumerate(tensors):
|
||||||
|
arr = tensor.cpu().numpy()
|
||||||
|
grp.create_dataset(f"data_{idx}", data=arr)
|
||||||
|
|
||||||
|
|
||||||
|
def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
|
||||||
|
tensor_group: Dict[str, List[Tensor]] = {}
|
||||||
|
|
||||||
|
root_path = Path(file_path)
|
||||||
|
h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5"))
|
||||||
|
|
||||||
|
for h5_file in h5_files:
|
||||||
|
with h5py.File(h5_file, "r") as f:
|
||||||
|
for key in f.keys():
|
||||||
|
grp = f[key]
|
||||||
|
dsets = []
|
||||||
|
for dset_name in grp.keys():
|
||||||
|
dset = grp[dset_name]
|
||||||
|
tensor = torch.from_numpy(dset[:])
|
||||||
|
if share_memory:
|
||||||
|
tensor = tensor.share_memory_()
|
||||||
|
dsets.append(tensor)
|
||||||
|
|
||||||
|
if tensor_group.get(key) is None:
|
||||||
|
tensor_group[key] = []
|
||||||
|
tensor_group[key].extend(dsets)
|
||||||
|
|
||||||
|
return tensor_group
|
||||||
|
|
||||||
|
|
||||||
|
class Checkpoint:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
state_dict: Dict[str, Any],
|
||||||
|
epoch: int = 0,
|
||||||
|
iteration: int = 0,
|
||||||
|
extra: Optional[Dict[str, Any]] = None,
|
||||||
|
):
|
||||||
|
self.state_dict = state_dict
|
||||||
|
self.epoch = epoch
|
||||||
|
self.iteration = iteration
|
||||||
|
self.extra = extra or {}
|
||||||
|
|
||||||
|
def save(
|
||||||
|
self,
|
||||||
|
save_dir: str,
|
||||||
|
) -> None:
|
||||||
|
|
||||||
|
save_path = Path(save_dir)
|
||||||
|
save_path.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
rank = get_rank()
|
||||||
|
if rank == 0:
|
||||||
|
meta = {
|
||||||
|
"epoch": self.epoch,
|
||||||
|
"iteration": self.iteration,
|
||||||
|
}
|
||||||
|
with open(save_path / "meta.json", "w") as f:
|
||||||
|
json.dump(meta, f, indent=2)
|
||||||
|
|
||||||
|
st.save_file(self.state_dict, save_path / "state_dict.safetensors")
|
||||||
|
if self.extra:
|
||||||
|
torch.save(self.extra, save_path / "extra.pt")
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def load(
|
||||||
|
cls,
|
||||||
|
save_dir: str,
|
||||||
|
) -> "Checkpoint":
|
||||||
|
|
||||||
|
rank = get_rank()
|
||||||
|
save_path = Path(save_dir)
|
||||||
|
|
||||||
|
meta = {}
|
||||||
|
if rank == 0:
|
||||||
|
with open(Path(save_dir) / "meta.json", "r") as f:
|
||||||
|
meta = json.load(f)
|
||||||
|
|
||||||
|
if dist.is_initialized():
|
||||||
|
meta_list = [meta]
|
||||||
|
dist.broadcast_object_list(meta_list, src=0)
|
||||||
|
meta = meta_list[0]
|
||||||
|
|
||||||
|
state_dict = st.load_file(save_path / "state_dict.safetensors")
|
||||||
|
|
||||||
|
extra = None
|
||||||
|
extra_path = save_path / "extra.pt"
|
||||||
|
if extra_path.exists():
|
||||||
|
extra = torch.load(extra_path, map_location="cpu", weights_only=False)
|
||||||
|
|
||||||
|
return cls(
|
||||||
|
state_dict=state_dict,
|
||||||
|
epoch=meta["epoch"],
|
||||||
|
iteration=meta["iteration"],
|
||||||
|
extra=extra,
|
||||||
|
)
|
||||||
@@ -0,0 +1,8 @@
|
|||||||
|
from astrai.tokenize.chat_template import ChatTemplate, MessageType
|
||||||
|
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"AutoTokenizer",
|
||||||
|
"ChatTemplate",
|
||||||
|
"MessageType",
|
||||||
|
]
|
||||||
@@ -0,0 +1,77 @@
|
|||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
from jinja2 import Template
|
||||||
|
|
||||||
|
# Message type for chat messages
|
||||||
|
type MessageType = Dict[str, Any]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ChatTemplate:
|
||||||
|
"""A chat template with Jinja2 rendering support.
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
name: Unique identifier for the template.
|
||||||
|
template_str: Jinja2 template string.
|
||||||
|
description: Optional description.
|
||||||
|
default_variables: Optional dictionary of default variable values
|
||||||
|
that will be passed to the template if not overridden during rendering.
|
||||||
|
special_tokens: Optional dictionary mapping token names to their string values.
|
||||||
|
These tokens are automatically added to the template variables.
|
||||||
|
"""
|
||||||
|
|
||||||
|
name: str
|
||||||
|
template_str: str
|
||||||
|
description: str = ""
|
||||||
|
default_variables: Dict[str, Any] = None
|
||||||
|
special_tokens: Dict[str, str] = None
|
||||||
|
|
||||||
|
def __post_init__(self):
|
||||||
|
if self.default_variables is None:
|
||||||
|
self.default_variables = {}
|
||||||
|
if self.special_tokens is None:
|
||||||
|
self.special_tokens = {}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_string(
|
||||||
|
cls,
|
||||||
|
template_str: str,
|
||||||
|
description: str = "",
|
||||||
|
default_variables: Optional[Dict[str, Any]] = None,
|
||||||
|
special_tokens: Optional[Dict[str, str]] = None,
|
||||||
|
) -> "ChatTemplate":
|
||||||
|
"""Create a ChatTemplate instance directly from a template string."""
|
||||||
|
return cls(
|
||||||
|
name="", # empty name for ad‑hoc templates
|
||||||
|
template_str=template_str,
|
||||||
|
description=description,
|
||||||
|
default_variables=default_variables,
|
||||||
|
special_tokens=special_tokens,
|
||||||
|
)
|
||||||
|
|
||||||
|
def render(
|
||||||
|
self,
|
||||||
|
messages: List[MessageType],
|
||||||
|
system_prompt: Optional[str] = None,
|
||||||
|
**extra_variables: Any,
|
||||||
|
) -> str:
|
||||||
|
"""Render the template with given messages and variables.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
messages: List of message dicts with 'role' and 'content'.
|
||||||
|
system_prompt: Optional system prompt string.
|
||||||
|
**extra_variables: Additional variables to pass to the template.
|
||||||
|
These override default_variables and special_tokens.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Rendered prompt string.
|
||||||
|
"""
|
||||||
|
# Merge default variables, special tokens, and extra variables
|
||||||
|
variables = {**self.default_variables, **self.special_tokens, **extra_variables}
|
||||||
|
variables["messages"] = messages
|
||||||
|
if system_prompt is not None:
|
||||||
|
variables["system_prompt"] = system_prompt
|
||||||
|
|
||||||
|
jinja_template = Template(self.template_str)
|
||||||
|
return jinja_template.render(**variables)
|
||||||
@@ -0,0 +1,247 @@
|
|||||||
|
"""
|
||||||
|
Tokenizer module with implementation and auto-loading support.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Dict, List, Optional, Union
|
||||||
|
|
||||||
|
from tokenizers import Tokenizer
|
||||||
|
|
||||||
|
from astrai.tokenize.chat_template import ChatTemplate
|
||||||
|
|
||||||
|
|
||||||
|
class AutoTokenizer:
|
||||||
|
"""Base tokenizer class with automatic loading support"""
|
||||||
|
|
||||||
|
TOKENIZER_CLASSES = {} # Registry for auto-loading
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
path: Optional[Union[str, Path]] = None,
|
||||||
|
special_token_map: Optional[Dict[str, str]] = None,
|
||||||
|
chat_template: Optional[str] = None,
|
||||||
|
):
|
||||||
|
self._tokenizer: Tokenizer = None
|
||||||
|
self._chat_template: Optional[ChatTemplate] = None
|
||||||
|
self._special_token_map: Optional[Dict] = special_token_map or {}
|
||||||
|
|
||||||
|
if chat_template:
|
||||||
|
self.set_chat_template(chat_template)
|
||||||
|
|
||||||
|
if path:
|
||||||
|
self.load(path)
|
||||||
|
|
||||||
|
def load(self, path: Union[str, Path]):
|
||||||
|
"""Load tokenizer from directory."""
|
||||||
|
path = Path(path)
|
||||||
|
tokenizer_file = path / "tokenizer.json"
|
||||||
|
config_file = path / "tokenizer_config.json"
|
||||||
|
self._tokenizer = Tokenizer.from_file(str(tokenizer_file))
|
||||||
|
|
||||||
|
if config_file.exists():
|
||||||
|
with open(config_file, "r", encoding="utf-8") as f:
|
||||||
|
config = json.load(f)
|
||||||
|
|
||||||
|
if "special_tokens" in config:
|
||||||
|
self._special_token_map.update(config["special_tokens"])
|
||||||
|
|
||||||
|
# Load chat template from config
|
||||||
|
if "chat_template" in config:
|
||||||
|
self.set_chat_template(config["chat_template"])
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_pretrained(cls, path: Union[str, Path], **kwargs) -> "AutoTokenizer":
|
||||||
|
"""Load tokenizer from pretrained directory."""
|
||||||
|
instance = cls(path)
|
||||||
|
return instance
|
||||||
|
|
||||||
|
def save_pretrained(self, save_path: str):
|
||||||
|
"""
|
||||||
|
Save tokenizer to pretrained directory.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
save_path: Path to save the tokenizer
|
||||||
|
"""
|
||||||
|
|
||||||
|
if self._tokenizer is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Tokenizer not initialized. Load or create a tokenizer first."
|
||||||
|
)
|
||||||
|
|
||||||
|
save_path = Path(save_path)
|
||||||
|
save_path.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
# Save tokenizer
|
||||||
|
self._tokenizer.save(str(save_path / "tokenizer.json"))
|
||||||
|
|
||||||
|
# Save tokenizer config
|
||||||
|
config = {}
|
||||||
|
if self._special_token_map is not None:
|
||||||
|
config["special_tokens"] = self._special_token_map
|
||||||
|
if self._chat_template is not None:
|
||||||
|
config["chat_template"] = self._chat_template.template_str
|
||||||
|
|
||||||
|
with open(save_path / "tokenizer_config.json", "w", encoding="utf-8") as f:
|
||||||
|
json.dump(config, f, ensure_ascii=False, indent=2)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def register_tokenizer(cls, name: str, tokenizer_class: type):
|
||||||
|
"""
|
||||||
|
Register a new tokenizer class.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
name: Name to register the tokenizer class under
|
||||||
|
tokenizer_class: The tokenizer class to register
|
||||||
|
"""
|
||||||
|
cls.TOKENIZER_CLASSES[name] = tokenizer_class
|
||||||
|
|
||||||
|
def encode(
|
||||||
|
self,
|
||||||
|
tokens: Union[str, List[str]],
|
||||||
|
out_ids: bool = True,
|
||||||
|
is_pretokenized: bool = False,
|
||||||
|
add_special_tokens: bool = True,
|
||||||
|
) -> List:
|
||||||
|
"""Encode text to tokens or token IDs."""
|
||||||
|
if self._tokenizer is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Tokenizer not initialized. Load or create a tokenizer first."
|
||||||
|
)
|
||||||
|
|
||||||
|
if isinstance(tokens, str):
|
||||||
|
encoded = self._tokenizer.encode(
|
||||||
|
tokens,
|
||||||
|
is_pretokenized=is_pretokenized,
|
||||||
|
add_special_tokens=add_special_tokens,
|
||||||
|
)
|
||||||
|
return encoded.ids if out_ids else encoded.tokens
|
||||||
|
else:
|
||||||
|
encoded_list = self._tokenizer.encode_batch(
|
||||||
|
tokens,
|
||||||
|
is_pretokenized=is_pretokenized,
|
||||||
|
add_special_tokens=add_special_tokens,
|
||||||
|
)
|
||||||
|
return [
|
||||||
|
encoded.ids if out_ids else encoded.tokens for encoded in encoded_list
|
||||||
|
]
|
||||||
|
|
||||||
|
def decode(self, tokens: List[int], skip_special_tokens: bool = True) -> str:
|
||||||
|
"""Decode token IDs to text."""
|
||||||
|
if self._tokenizer is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Tokenizer not initialized. Load or create a tokenizer first."
|
||||||
|
)
|
||||||
|
|
||||||
|
return self._tokenizer.decode(tokens, skip_special_tokens=skip_special_tokens)
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
if self._tokenizer is None:
|
||||||
|
return 0
|
||||||
|
return self._tokenizer.get_vocab_size()
|
||||||
|
|
||||||
|
def __getattr__(self, key: str):
|
||||||
|
"""
|
||||||
|
Dynamically intercept special token attribute access.
|
||||||
|
Supports three forms:
|
||||||
|
- tokenizer.bos_token → returns string
|
||||||
|
- tokenizer.bos_token_id → returns corresponding integer ID
|
||||||
|
- tokenizer.stop_ids → returns list of corresponding integer IDs for all special tokens
|
||||||
|
"""
|
||||||
|
# Handle stop_ids - return IDs for all special tokens
|
||||||
|
if key == "stop_ids":
|
||||||
|
stop_ids = []
|
||||||
|
|
||||||
|
if self._tokenizer is None:
|
||||||
|
return stop_ids
|
||||||
|
|
||||||
|
for val in self._special_token_map.values():
|
||||||
|
token_id = self._tokenizer.token_to_id(val)
|
||||||
|
if token_id is not None:
|
||||||
|
stop_ids.append(token_id)
|
||||||
|
|
||||||
|
return stop_ids
|
||||||
|
|
||||||
|
# Handle _id suffix (e.g., bos_token_id -> bos_token)
|
||||||
|
if key.endswith("_id"):
|
||||||
|
base_attr = key[:-3] # Remove "_id"
|
||||||
|
token_str = self._special_token_map.get(base_attr)
|
||||||
|
if token_str is None:
|
||||||
|
return None
|
||||||
|
if self._tokenizer is None:
|
||||||
|
raise RuntimeError("Tokenizer not loaded, cannot convert token to id.")
|
||||||
|
return self._tokenizer.token_to_id(token_str)
|
||||||
|
|
||||||
|
# Handle regular string attributes
|
||||||
|
if key in self._special_token_map:
|
||||||
|
return self._special_token_map.get(key)
|
||||||
|
|
||||||
|
# Other attributes trigger default AttributeError
|
||||||
|
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{key}'")
|
||||||
|
|
||||||
|
@property
|
||||||
|
def vocab_size(self) -> int:
|
||||||
|
return len(self)
|
||||||
|
|
||||||
|
def set_chat_template(self, template: Union[str, ChatTemplate]):
|
||||||
|
"""
|
||||||
|
Set the chat template for the tokenizer.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
template: Either a template name (str) registered in the global registry,
|
||||||
|
or a ChatTemplate instance, or a Jinja2 template string.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
KeyError: If template name is not registered.
|
||||||
|
"""
|
||||||
|
if isinstance(template, str):
|
||||||
|
self._chat_template = ChatTemplate.from_string(template)
|
||||||
|
elif isinstance(template, ChatTemplate):
|
||||||
|
self._chat_template = template
|
||||||
|
else:
|
||||||
|
raise ValueError("Invalid template type, must be str or ChatTemplate.")
|
||||||
|
|
||||||
|
def apply_chat_template(
|
||||||
|
self,
|
||||||
|
messages: List[Dict[str, str]],
|
||||||
|
system_prompt: Optional[str] = None,
|
||||||
|
tokenize: bool = True,
|
||||||
|
add_generation_prompt: bool = True,
|
||||||
|
**kwargs,
|
||||||
|
) -> Union[str, List[int]]:
|
||||||
|
"""
|
||||||
|
Apply the chat template to messages and optionally tokenize the result.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
messages: List of message dicts with 'role' and 'content'.
|
||||||
|
system_prompt: Optional system prompt string (auto-converted to first message).
|
||||||
|
tokenize: Whether to return token IDs (True) or raw string (False).
|
||||||
|
add_generation_prompt: Whether to add the generation prompt (default: True).
|
||||||
|
**kwargs: Additional variables to pass to the template.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Either the rendered string or list of token IDs.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
RuntimeError: If chat template is not set.
|
||||||
|
"""
|
||||||
|
if self._chat_template is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Chat template not set. Use set_chat_template() to set a template first."
|
||||||
|
)
|
||||||
|
|
||||||
|
# Auto-convert system_prompt to first message if provided
|
||||||
|
if system_prompt:
|
||||||
|
messages = [{"role": "system", "content": system_prompt}] + list(messages)
|
||||||
|
|
||||||
|
# Render the template
|
||||||
|
rendered = self._chat_template.render(
|
||||||
|
messages=messages,
|
||||||
|
add_generation_prompt=add_generation_prompt,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
if tokenize:
|
||||||
|
return self.encode(rendered)
|
||||||
|
|
||||||
|
return rendered
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
from astrai.trainer.schedule import BaseScheduler, SchedulerFactory
|
||||||
|
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
||||||
|
from astrai.trainer.train_callback import (
|
||||||
|
CallbackFactory,
|
||||||
|
TrainCallback,
|
||||||
|
)
|
||||||
|
from astrai.trainer.trainer import Trainer
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
# Main trainer
|
||||||
|
"Trainer",
|
||||||
|
# Strategy factory
|
||||||
|
"StrategyFactory",
|
||||||
|
"BaseStrategy",
|
||||||
|
# Scheduler factory
|
||||||
|
"SchedulerFactory",
|
||||||
|
"BaseScheduler",
|
||||||
|
# Callback factory
|
||||||
|
"TrainCallback",
|
||||||
|
"CallbackFactory",
|
||||||
|
]
|
||||||
@@ -1,8 +1,10 @@
|
|||||||
import torch.nn as nn
|
|
||||||
from typing import Dict
|
from typing import Dict
|
||||||
|
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
|
||||||
def grad_norm(model: nn.Module, norm_type: int = 2) -> Dict[str, float]:
|
def grad_norm(model: nn.Module, norm_type: int = 2) -> Dict[str, float]:
|
||||||
""" Compute gradient norm for each parameter in the model. """
|
"""Compute gradient norm for each parameter in the model."""
|
||||||
norms = {}
|
norms = {}
|
||||||
for name, param in model.named_parameters():
|
for name, param in model.named_parameters():
|
||||||
norms[name] = 0.0
|
norms[name] = 0.0
|
||||||
@@ -11,8 +13,9 @@ def grad_norm(model: nn.Module, norm_type: int = 2) -> Dict[str, float]:
|
|||||||
norms[name] = norm
|
norms[name] = norm
|
||||||
return norms
|
return norms
|
||||||
|
|
||||||
|
|
||||||
def grad_std(model: nn.Module) -> Dict[str, float]:
|
def grad_std(model: nn.Module) -> Dict[str, float]:
|
||||||
""" Compute standard deviation of gradients for each parameter. """
|
"""Compute standard deviation of gradients for each parameter."""
|
||||||
stds = {}
|
stds = {}
|
||||||
for name, param in model.named_parameters():
|
for name, param in model.named_parameters():
|
||||||
stds[name] = 0.0
|
stds[name] = 0.0
|
||||||
@@ -21,45 +24,81 @@ def grad_std(model: nn.Module) -> Dict[str, float]:
|
|||||||
stds[name] = std
|
stds[name] = std
|
||||||
return stds
|
return stds
|
||||||
|
|
||||||
|
|
||||||
def grad_max(model: nn.Module) -> Dict[str, float]:
|
def grad_max(model: nn.Module) -> Dict[str, float]:
|
||||||
""" Find the maximum absolute gradient value for each parameter. """
|
"""Find the maximum absolute gradient value for each parameter."""
|
||||||
max_vals = {}
|
max_vals = {}
|
||||||
for name, param in model.named_parameters():
|
for name, param in model.named_parameters():
|
||||||
max_vals[name] = -float('inf')
|
max_vals[name] = -float("inf")
|
||||||
if param.grad:
|
if param.grad:
|
||||||
max_val = param.grad.data.max().item()
|
max_val = param.grad.data.max().item()
|
||||||
max_vals[name] = max_val
|
max_vals[name] = max_val
|
||||||
|
|
||||||
return max_vals
|
return max_vals
|
||||||
|
|
||||||
|
|
||||||
def grad_min(model: nn.Module) -> Dict[str, float]:
|
def grad_min(model: nn.Module) -> Dict[str, float]:
|
||||||
""" Find the minimum absolute gradient value for each parameter. """
|
"""Find the minimum absolute gradient value for each parameter."""
|
||||||
min_vals = {}
|
min_vals = {}
|
||||||
for name, param in model.named_parameters():
|
for name, param in model.named_parameters():
|
||||||
min_vals[name] = float('inf')
|
min_vals[name] = float("inf")
|
||||||
if param.grad:
|
if param.grad:
|
||||||
min_val = param.grad.data.min().item()
|
min_val = param.grad.data.min().item()
|
||||||
min_vals[name] = min_val
|
min_vals[name] = min_val
|
||||||
|
|
||||||
return min_vals
|
return min_vals
|
||||||
|
|
||||||
|
|
||||||
def grad_mean(model: nn.Module) -> Dict[str, float]:
|
def grad_mean(model: nn.Module) -> Dict[str, float]:
|
||||||
""" Compute mean of gradients for each parameter. """
|
"""Compute mean of gradients for each parameter."""
|
||||||
means = {}
|
means = {}
|
||||||
for name, param in model.named_parameters():
|
for name, param in model.named_parameters():
|
||||||
means[name] = 0.0
|
means[name] = 0.0
|
||||||
if param.grad:
|
if param.grad:
|
||||||
mean = param.grad.data.mean().item()
|
mean = param.grad.data.mean().item()
|
||||||
means[name] = mean
|
means[name] = mean
|
||||||
|
|
||||||
return means
|
return means
|
||||||
|
|
||||||
|
|
||||||
def grad_nan_num(model: nn.Module) -> Dict[str, int]:
|
def grad_nan_num(model: nn.Module) -> Dict[str, int]:
|
||||||
""" Count the number of NaNs in gradients for each parameter. """
|
"""Count the number of NaNs in gradients for each parameter."""
|
||||||
nan_nums = {}
|
nan_nums = {}
|
||||||
for name, param in model.named_parameters():
|
for name, param in model.named_parameters():
|
||||||
nan_nums[name] = 0
|
nan_nums[name] = 0
|
||||||
if param.grad:
|
if param.grad:
|
||||||
nan_num = param.grad.isnan().sum().item()
|
nan_num = param.grad.isnan().sum().item()
|
||||||
nan_nums[name] = nan_num
|
nan_nums[name] = nan_num
|
||||||
return nan_nums
|
return nan_nums
|
||||||
|
|
||||||
|
|
||||||
|
def ctx_get_loss(ctx):
|
||||||
|
return ctx.loss
|
||||||
|
|
||||||
|
|
||||||
|
def ctx_get_lr(ctx):
|
||||||
|
return ctx.optimizer.param_groups[-1]["lr"]
|
||||||
|
|
||||||
|
|
||||||
|
def ctx_get_grad_norm(ctx):
|
||||||
|
return grad_norm(ctx.model)
|
||||||
|
|
||||||
|
|
||||||
|
def ctx_get_grad_std(ctx):
|
||||||
|
return grad_std(ctx.model)
|
||||||
|
|
||||||
|
|
||||||
|
def ctx_get_grad_max(ctx):
|
||||||
|
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,194 @@
|
|||||||
|
"""Learning rate scheduler implementations with factory pattern."""
|
||||||
|
|
||||||
|
import math
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from typing import Any, Dict, List, Type
|
||||||
|
|
||||||
|
from torch.optim.lr_scheduler import LRScheduler
|
||||||
|
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
|
||||||
|
|
||||||
|
class BaseScheduler(LRScheduler, ABC):
|
||||||
|
"""Base scheduler class for all other schedulers."""
|
||||||
|
|
||||||
|
def __init__(self, optimizer, last_epoch: int = -1):
|
||||||
|
super().__init__(optimizer, last_epoch)
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_lr(self) -> List[float]:
|
||||||
|
"""Calculate the current learning rate."""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def state_dict(self) -> Dict[str, Any]:
|
||||||
|
return super().state_dict()
|
||||||
|
|
||||||
|
def load_state_dict(self, state_dict: Dict[str, Any]):
|
||||||
|
super().load_state_dict(state_dict)
|
||||||
|
|
||||||
|
|
||||||
|
class SchedulerFactory(BaseFactory["BaseScheduler"]):
|
||||||
|
"""Factory class for creating learning rate schedulers.
|
||||||
|
|
||||||
|
Supports decorator-based registration for extensible scheduler types.
|
||||||
|
Also supports creation from ScheduleConfig objects.
|
||||||
|
|
||||||
|
Example usage:
|
||||||
|
@SchedulerFactory.register("custom")
|
||||||
|
class CustomScheduler(BaseScheduler):
|
||||||
|
...
|
||||||
|
|
||||||
|
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 -----------
|
||||||
|
|
||||||
|
|
||||||
|
@SchedulerFactory.register("cosine")
|
||||||
|
class CosineScheduler(BaseScheduler):
|
||||||
|
"""Cosine decay scheduler with warmup, implemented as PyTorch LRScheduler."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
optimizer,
|
||||||
|
warmup_steps: int,
|
||||||
|
lr_decay_steps: int,
|
||||||
|
min_rate: float = 0.05,
|
||||||
|
last_epoch: int = -1,
|
||||||
|
):
|
||||||
|
self.warmup_steps = warmup_steps
|
||||||
|
self.lr_decay_steps = lr_decay_steps
|
||||||
|
self.min_rate = min_rate
|
||||||
|
self.total_steps = warmup_steps + lr_decay_steps
|
||||||
|
super().__init__(optimizer, last_epoch)
|
||||||
|
|
||||||
|
def get_lr(self) -> List[float]:
|
||||||
|
# warmup
|
||||||
|
if self.last_epoch < self.warmup_steps:
|
||||||
|
warmup_factor = max(self.min_rate, self.last_epoch / self.warmup_steps)
|
||||||
|
return [base_lr * warmup_factor for base_lr in self.base_lrs]
|
||||||
|
|
||||||
|
# cosine decay
|
||||||
|
decay_progress = (self.last_epoch - self.warmup_steps) / self.lr_decay_steps
|
||||||
|
decay_progress = min(decay_progress, 1.0)
|
||||||
|
cosine_decay = 0.5 * (1.0 + math.cos(math.pi * decay_progress))
|
||||||
|
decay_factor = max(self.min_rate, cosine_decay)
|
||||||
|
return [base_lr * decay_factor for base_lr in self.base_lrs]
|
||||||
|
|
||||||
|
def state_dict(self):
|
||||||
|
state = super().state_dict()
|
||||||
|
state.update(
|
||||||
|
{
|
||||||
|
"warmup_steps": self.warmup_steps,
|
||||||
|
"lr_decay_steps": self.lr_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.lr_decay_steps = state_dict.pop("lr_decay_steps")
|
||||||
|
self.min_rate = state_dict.pop("min_rate")
|
||||||
|
self.total_steps = state_dict.pop("total_steps")
|
||||||
|
super().load_state_dict(state_dict)
|
||||||
|
|
||||||
|
|
||||||
|
@SchedulerFactory.register("sgdr")
|
||||||
|
class SGDRScheduler(BaseScheduler):
|
||||||
|
"""SGDR (Stochastic Gradient Descent with Warm Restarts) scheduler."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
optimizer,
|
||||||
|
warmup_steps: int,
|
||||||
|
cycle_length: int,
|
||||||
|
min_rate: float = 0.05,
|
||||||
|
t_mult: int = 2,
|
||||||
|
last_epoch: int = -1,
|
||||||
|
):
|
||||||
|
self.warmup_steps = warmup_steps
|
||||||
|
self.cycle_length = cycle_length
|
||||||
|
self.min_rate = min_rate
|
||||||
|
self.t_mult = t_mult
|
||||||
|
|
||||||
|
super().__init__(optimizer, last_epoch)
|
||||||
|
|
||||||
|
def get_lr(self):
|
||||||
|
# warmup
|
||||||
|
if self.last_epoch < self.warmup_steps:
|
||||||
|
warmup_factor = max(self.min_rate, self.last_epoch / self.warmup_steps)
|
||||||
|
return [base_lr * warmup_factor for base_lr in self.base_lrs]
|
||||||
|
|
||||||
|
# SGDR
|
||||||
|
steps_since_warmup = self.last_epoch - self.warmup_steps
|
||||||
|
|
||||||
|
# 1. Calculate current cycle and position within cycle
|
||||||
|
current_cycle_length = self.cycle_length
|
||||||
|
total_cycles_length = 0
|
||||||
|
cycle_num = 0
|
||||||
|
|
||||||
|
while total_cycles_length + current_cycle_length <= steps_since_warmup:
|
||||||
|
total_cycles_length += current_cycle_length
|
||||||
|
current_cycle_length *= self.t_mult
|
||||||
|
cycle_num += 1
|
||||||
|
|
||||||
|
steps_in_cycle = steps_since_warmup - total_cycles_length
|
||||||
|
|
||||||
|
# 2. Cosine annealing within the current cycle
|
||||||
|
cosine_factor = 0.5 * (
|
||||||
|
1 + math.cos(math.pi * steps_in_cycle / current_cycle_length)
|
||||||
|
)
|
||||||
|
learning_rate_factor = self.min_rate + (1 - self.min_rate) * cosine_factor
|
||||||
|
|
||||||
|
return [base_lr * learning_rate_factor for base_lr in self.base_lrs]
|
||||||
|
|
||||||
|
def state_dict(self):
|
||||||
|
"""Returns the state of the scheduler as a dict."""
|
||||||
|
state = super().state_dict()
|
||||||
|
state.update(
|
||||||
|
{
|
||||||
|
"warmup_steps": self.warmup_steps,
|
||||||
|
"cycle_length": self.cycle_length,
|
||||||
|
"min_rate": self.min_rate,
|
||||||
|
"t_mult": self.t_mult,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return state
|
||||||
|
|
||||||
|
def load_state_dict(self, state_dict):
|
||||||
|
"""Loads the scheduler's state."""
|
||||||
|
self.warmup_steps = state_dict.pop("warmup_steps")
|
||||||
|
self.cycle_length = state_dict.pop("cycle_length")
|
||||||
|
self.min_rate = state_dict.pop("min_rate")
|
||||||
|
self.t_mult = state_dict.pop("t_mult")
|
||||||
|
super().load_state_dict(state_dict)
|
||||||
@@ -0,0 +1,342 @@
|
|||||||
|
"""Training strategy implementations with factory pattern."""
|
||||||
|
|
||||||
|
import copy
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from typing import Any, Callable, Dict, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||||
|
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
|
||||||
|
|
||||||
|
def unwrap_model(model: nn.Module) -> nn.Module:
|
||||||
|
"""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."""
|
||||||
|
return {key: value.to(device, non_blocking=True) for key, value in batch.items()}
|
||||||
|
|
||||||
|
|
||||||
|
def get_logprobs(
|
||||||
|
model: Union[nn.Module, Callable[..., Dict[str, Tensor]]],
|
||||||
|
input_ids: Tensor,
|
||||||
|
mask: Tensor,
|
||||||
|
reduction: str,
|
||||||
|
):
|
||||||
|
"""Compute token-wise log probabilities from model outputs.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model: The language model
|
||||||
|
input_ids: Input token IDs of shape [batch_size, seq_len]
|
||||||
|
mask: Attention mask of shape [batch_size, seq_len]
|
||||||
|
reduction: How to reduce over sequence dimension ("mean", "sum", "none")
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Log probabilities with reduction applied over sequence dimension
|
||||||
|
"""
|
||||||
|
allowed_reductions = ["mean", "sum", "none"]
|
||||||
|
if reduction not in allowed_reductions:
|
||||||
|
raise ValueError(
|
||||||
|
f"reduction must be one of {allowed_reductions}, got '{reduction}'"
|
||||||
|
)
|
||||||
|
|
||||||
|
shifted_input_ids = input_ids[:, 1:]
|
||||||
|
shifted_mask = mask[:, 1:]
|
||||||
|
|
||||||
|
logits = model(input_ids[:, :-1], mask[:, :-1])["logits"]
|
||||||
|
log_probs = torch.log_softmax(logits.float(), dim=-1)
|
||||||
|
|
||||||
|
token_logprobs = torch.gather(
|
||||||
|
log_probs, dim=-1, index=shifted_input_ids.unsqueeze(-1)
|
||||||
|
).squeeze(-1)
|
||||||
|
|
||||||
|
if reduction == "mean":
|
||||||
|
return (token_logprobs * shifted_mask).sum(dim=-1) / shifted_mask.sum(
|
||||||
|
dim=-1
|
||||||
|
).clamp(min=1.0)
|
||||||
|
elif reduction == "sum":
|
||||||
|
return (token_logprobs * shifted_mask).sum(dim=-1)
|
||||||
|
else:
|
||||||
|
return token_logprobs * shifted_mask
|
||||||
|
|
||||||
|
|
||||||
|
class BaseStrategy(ABC):
|
||||||
|
"""Abstract base class for training strategies."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self, model: Union[Callable[..., Dict[str, Tensor]]], device: str, **kwargs
|
||||||
|
):
|
||||||
|
self.model = model
|
||||||
|
self.device = device
|
||||||
|
self.extra_kwargs = kwargs
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
|
"""Compute loss for the given batch.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
batch: Dictionary containing batch tensors
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Computed loss tensor
|
||||||
|
"""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def __call__(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
|
"""Allow calling strategy directly as a callable."""
|
||||||
|
return self.compute_loss(batch)
|
||||||
|
|
||||||
|
|
||||||
|
class StrategyFactory(BaseFactory["BaseStrategy"]):
|
||||||
|
"""Factory class for creating training strategy instances.
|
||||||
|
|
||||||
|
Supports decorator-based registration for extensible strategy types.
|
||||||
|
All default strategies (seq, sft, dpo, grpo) are automatically registered.
|
||||||
|
|
||||||
|
Example usage:
|
||||||
|
@StrategyFactory.register("custom")
|
||||||
|
class CustomStrategy(BaseStrategy):
|
||||||
|
...
|
||||||
|
|
||||||
|
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 ==============
|
||||||
|
# All strategies are registered at class definition time using the decorator
|
||||||
|
|
||||||
|
|
||||||
|
@StrategyFactory.register("seq")
|
||||||
|
class SEQStrategy(BaseStrategy):
|
||||||
|
"""Standard next-token prediction training strategy.
|
||||||
|
|
||||||
|
Computes cross-entropy loss for next token prediction.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, model, device, label_smoothing: float = 0.0, **kwargs):
|
||||||
|
super().__init__(model, device, **kwargs)
|
||||||
|
self.label_smoothing = label_smoothing
|
||||||
|
|
||||||
|
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
|
batch = move_to_device(batch, self.device)
|
||||||
|
input_ids, target_ids = batch["input_ids"], batch["target_ids"]
|
||||||
|
logits = self.model(input_ids=input_ids)["logits"]
|
||||||
|
|
||||||
|
loss = F.cross_entropy(
|
||||||
|
input=logits.flatten(0, 1).float(),
|
||||||
|
target=target_ids.flatten(),
|
||||||
|
label_smoothing=self.label_smoothing,
|
||||||
|
)
|
||||||
|
|
||||||
|
return loss
|
||||||
|
|
||||||
|
|
||||||
|
@StrategyFactory.register("sft")
|
||||||
|
class SFTStrategy(BaseStrategy):
|
||||||
|
"""Supervised Fine-tuning strategy with loss masking.
|
||||||
|
|
||||||
|
Applies cross-entropy loss only to tokens where loss_mask is True.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, model, device, label_smoothing: float = 0.0, **kwargs):
|
||||||
|
super().__init__(model, device, **kwargs)
|
||||||
|
self.label_smoothing = label_smoothing
|
||||||
|
|
||||||
|
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
|
batch = move_to_device(batch, self.device)
|
||||||
|
input_ids, target_ids, loss_mask = (
|
||||||
|
batch["input_ids"],
|
||||||
|
batch["target_ids"],
|
||||||
|
batch["loss_mask"],
|
||||||
|
)
|
||||||
|
|
||||||
|
ignore_index = -100
|
||||||
|
logits = self.model(input_ids=input_ids)["logits"]
|
||||||
|
target_ids = target_ids.masked_fill(loss_mask == 0, ignore_index)
|
||||||
|
|
||||||
|
loss = F.cross_entropy(
|
||||||
|
input=logits.flatten(0, 1).float(),
|
||||||
|
target=target_ids.flatten(),
|
||||||
|
ignore_index=ignore_index,
|
||||||
|
label_smoothing=self.label_smoothing,
|
||||||
|
)
|
||||||
|
|
||||||
|
return loss
|
||||||
|
|
||||||
|
|
||||||
|
@StrategyFactory.register("dpo")
|
||||||
|
class DPOStrategy(BaseStrategy):
|
||||||
|
"""Direct Preference Optimization strategy.
|
||||||
|
|
||||||
|
Implements the DPO loss from the paper "Direct Preference Optimization".
|
||||||
|
Uses a reference model to compute KL divergence penalty.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: nn.Module,
|
||||||
|
device: str,
|
||||||
|
beta: float = 0.1,
|
||||||
|
reduction: str = "mean",
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
super().__init__(model, device, **kwargs)
|
||||||
|
self.ref_model = create_ref_model(model)
|
||||||
|
self.beta = beta
|
||||||
|
self.reduction = reduction
|
||||||
|
|
||||||
|
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||||
|
batch = move_to_device(batch, self.device)
|
||||||
|
chosen_ids, rejected_ids = batch["chosen"], batch["rejected"]
|
||||||
|
chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"]
|
||||||
|
|
||||||
|
concat_ids = torch.cat([chosen_ids, rejected_ids], dim=0)
|
||||||
|
concat_mask = torch.cat([chosen_mask, rejected_mask], dim=0)
|
||||||
|
|
||||||
|
log_pi = get_logprobs(self.model, concat_ids, concat_mask, self.reduction)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
log_ref = get_logprobs(
|
||||||
|
self.ref_model, concat_ids, concat_mask, self.reduction
|
||||||
|
)
|
||||||
|
|
||||||
|
log_pi_chosen = log_pi[: chosen_ids.shape[0]]
|
||||||
|
log_pi_rejected = log_pi[chosen_ids.shape[0] :]
|
||||||
|
log_ref_chosen = log_ref[: chosen_ids.shape[0]]
|
||||||
|
log_ref_rejected = log_ref[chosen_ids.shape[0] :]
|
||||||
|
|
||||||
|
pi_log_ratio = log_pi_chosen - log_pi_rejected
|
||||||
|
ref_log_ratio = log_ref_chosen - log_ref_rejected
|
||||||
|
|
||||||
|
ratio_diff = pi_log_ratio - ref_log_ratio
|
||||||
|
dpo_loss = -F.logsigmoid(self.beta * ratio_diff).mean()
|
||||||
|
|
||||||
|
return dpo_loss
|
||||||
|
|
||||||
|
|
||||||
|
@StrategyFactory.register("grpo")
|
||||||
|
class GRPOStrategy(BaseStrategy):
|
||||||
|
"""Group Relative Policy Optimization strategy.
|
||||||
|
|
||||||
|
On-policy GRPO following DeepSeek-R1: the policy model is updated while
|
||||||
|
a frozen ref_model stores the old-policy log-probs. ratio = exp(logπ_θ - logπ_ref),
|
||||||
|
clipped PPO objective. Call ``sync_ref_model()`` after each data-generation round.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: nn.Module,
|
||||||
|
device: str,
|
||||||
|
clip_eps: float = 0.2,
|
||||||
|
kl_coef: float = 0.01,
|
||||||
|
group_size: int = 4,
|
||||||
|
reduction: str = "mean",
|
||||||
|
sync_interval: int = 200,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
super().__init__(model, device, **kwargs)
|
||||||
|
self.ref_model = create_ref_model(model)
|
||||||
|
self.clip_eps = clip_eps
|
||||||
|
self.kl_coef = kl_coef
|
||||||
|
self.group_size = group_size
|
||||||
|
self.reduction = reduction
|
||||||
|
self.sync_interval = sync_interval
|
||||||
|
self._step = 0
|
||||||
|
|
||||||
|
def sync_ref_model(self):
|
||||||
|
"""Copy current model weights to ref model."""
|
||||||
|
ref_state = self.model.state_dict()
|
||||||
|
self.ref_model.load_state_dict(ref_state)
|
||||||
|
|
||||||
|
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)
|
||||||
|
prompts = batch["prompts"]
|
||||||
|
responses = batch["responses"]
|
||||||
|
masks = batch["masks"]
|
||||||
|
rewards = batch["rewards"]
|
||||||
|
|
||||||
|
batch_size, group_size, response_len = responses.shape
|
||||||
|
responses_flat = responses.view(-1, response_len)
|
||||||
|
masks_flat = masks.view(-1, response_len)
|
||||||
|
prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1)
|
||||||
|
|
||||||
|
full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1)
|
||||||
|
full_masks = torch.cat([torch.ones_like(prompt_expanded), masks_flat], dim=-1)
|
||||||
|
|
||||||
|
log_probs_policy = get_logprobs(
|
||||||
|
self.model, full_sequences, full_masks, self.reduction
|
||||||
|
)
|
||||||
|
log_probs_policy = log_probs_policy.view(batch_size, group_size)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
log_probs_ref = get_logprobs(
|
||||||
|
self.ref_model, full_sequences, full_masks, self.reduction
|
||||||
|
)
|
||||||
|
log_probs_ref = log_probs_ref.view(batch_size, group_size)
|
||||||
|
|
||||||
|
eps = torch.finfo(log_probs_policy.dtype).eps
|
||||||
|
mean = rewards.mean(dim=-1, keepdim=True)
|
||||||
|
std = rewards.std(dim=-1, keepdim=True)
|
||||||
|
advantages = (rewards - mean) / (std + eps)
|
||||||
|
|
||||||
|
ratio = torch.exp(log_probs_policy - log_probs_ref)
|
||||||
|
|
||||||
|
surr1 = ratio * advantages
|
||||||
|
surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * advantages
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
return total_loss
|
||||||
@@ -0,0 +1,266 @@
|
|||||||
|
import json
|
||||||
|
import os
|
||||||
|
import time
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Callable, List, Optional, Protocol, runtime_checkable
|
||||||
|
|
||||||
|
import torch.nn as nn
|
||||||
|
from torch.nn.utils import clip_grad_norm_
|
||||||
|
from tqdm import tqdm
|
||||||
|
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.parallel import only_on_rank
|
||||||
|
from astrai.serialization import Checkpoint
|
||||||
|
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_std,
|
||||||
|
ctx_get_loss,
|
||||||
|
ctx_get_lr,
|
||||||
|
)
|
||||||
|
from astrai.trainer.train_context import TrainContext
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class TrainCallback(Protocol):
|
||||||
|
"""
|
||||||
|
Callback interface for trainer.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def on_train_begin(self, context: TrainContext):
|
||||||
|
"""Called at the beginning of training."""
|
||||||
|
|
||||||
|
def on_train_end(self, context: TrainContext):
|
||||||
|
"""Called at the end of training."""
|
||||||
|
|
||||||
|
def on_epoch_begin(self, context: TrainContext):
|
||||||
|
"""Called at the beginning of each epoch."""
|
||||||
|
|
||||||
|
def on_epoch_end(self, context: TrainContext):
|
||||||
|
"""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):
|
||||||
|
"""Called at the beginning of each batch."""
|
||||||
|
|
||||||
|
def on_batch_end(self, context: TrainContext):
|
||||||
|
"""Called at the end of each batch."""
|
||||||
|
|
||||||
|
def on_error(self, context: TrainContext):
|
||||||
|
"""Called when an error occurs during training."""
|
||||||
|
|
||||||
|
|
||||||
|
class CallbackFactory(BaseFactory[TrainCallback]):
|
||||||
|
"""Factory for registering and creating training callbacks.
|
||||||
|
|
||||||
|
Example:
|
||||||
|
@CallbackFactory.register("my_callback")
|
||||||
|
class MyCallback(TrainCallback):
|
||||||
|
...
|
||||||
|
|
||||||
|
callback = CallbackFactory.create("my_callback", **kwargs)
|
||||||
|
"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _validate_component(cls, callback_cls: type) -> None:
|
||||||
|
"""Validate that the callback class inherits from TrainCallback."""
|
||||||
|
if not issubclass(callback_cls, TrainCallback):
|
||||||
|
raise TypeError(f"{callback_cls.__name__} must inherit from TrainCallback")
|
||||||
|
|
||||||
|
|
||||||
|
@CallbackFactory.register("gradient_clipping")
|
||||||
|
class GradientClippingCallback(TrainCallback):
|
||||||
|
"""
|
||||||
|
Gradient clipping callback for trainer.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, max_grad_norm: float):
|
||||||
|
self.max_grad_norm = max_grad_norm
|
||||||
|
|
||||||
|
def on_step_begin(self, context: TrainContext):
|
||||||
|
_ = context
|
||||||
|
clip_grad_norm_(context.model.parameters(), self.max_grad_norm)
|
||||||
|
|
||||||
|
|
||||||
|
@CallbackFactory.register("scheduler")
|
||||||
|
class SchedulerCallback(TrainCallback):
|
||||||
|
"""
|
||||||
|
Scheduler callback for trainer.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def on_train_begin(self, context: TrainContext):
|
||||||
|
for group in context.optimizer.param_groups:
|
||||||
|
if "initial_lr" not in group:
|
||||||
|
group["initial_lr"] = group["lr"]
|
||||||
|
|
||||||
|
def on_batch_end(self, context: TrainContext):
|
||||||
|
if context.scheduler:
|
||||||
|
context.scheduler.step()
|
||||||
|
|
||||||
|
|
||||||
|
@CallbackFactory.register("checkpoint")
|
||||||
|
class CheckpointCallback(TrainCallback):
|
||||||
|
"""
|
||||||
|
Checkpoint callback for trainer.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
save_dir: str,
|
||||||
|
interval: int,
|
||||||
|
weight_only: bool = False,
|
||||||
|
state_dict_fn: Optional[Callable[[nn.Module], dict]] = None,
|
||||||
|
save_extra_fn: Optional[Callable[["TrainContext"], dict]] = None,
|
||||||
|
):
|
||||||
|
self.save_dir = save_dir
|
||||||
|
self.interval = interval
|
||||||
|
self.weight_only = weight_only
|
||||||
|
self.state_dict_fn = state_dict_fn
|
||||||
|
self.save_extra_fn = save_extra_fn
|
||||||
|
self.last_ckpt_iter = 0
|
||||||
|
|
||||||
|
@only_on_rank(0)
|
||||||
|
def _save_checkpoint(self, context: TrainContext):
|
||||||
|
save_path = os.path.join(
|
||||||
|
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
|
||||||
|
context.checkpoint = Checkpoint(
|
||||||
|
state_dict=state_dict,
|
||||||
|
epoch=context.epoch,
|
||||||
|
iteration=context.iteration,
|
||||||
|
extra=extra,
|
||||||
|
)
|
||||||
|
|
||||||
|
context.checkpoint.save(save_path)
|
||||||
|
self.last_ckpt_iter = context.iteration
|
||||||
|
|
||||||
|
def on_batch_end(self, context: TrainContext):
|
||||||
|
if context.iteration - self.last_ckpt_iter >= self.interval:
|
||||||
|
self._save_checkpoint(context)
|
||||||
|
|
||||||
|
def on_train_end(self, context: TrainContext):
|
||||||
|
if context.iteration != self.last_ckpt_iter:
|
||||||
|
self._save_checkpoint(context)
|
||||||
|
|
||||||
|
def on_error(self, context: TrainContext):
|
||||||
|
self._save_checkpoint(context)
|
||||||
|
|
||||||
|
|
||||||
|
@CallbackFactory.register("progress_bar")
|
||||||
|
class ProgressBarCallback(TrainCallback):
|
||||||
|
"""
|
||||||
|
Progress bar callback for trainer.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, num_epoch: int):
|
||||||
|
self.num_epoch = num_epoch
|
||||||
|
self.progress_bar: tqdm = None
|
||||||
|
|
||||||
|
@only_on_rank(0)
|
||||||
|
def on_epoch_begin(self, context: TrainContext):
|
||||||
|
self.progress_bar = tqdm(
|
||||||
|
context.dataloader,
|
||||||
|
desc=f"Epoch {context.epoch + 1}/{self.num_epoch}",
|
||||||
|
dynamic_ncols=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
@only_on_rank(0)
|
||||||
|
def on_batch_end(self, context: TrainContext):
|
||||||
|
self.progress_bar.set_postfix(
|
||||||
|
{
|
||||||
|
"loss": f"{context.loss:.4f}",
|
||||||
|
"lr": f"{context.optimizer.param_groups[-1]['lr']:.2e}",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
self.progress_bar.update(1)
|
||||||
|
|
||||||
|
@only_on_rank(0)
|
||||||
|
def on_epoch_end(self, context: TrainContext):
|
||||||
|
_ = context
|
||||||
|
if self.progress_bar:
|
||||||
|
self.progress_bar.close()
|
||||||
|
|
||||||
|
|
||||||
|
@CallbackFactory.register("metric_logger")
|
||||||
|
class MetricLoggerCallback(TrainCallback):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
log_dir: str,
|
||||||
|
save_interval: int,
|
||||||
|
log_interval: int = 10,
|
||||||
|
metrics: List[str] = None,
|
||||||
|
):
|
||||||
|
self.last_log_iter = 0
|
||||||
|
self.save_interval = save_interval
|
||||||
|
self.log_interval = log_interval
|
||||||
|
self.metrics = metrics or ["loss", "lr"]
|
||||||
|
|
||||||
|
self.log_dir = Path(log_dir) if log_dir else Path.cwd() / "logs"
|
||||||
|
self.log_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
self.log_cache = []
|
||||||
|
|
||||||
|
self._metric_funcs = {
|
||||||
|
"loss": ctx_get_loss,
|
||||||
|
"lr": ctx_get_lr,
|
||||||
|
"grad_norm": ctx_get_grad_norm,
|
||||||
|
"grad_std": ctx_get_grad_std,
|
||||||
|
"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):
|
||||||
|
return {
|
||||||
|
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S"),
|
||||||
|
"epoch": context.epoch,
|
||||||
|
"iter": context.iteration,
|
||||||
|
**{m: self._metric_funcs[m](context) for m in self.metrics},
|
||||||
|
}
|
||||||
|
|
||||||
|
@only_on_rank(0)
|
||||||
|
def _add_log(self, log_data):
|
||||||
|
self.log_cache.append(log_data)
|
||||||
|
|
||||||
|
@only_on_rank(0)
|
||||||
|
def _save_log(self, epoch, iter):
|
||||||
|
log_file = self.log_dir / f"epoch_{epoch}_iter_{iter}_metric.jsonl"
|
||||||
|
|
||||||
|
with open(log_file, "w") as f:
|
||||||
|
for log in self.log_cache:
|
||||||
|
f.write(json.dumps(log) + "\n")
|
||||||
|
|
||||||
|
def on_batch_end(self, context):
|
||||||
|
if context.iteration % self.log_interval == 0:
|
||||||
|
log_data = self._get_log_data(context)
|
||||||
|
self._add_log(log_data)
|
||||||
|
|
||||||
|
if context.iteration - self.last_log_iter >= self.save_interval:
|
||||||
|
self._save_log(context.epoch, context.iteration)
|
||||||
|
self.last_log_iter = context.iteration
|
||||||
|
|
||||||
|
def on_train_end(self, context):
|
||||||
|
if context.iteration != self.last_log_iter:
|
||||||
|
self._save_log(context.epoch, context.iteration)
|
||||||
|
|
||||||
|
def on_error(self, context):
|
||||||
|
self._save_log(context.epoch, context.iteration)
|
||||||
@@ -0,0 +1,101 @@
|
|||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Callable, Optional, Self
|
||||||
|
|
||||||
|
import torch.nn as nn
|
||||||
|
from torch.optim import Optimizer
|
||||||
|
from torch.optim.lr_scheduler import LRScheduler
|
||||||
|
from torch.utils.data import DataLoader
|
||||||
|
|
||||||
|
from astrai.config.train_config import TrainConfig
|
||||||
|
from astrai.dataset import ResumableDistributedSampler
|
||||||
|
from astrai.parallel.setup import get_current_device, get_rank, get_world_size
|
||||||
|
from astrai.serialization import Checkpoint
|
||||||
|
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class TrainContext:
|
||||||
|
model: nn.Module = field(default=None)
|
||||||
|
strategy: BaseStrategy = field(default=None)
|
||||||
|
dataloader: DataLoader = field(default=None)
|
||||||
|
optimizer: Optimizer = field(default=None)
|
||||||
|
scheduler: LRScheduler = field(default=None)
|
||||||
|
checkpoint: Checkpoint = field(default=None)
|
||||||
|
|
||||||
|
epoch: int = field(default=0)
|
||||||
|
iteration: int = field(default=0)
|
||||||
|
loss: float = field(default=0.0)
|
||||||
|
|
||||||
|
world_size: int = field(default=1)
|
||||||
|
rank: int = field(default=0)
|
||||||
|
kwargs: dict = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
|
class TrainContextBuilder:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: TrainConfig,
|
||||||
|
load_extra_fn: Optional[Callable[[dict, "TrainContext"], None]] = None,
|
||||||
|
):
|
||||||
|
self.config = config
|
||||||
|
self._checkpoint: Optional[Checkpoint] = None
|
||||||
|
self._load_extra_fn = load_extra_fn
|
||||||
|
|
||||||
|
def with_checkpoint(self, checkpoint: Optional[Checkpoint]) -> Self:
|
||||||
|
self._checkpoint = checkpoint
|
||||||
|
return self
|
||||||
|
|
||||||
|
def build(self) -> TrainContext:
|
||||||
|
context = TrainContext(
|
||||||
|
model=self.config.model,
|
||||||
|
world_size=get_world_size(),
|
||||||
|
rank=get_rank(),
|
||||||
|
)
|
||||||
|
|
||||||
|
device = get_current_device()
|
||||||
|
context.model = context.model.to(device=device)
|
||||||
|
|
||||||
|
if self.config.nprocs > 1 and self.config.parallel_wrapper:
|
||||||
|
context.model = self.config.parallel_wrapper(context.model)
|
||||||
|
|
||||||
|
if self._checkpoint is not None:
|
||||||
|
context.epoch = max(self._checkpoint.epoch, self.config.start_epoch)
|
||||||
|
context.iteration = max(self._checkpoint.iteration, self.config.start_batch)
|
||||||
|
context.model.load_state_dict(self._checkpoint.state_dict)
|
||||||
|
context.checkpoint = self._checkpoint
|
||||||
|
else:
|
||||||
|
context.checkpoint = Checkpoint(
|
||||||
|
state_dict=context.model.state_dict(),
|
||||||
|
)
|
||||||
|
|
||||||
|
context.optimizer = self.config.optimizer_fn(context.model)
|
||||||
|
context.scheduler = self.config.scheduler_fn(context.optimizer)
|
||||||
|
|
||||||
|
if self._checkpoint and self._checkpoint.extra and self._load_extra_fn:
|
||||||
|
self._load_extra_fn(self._checkpoint.extra, context)
|
||||||
|
|
||||||
|
cfg = self.config
|
||||||
|
sampler_offset = context.iteration * cfg.batch_size
|
||||||
|
sampler = ResumableDistributedSampler(
|
||||||
|
data_source=cfg.dataset,
|
||||||
|
start_epoch=context.epoch,
|
||||||
|
start_iter=sampler_offset,
|
||||||
|
seed=cfg.random_seed,
|
||||||
|
)
|
||||||
|
context.dataloader = DataLoader(
|
||||||
|
cfg.dataset,
|
||||||
|
batch_size=cfg.batch_size,
|
||||||
|
sampler=sampler,
|
||||||
|
num_workers=cfg.num_workers,
|
||||||
|
pin_memory=cfg.pin_memory,
|
||||||
|
prefetch_factor=cfg.prefetch_factor,
|
||||||
|
)
|
||||||
|
|
||||||
|
context.strategy = StrategyFactory.create(
|
||||||
|
model=context.model,
|
||||||
|
train_type=self.config.strategy,
|
||||||
|
device=device,
|
||||||
|
**self.config.extra_kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
return context
|
||||||
@@ -0,0 +1,98 @@
|
|||||||
|
import logging
|
||||||
|
from typing import List, Optional
|
||||||
|
|
||||||
|
from astrai.config import TrainConfig
|
||||||
|
from astrai.parallel.setup import spawn_parallel_fn
|
||||||
|
from astrai.serialization import Checkpoint
|
||||||
|
from astrai.trainer.train_callback import (
|
||||||
|
CallbackFactory,
|
||||||
|
TrainCallback,
|
||||||
|
)
|
||||||
|
from astrai.trainer.train_context import TrainContext, TrainContextBuilder
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class Trainer:
|
||||||
|
def __init__(
|
||||||
|
self, train_config: TrainConfig, callbacks: Optional[List[TrainCallback]] = None
|
||||||
|
):
|
||||||
|
self.train_config = train_config
|
||||||
|
default_callbacks = self._get_default_callbacks()
|
||||||
|
self.callbacks = (
|
||||||
|
default_callbacks + callbacks if callbacks else default_callbacks
|
||||||
|
)
|
||||||
|
|
||||||
|
def _get_default_callbacks(self) -> List[TrainCallback]:
|
||||||
|
cfg = self.train_config
|
||||||
|
return [
|
||||||
|
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("scheduler"),
|
||||||
|
]
|
||||||
|
|
||||||
|
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):
|
||||||
|
for callback in self.callbacks:
|
||||||
|
method = getattr(callback, method_name, None)
|
||||||
|
if method:
|
||||||
|
method(context)
|
||||||
|
|
||||||
|
def train(self, checkpoint: Optional[Checkpoint] = None):
|
||||||
|
config = self.train_config
|
||||||
|
spawn_parallel_fn(
|
||||||
|
self._train_impl,
|
||||||
|
backend=config.backend,
|
||||||
|
world_size=config.nprocs,
|
||||||
|
master_addr=config.master_addr,
|
||||||
|
master_port=config.master_port,
|
||||||
|
device_type=config.device_type,
|
||||||
|
checkpoint=checkpoint,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _train_impl(self, checkpoint: Optional[Checkpoint] = None) -> Checkpoint:
|
||||||
|
context = self._build_context(checkpoint)
|
||||||
|
self._call_callbacks("on_train_begin", context)
|
||||||
|
|
||||||
|
try:
|
||||||
|
context.model.train()
|
||||||
|
# 1.epoch
|
||||||
|
for epoch in range(context.epoch, self.train_config.n_epoch):
|
||||||
|
context.epoch = epoch
|
||||||
|
self._call_callbacks("on_epoch_begin", context)
|
||||||
|
|
||||||
|
accumulation_steps = max(self.train_config.accumulation_steps, 1)
|
||||||
|
for batch in context.dataloader:
|
||||||
|
if context.iteration % accumulation_steps == 0:
|
||||||
|
# 2. step
|
||||||
|
self._call_callbacks("on_step_begin", context)
|
||||||
|
context.optimizer.step()
|
||||||
|
context.optimizer.zero_grad()
|
||||||
|
self._call_callbacks("on_step_end", context)
|
||||||
|
|
||||||
|
# 3. batch
|
||||||
|
self._call_callbacks("on_batch_begin", context)
|
||||||
|
loss = context.strategy(batch)
|
||||||
|
context.loss = loss.item()
|
||||||
|
context.iteration += 1
|
||||||
|
|
||||||
|
# to make the loss normalized by accumulation steps
|
||||||
|
stand_loss = loss / accumulation_steps
|
||||||
|
stand_loss.backward()
|
||||||
|
|
||||||
|
self._call_callbacks("on_batch_end", context)
|
||||||
|
|
||||||
|
self._call_callbacks("on_epoch_end", context)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Training failed: {str(e)}", exc_info=True)
|
||||||
|
self._call_callbacks("on_error", context)
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
self._call_callbacks("on_train_end", context)
|
||||||
-198
@@ -1,198 +0,0 @@
|
|||||||
import torch
|
|
||||||
from typing import Dict, Any
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from khaosz.model.transformer import ModelConfig, Transformer
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class BenchmarkResult:
|
|
||||||
total_tokens: int
|
|
||||||
total_time: float
|
|
||||||
tokens_per_second: float
|
|
||||||
metadata: Dict[str, Any]
|
|
||||||
|
|
||||||
|
|
||||||
class GenerationBenchmark:
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
config: ModelConfig,
|
|
||||||
device: str = "cuda",
|
|
||||||
dtype: torch.dtype = torch.float16
|
|
||||||
):
|
|
||||||
self.config = config
|
|
||||||
self.device = device
|
|
||||||
self.dtype = dtype
|
|
||||||
self.model = Transformer(config).to(device=device, dtype=dtype)
|
|
||||||
self.model.eval()
|
|
||||||
|
|
||||||
def _initialize_kv_cache(self, batch_size: int) -> list:
|
|
||||||
"""初始化KV缓存"""
|
|
||||||
config = self.config
|
|
||||||
shape = (batch_size, config.n_layer, config.m_len, config.n_kvhead, config.n_dim // config.n_head)
|
|
||||||
k_cache = torch.zeros(shape, device=self.device, dtype=self.dtype)
|
|
||||||
v_cache = torch.zeros(shape, device=self.device, dtype=self.dtype)
|
|
||||||
return (k_cache, v_cache)
|
|
||||||
|
|
||||||
def _prepare_inputs(self, batch_size: int, prompt_length: int, total_length: int):
|
|
||||||
prompt_ids = torch.randint(
|
|
||||||
low=0,
|
|
||||||
high=self.config.vocab_size,
|
|
||||||
size=(batch_size, prompt_length),
|
|
||||||
device=self.device,
|
|
||||||
dtype=torch.long
|
|
||||||
)
|
|
||||||
|
|
||||||
gen_ids = torch.randint(
|
|
||||||
low=0,
|
|
||||||
high=self.config.vocab_size,
|
|
||||||
size=(batch_size, total_length - prompt_length),
|
|
||||||
device=self.device,
|
|
||||||
dtype=torch.long
|
|
||||||
)
|
|
||||||
|
|
||||||
return prompt_ids, gen_ids
|
|
||||||
|
|
||||||
@torch.inference_mode()
|
|
||||||
def run_prefill_benchmark(
|
|
||||||
self,
|
|
||||||
batch_size: int = 1,
|
|
||||||
prompt_length: int = 512,
|
|
||||||
num_trials: int = 10,
|
|
||||||
) -> BenchmarkResult:
|
|
||||||
|
|
||||||
for _ in range(3):
|
|
||||||
prompt_ids, _ = self._prepare_inputs(batch_size, prompt_length, prompt_length)
|
|
||||||
_ = self.model(prompt_ids)
|
|
||||||
|
|
||||||
torch.cuda.synchronize()
|
|
||||||
|
|
||||||
total_time = 0.0
|
|
||||||
total_tokens = batch_size * prompt_length * num_trials
|
|
||||||
|
|
||||||
for trial in range(num_trials):
|
|
||||||
prompt_ids, _ = self._prepare_inputs(batch_size, prompt_length, prompt_length)
|
|
||||||
start_event = torch.cuda.Event(enable_timing=True)
|
|
||||||
end_event = torch.cuda.Event(enable_timing=True)
|
|
||||||
|
|
||||||
start_event.record()
|
|
||||||
_ = self.model(prompt_ids)
|
|
||||||
end_event.record()
|
|
||||||
torch.cuda.synchronize()
|
|
||||||
|
|
||||||
trial_time = start_event.elapsed_time(end_event) / 1000
|
|
||||||
total_time += trial_time
|
|
||||||
|
|
||||||
print(f"Trial {trial + 1}/{num_trials}: {prompt_length} tokens in {trial_time:.3f}s "
|
|
||||||
f"({prompt_length / trial_time:.1f} tokens/s)")
|
|
||||||
|
|
||||||
return BenchmarkResult(
|
|
||||||
total_tokens=total_tokens,
|
|
||||||
total_time=total_time,
|
|
||||||
tokens_per_second=total_tokens / total_time,
|
|
||||||
metadata={
|
|
||||||
"benchmark_type": "prefill",
|
|
||||||
"batch_size": batch_size,
|
|
||||||
"prompt_length": prompt_length,
|
|
||||||
"dtype": self.dtype,
|
|
||||||
"device": self.device,
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
@torch.inference_mode()
|
|
||||||
def run_decoding_benchmark(
|
|
||||||
self,
|
|
||||||
batch_size: int = 1,
|
|
||||||
prompt_length: int = 512,
|
|
||||||
gen_length: int = 128,
|
|
||||||
num_trials: int = 5,
|
|
||||||
) -> BenchmarkResult:
|
|
||||||
|
|
||||||
total_time = 0.0
|
|
||||||
total_tokens = batch_size * gen_length * num_trials
|
|
||||||
|
|
||||||
for trial in range(num_trials):
|
|
||||||
|
|
||||||
prompt_ids, gen_ids = self._prepare_inputs(batch_size, prompt_length, prompt_length + gen_length)
|
|
||||||
kv_cache = self._initialize_kv_cache(batch_size)
|
|
||||||
_ = self.model(prompt_ids, persistent_key_values=kv_cache, start_pos=0)
|
|
||||||
|
|
||||||
torch.cuda.synchronize()
|
|
||||||
|
|
||||||
start_event = torch.cuda.Event(enable_timing=True)
|
|
||||||
end_event = torch.cuda.Event(enable_timing=True)
|
|
||||||
|
|
||||||
start_event.record()
|
|
||||||
|
|
||||||
current_pos = prompt_length
|
|
||||||
for i in range(gen_length):
|
|
||||||
input_token = gen_ids[:, i:i+1]
|
|
||||||
_ = self.model(input_token, persistent_key_values=kv_cache, start_pos=current_pos)
|
|
||||||
current_pos += 1
|
|
||||||
|
|
||||||
end_event.record()
|
|
||||||
torch.cuda.synchronize()
|
|
||||||
|
|
||||||
trial_time = start_event.elapsed_time(end_event) / 1000
|
|
||||||
total_time += trial_time
|
|
||||||
|
|
||||||
print(f"Trial {trial + 1}/{num_trials}: {gen_length} tokens in {trial_time:.3f}s "
|
|
||||||
f"({gen_length / trial_time:.1f} tokens/s)")
|
|
||||||
|
|
||||||
|
|
||||||
return BenchmarkResult(
|
|
||||||
total_tokens=total_tokens,
|
|
||||||
total_time=total_time,
|
|
||||||
tokens_per_second=total_tokens / total_time,
|
|
||||||
metadata={
|
|
||||||
"benchmark_type": "decoding",
|
|
||||||
"batch_size": batch_size,
|
|
||||||
"prompt_length": prompt_length,
|
|
||||||
"gen_length": gen_length,
|
|
||||||
"dtype": self.dtype,
|
|
||||||
"device": self.device,
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def print_benchmark_result(result: BenchmarkResult):
|
|
||||||
"""打印基准测试结果"""
|
|
||||||
benchmark_type = result.metadata["benchmark_type"]
|
|
||||||
|
|
||||||
print(f"\n{' ' + benchmark_type.upper().replace('_', ' ') + ' Benchmark ':-^80}")
|
|
||||||
print(f"Total Tokens Processed: {result.total_tokens:,}")
|
|
||||||
print(f"Time Consumed: {result.total_time:.3f}s")
|
|
||||||
print(f"Throughput: {result.tokens_per_second:,.1f} tokens/s")
|
|
||||||
|
|
||||||
if benchmark_type == "prefill":
|
|
||||||
print(f"Batch Size: {result.metadata['batch_size']} | Prompt Length: {result.metadata['prompt_length']}")
|
|
||||||
elif benchmark_type == "decoding":
|
|
||||||
print(f"Batch Size: {result.metadata['batch_size']} | Gen Length: {result.metadata['gen_length']}")
|
|
||||||
|
|
||||||
print(f"Device: {result.metadata['device']} | Dtype: {result.metadata['dtype']}")
|
|
||||||
print("-" * 80)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
config = ModelConfig(
|
|
||||||
vocab_size=10000,
|
|
||||||
n_dim=1536,
|
|
||||||
n_head=24,
|
|
||||||
n_kvhead=4,
|
|
||||||
d_ffn=6912,
|
|
||||||
m_len=2048,
|
|
||||||
n_layer=24,
|
|
||||||
norm_eps=1e-5,
|
|
||||||
)
|
|
||||||
|
|
||||||
benchmark = GenerationBenchmark(config)
|
|
||||||
|
|
||||||
print("=" * 80)
|
|
||||||
print("Running Transformer Generation Benchmark")
|
|
||||||
print("=" * 80)
|
|
||||||
|
|
||||||
prefill_result = benchmark.run_prefill_benchmark(batch_size=4, prompt_length=512, num_trials=5)
|
|
||||||
print_benchmark_result(prefill_result)
|
|
||||||
|
|
||||||
gen_result = benchmark.run_decoding_benchmark(batch_size=4, prompt_length=512, gen_length=128, num_trials=5)
|
|
||||||
print_benchmark_result(gen_result)
|
|
||||||
|
|
||||||
@@ -0,0 +1,42 @@
|
|||||||
|
services:
|
||||||
|
server:
|
||||||
|
build: .
|
||||||
|
image: astrai:latest
|
||||||
|
ports:
|
||||||
|
- "8000:8000"
|
||||||
|
volumes:
|
||||||
|
- ./params:/app/params:ro
|
||||||
|
- ./checkpoints:/app/checkpoints
|
||||||
|
command: python -m scripts.tools.server --port 8000 --device cuda
|
||||||
|
deploy:
|
||||||
|
resources:
|
||||||
|
reservations:
|
||||||
|
devices:
|
||||||
|
- driver: nvidia
|
||||||
|
count: 1
|
||||||
|
capabilities: [gpu]
|
||||||
|
healthcheck:
|
||||||
|
test: ["CMD", "curl", "-f", "http://localhost:8000/health"]
|
||||||
|
interval: 30s
|
||||||
|
timeout: 10s
|
||||||
|
retries: 3
|
||||||
|
start_period: 60s
|
||||||
|
restart: unless-stopped
|
||||||
|
|
||||||
|
server-cpu:
|
||||||
|
profiles: [cpu]
|
||||||
|
build: .
|
||||||
|
image: astrai:latest
|
||||||
|
ports:
|
||||||
|
- "8000:8000"
|
||||||
|
volumes:
|
||||||
|
- ./params:/app/params:ro
|
||||||
|
- ./checkpoints:/app/checkpoints
|
||||||
|
command: python -m scripts.tools.server --port 8000 --device cpu
|
||||||
|
healthcheck:
|
||||||
|
test: ["CMD", "curl", "-f", "http://localhost:8000/health"]
|
||||||
|
interval: 30s
|
||||||
|
timeout: 10s
|
||||||
|
retries: 3
|
||||||
|
start_period: 120s
|
||||||
|
restart: unless-stopped
|
||||||
-101
@@ -1,101 +0,0 @@
|
|||||||
import os
|
|
||||||
import torch
|
|
||||||
import json
|
|
||||||
import torch
|
|
||||||
import argparse
|
|
||||||
|
|
||||||
from khaosz import Khaosz
|
|
||||||
from typing import List
|
|
||||||
from tqdm import tqdm
|
|
||||||
|
|
||||||
|
|
||||||
PROJECT_ROOT = os.path.dirname(os.path.abspath(__file__))
|
|
||||||
|
|
||||||
def batch_generate(
|
|
||||||
model: Khaosz,
|
|
||||||
queries: List[str],
|
|
||||||
temperature: float,
|
|
||||||
top_k: int,
|
|
||||||
top_p: float,
|
|
||||||
batch_size: int,
|
|
||||||
) -> List:
|
|
||||||
assert batch_size > 0
|
|
||||||
sorted_queries = sorted(queries, key=lambda x: len(x), reverse=True)
|
|
||||||
original_indices = {query: idx for idx, query in enumerate(queries)}
|
|
||||||
|
|
||||||
responses = [None] * len(queries)
|
|
||||||
total_batches = (len(sorted_queries) + batch_size - 1) // batch_size
|
|
||||||
|
|
||||||
for i in tqdm(range(0, total_batches * batch_size, batch_size), desc="Generating responses"):
|
|
||||||
batch_queries = sorted_queries[i: min(i + batch_size, len(queries))]
|
|
||||||
if not isinstance(batch_queries, list):
|
|
||||||
batch_queries = [batch_queries]
|
|
||||||
|
|
||||||
batch_responses = model.batch_generate(
|
|
||||||
queries=batch_queries,
|
|
||||||
temperature=temperature,
|
|
||||||
top_k=top_k,
|
|
||||||
top_p=top_p
|
|
||||||
)
|
|
||||||
|
|
||||||
for batch_query, batch_response in zip(batch_queries, batch_responses):
|
|
||||||
print(f"Q: {batch_query[:50]} \nR: {batch_response[:50]})")
|
|
||||||
|
|
||||||
for query, response in zip(batch_queries, batch_responses):
|
|
||||||
original_idx = original_indices[query]
|
|
||||||
responses[original_idx] = response
|
|
||||||
|
|
||||||
return responses
|
|
||||||
|
|
||||||
|
|
||||||
def processor(
|
|
||||||
model: Khaosz,
|
|
||||||
input_json_file: str,
|
|
||||||
output_json_file: str,
|
|
||||||
batch_size: int,
|
|
||||||
temperature: float,
|
|
||||||
top_p: float,
|
|
||||||
top_k: int,
|
|
||||||
question_key: str="question",
|
|
||||||
):
|
|
||||||
with open(input_json_file, "r", encoding='utf-8') as f:
|
|
||||||
input_dict = [json.loads(line) for line in f]
|
|
||||||
queries = [item[question_key] for item in input_dict]
|
|
||||||
|
|
||||||
output_dict = batch_generate(
|
|
||||||
model=model,
|
|
||||||
queries=queries,
|
|
||||||
temperature=temperature,
|
|
||||||
top_k=top_k,
|
|
||||||
top_p=top_p,
|
|
||||||
batch_size=batch_size
|
|
||||||
)
|
|
||||||
|
|
||||||
with open(output_json_file, "w", encoding='utf-8') as f:
|
|
||||||
json.dump(output_dict, f, indent=4, ensure_ascii=False)
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
parser = argparse.ArgumentParser(description="Run generate with a Khaosz model.")
|
|
||||||
|
|
||||||
parser.add_argument("--model_dir", type=str, required=True, help="Path to the model directory.")
|
|
||||||
parser.add_argument("--input_json_file", type=str, required=True, help="Path to the input JSONL file.")
|
|
||||||
parser.add_argument("--output_json_file", type=str, required=True, help="Path to the output JSONL file.")
|
|
||||||
parser.add_argument("--question_key", type=str, default="question", help="Key for the question in the input JSON.")
|
|
||||||
parser.add_argument("--temperature", type=float, default=0.60, help="Temperature for generating responses.")
|
|
||||||
parser.add_argument("--top_p", type=float, default=0.95, help="Top-p value for generating responses.")
|
|
||||||
parser.add_argument("--top_k", type=int, default=30, help="Top-k value for generating responses.")
|
|
||||||
parser.add_argument("--batch_size", type=int, default=1, help="Batch size for generating responses.")
|
|
||||||
|
|
||||||
args = parser.parse_args()
|
|
||||||
model = Khaosz(args.model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
|
|
||||||
processor(
|
|
||||||
model,
|
|
||||||
input_json_file=args.input_json_file,
|
|
||||||
output_json_file=args.output_json_file,
|
|
||||||
question_key=args.question_key,
|
|
||||||
batch_size=args.batch_size,
|
|
||||||
temperature=args.temperature,
|
|
||||||
top_k=args.top_k,
|
|
||||||
top_p=args.top_p
|
|
||||||
)
|
|
||||||
@@ -1,61 +0,0 @@
|
|||||||
__version__ = "1.3.1"
|
|
||||||
__author__ = "ViperEkura"
|
|
||||||
|
|
||||||
from khaosz.api import Khaosz
|
|
||||||
from khaosz.config import (
|
|
||||||
ModelConfig,
|
|
||||||
ParameterLoader,
|
|
||||||
TrainConfig,
|
|
||||||
)
|
|
||||||
from khaosz.model.transformer import Transformer
|
|
||||||
from khaosz.utils.retriever import Retriever
|
|
||||||
from khaosz.utils.splitter import (
|
|
||||||
SemanticTextSplitter,
|
|
||||||
PriorityTextSplitter
|
|
||||||
)
|
|
||||||
from khaosz.data import (
|
|
||||||
DatasetLoader,
|
|
||||||
BpeTokenizer
|
|
||||||
)
|
|
||||||
from khaosz.inference.generator import (
|
|
||||||
TextGenerator,
|
|
||||||
ChatGenerator,
|
|
||||||
StreamGenerator,
|
|
||||||
BatchGenerator,
|
|
||||||
RetrievalGenerator,
|
|
||||||
EmbeddingEncoder
|
|
||||||
)
|
|
||||||
|
|
||||||
from khaosz.trainer import (
|
|
||||||
Trainer,
|
|
||||||
StrategyFactory,
|
|
||||||
SchedulerFactory
|
|
||||||
)
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
"Khaosz",
|
|
||||||
|
|
||||||
"Transformer",
|
|
||||||
|
|
||||||
"Retriever",
|
|
||||||
"SemanticTextSplitter",
|
|
||||||
"PriorityTextSplitter",
|
|
||||||
|
|
||||||
"ModelConfig",
|
|
||||||
"ParameterLoader",
|
|
||||||
"TrainConfig",
|
|
||||||
|
|
||||||
"DatasetLoader",
|
|
||||||
"BpeTokenizer",
|
|
||||||
|
|
||||||
"TextGenerator",
|
|
||||||
"ChatGenerator",
|
|
||||||
"StreamGenerator",
|
|
||||||
"BatchGenerator",
|
|
||||||
"RetrievalGenerator",
|
|
||||||
"EmbeddingEncoder",
|
|
||||||
|
|
||||||
"Trainer",
|
|
||||||
"StrategyFactory",
|
|
||||||
"SchedulerFactory"
|
|
||||||
]
|
|
||||||
-112
@@ -1,112 +0,0 @@
|
|||||||
from torch import Tensor
|
|
||||||
from typing import List, Tuple, Generator, Union
|
|
||||||
|
|
||||||
from khaosz.inference.generator import (
|
|
||||||
TextGenerator,
|
|
||||||
ChatGenerator,
|
|
||||||
StreamGenerator,
|
|
||||||
BatchGenerator,
|
|
||||||
RetrievalGenerator,
|
|
||||||
EmbeddingEncoder
|
|
||||||
)
|
|
||||||
from khaosz.config.param_config import ParameterLoader
|
|
||||||
|
|
||||||
|
|
||||||
class Khaosz:
|
|
||||||
def __init__(self, model_dir: str):
|
|
||||||
self.parameter = ParameterLoader.load(model_dir)
|
|
||||||
|
|
||||||
def to(self, *args, **kwargs):
|
|
||||||
self.parameter.to(*args, **kwargs)
|
|
||||||
return self
|
|
||||||
|
|
||||||
def generate(
|
|
||||||
self,
|
|
||||||
query: str,
|
|
||||||
history: List[Tuple[str, str]]=None,
|
|
||||||
temperature: float=0.8,
|
|
||||||
top_k: int=50,
|
|
||||||
top_p: float=0.95,
|
|
||||||
) -> str:
|
|
||||||
generator = ChatGenerator(self.parameter)
|
|
||||||
return generator.generate(
|
|
||||||
query,
|
|
||||||
history=history,
|
|
||||||
temperature=temperature,
|
|
||||||
top_k=top_k,
|
|
||||||
top_p=top_p,
|
|
||||||
)
|
|
||||||
|
|
||||||
def batch_generate(
|
|
||||||
self,
|
|
||||||
queries: List[str],
|
|
||||||
histories: List[Tuple[str, str]]=None,
|
|
||||||
temperature: float=0.8,
|
|
||||||
top_k: int=50,
|
|
||||||
top_p: float=0.95,
|
|
||||||
) -> List[str]:
|
|
||||||
generator = BatchGenerator(self.parameter)
|
|
||||||
return generator.generate(
|
|
||||||
queries,
|
|
||||||
histories=histories,
|
|
||||||
temperature=temperature,
|
|
||||||
top_k=top_k,
|
|
||||||
top_p=top_p,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def stream_generate(
|
|
||||||
self,
|
|
||||||
query: str,
|
|
||||||
history: List[Tuple[str, str]]=None,
|
|
||||||
temperature: float=0.8,
|
|
||||||
top_k: int=50,
|
|
||||||
top_p: float=0.95,
|
|
||||||
) -> Generator[Tuple[str, List[Tuple[str, str]]], None, None]:
|
|
||||||
stream_generator = StreamGenerator(self.parameter)
|
|
||||||
return stream_generator.generate(
|
|
||||||
query,
|
|
||||||
history=history,
|
|
||||||
temperature=temperature,
|
|
||||||
top_k=top_k,
|
|
||||||
top_p=top_p,
|
|
||||||
)
|
|
||||||
|
|
||||||
def retrieve_generate(
|
|
||||||
self,
|
|
||||||
retrieved,
|
|
||||||
query: str,
|
|
||||||
history: List[Tuple[str, str]] = None,
|
|
||||||
temperature: float=0.8,
|
|
||||||
top_k: int=50,
|
|
||||||
top_p: float=0.95,
|
|
||||||
) -> str:
|
|
||||||
generator = RetrievalGenerator(self.parameter)
|
|
||||||
return generator.generate(
|
|
||||||
retrieved,
|
|
||||||
query,
|
|
||||||
history=history,
|
|
||||||
temperature=temperature,
|
|
||||||
top_k=top_k,
|
|
||||||
top_p=top_p,
|
|
||||||
)
|
|
||||||
|
|
||||||
def text_generate(
|
|
||||||
self,
|
|
||||||
query: str,
|
|
||||||
temperature: float=0.8,
|
|
||||||
top_k: int=50,
|
|
||||||
top_p: float=0.95,
|
|
||||||
) -> str:
|
|
||||||
generator = TextGenerator(self.parameter)
|
|
||||||
|
|
||||||
return generator.generate(
|
|
||||||
query,
|
|
||||||
temperature=temperature,
|
|
||||||
top_k=top_k,
|
|
||||||
top_p=top_p,
|
|
||||||
)
|
|
||||||
|
|
||||||
def encode(self, sentence: Union[str, List[str]]) -> Union[Tensor, List[Tensor]]:
|
|
||||||
encoder = EmbeddingEncoder(self.parameter)
|
|
||||||
return encoder.encode(sentence)
|
|
||||||
@@ -1,18 +0,0 @@
|
|||||||
from khaosz.config.model_config import ModelConfig
|
|
||||||
from khaosz.config.param_config import BaseModelIO, ModelParameter, Checkpoint, ParameterLoader
|
|
||||||
from khaosz.config.schedule_config import ScheduleConfig, CosineScheduleConfig, SGDRScheduleConfig
|
|
||||||
from khaosz.config.train_config import TrainConfig
|
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
"BaseModelIO",
|
|
||||||
"ModelParameter",
|
|
||||||
"Checkpoint",
|
|
||||||
"ParameterLoader",
|
|
||||||
"ModelConfig",
|
|
||||||
"TrainConfig",
|
|
||||||
|
|
||||||
"ScheduleConfig",
|
|
||||||
"CosineScheduleConfig",
|
|
||||||
"SGDRScheduleConfig",
|
|
||||||
]
|
|
||||||
@@ -1,37 +0,0 @@
|
|||||||
import json
|
|
||||||
|
|
||||||
from dataclasses import asdict, dataclass
|
|
||||||
from typing import Any, Dict, Optional, Self
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class ModelConfig:
|
|
||||||
# basic config
|
|
||||||
vocab_size: Optional[int] = None
|
|
||||||
n_dim: Optional[int] = None
|
|
||||||
n_head: Optional[int] = None
|
|
||||||
n_layer: Optional[int] = None
|
|
||||||
m_len: Optional[int] = None
|
|
||||||
norm_eps: Optional[float] = None
|
|
||||||
d_ffn: Optional[int] = None
|
|
||||||
tie_weight: Optional[bool] = None
|
|
||||||
|
|
||||||
# GQA
|
|
||||||
n_kvhead: Optional[int] = None
|
|
||||||
|
|
||||||
|
|
||||||
def load(self, config_path: str) -> Self:
|
|
||||||
with open(config_path, 'r') as f:
|
|
||||||
config: Dict[str, Any] = json.load(f)
|
|
||||||
for key, value in config.items():
|
|
||||||
if hasattr(self, key):
|
|
||||||
setattr(self, key, value)
|
|
||||||
|
|
||||||
return self
|
|
||||||
|
|
||||||
def save(self, config_path: str) -> None:
|
|
||||||
config_dict = asdict(self)
|
|
||||||
config_dict = {k: v for k, v in config_dict.items() if v is not None}
|
|
||||||
with open(config_path, 'w') as f:
|
|
||||||
json.dump(config_dict, f, indent=4)
|
|
||||||
|
|
||||||
|
|
||||||
@@ -1,244 +0,0 @@
|
|||||||
import pickle as pkl
|
|
||||||
import matplotlib.pyplot as plt
|
|
||||||
import safetensors.torch as st
|
|
||||||
import torch.nn as nn
|
|
||||||
import torch.optim as optim
|
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
|
||||||
from typing import Any, Dict, List, Optional, Self, Union
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
from khaosz.data.tokenizer import BpeTokenizer
|
|
||||||
from khaosz.config.model_config import ModelConfig
|
|
||||||
from khaosz.model.transformer import Transformer
|
|
||||||
|
|
||||||
|
|
||||||
class BaseModelIO:
|
|
||||||
"""Base class for model I/O operations."""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
model: Optional[nn.Module] = None,
|
|
||||||
tokenizer: Optional[BpeTokenizer] = None,
|
|
||||||
config: Optional[ModelConfig] = None
|
|
||||||
):
|
|
||||||
self.model = model
|
|
||||||
self.tokenizer = tokenizer or BpeTokenizer()
|
|
||||||
self.config = config or ModelConfig()
|
|
||||||
|
|
||||||
def _get_file_paths(self, directory: Union[str, Path]) -> dict[str, Path]:
|
|
||||||
"""Get standardized file paths for model components."""
|
|
||||||
dir_path = Path(directory)
|
|
||||||
return {
|
|
||||||
"model": dir_path / "model.safetensors",
|
|
||||||
"config": dir_path / "config.json",
|
|
||||||
"tokenizer": dir_path / "tokenizer.json"
|
|
||||||
}
|
|
||||||
|
|
||||||
def save_components(self, save_dir: Union[str, Path]):
|
|
||||||
"""Save core model components."""
|
|
||||||
paths = self._get_file_paths(save_dir)
|
|
||||||
paths["model"].parent.mkdir(parents=True, exist_ok=True)
|
|
||||||
|
|
||||||
if self.model is not None:
|
|
||||||
st.save_file(self.model.state_dict(), str(paths["model"]))
|
|
||||||
self.config.save(str(paths["config"]))
|
|
||||||
self.tokenizer.save(str(paths["tokenizer"]))
|
|
||||||
|
|
||||||
def load_components(self, load_dir: Union[str, Path]) -> Self:
|
|
||||||
"""Load core model components."""
|
|
||||||
paths = self._get_file_paths(load_dir)
|
|
||||||
|
|
||||||
self.config.load(str(paths["config"]))
|
|
||||||
self.tokenizer.load(str(paths["tokenizer"]))
|
|
||||||
|
|
||||||
if paths["model"].exists():
|
|
||||||
state_dict = st.load_file(str(paths["model"]))
|
|
||||||
if self.model is None:
|
|
||||||
self.model = Transformer(self.config)
|
|
||||||
self.model.load_state_dict(state_dict)
|
|
||||||
|
|
||||||
return self
|
|
||||||
|
|
||||||
def to(self, *args, **kwargs) -> Self:
|
|
||||||
"""Move model to device."""
|
|
||||||
if self.model is not None:
|
|
||||||
self.model.to(*args, **kwargs)
|
|
||||||
return self
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class ModelParameter(BaseModelIO):
|
|
||||||
"""Container for model parameters with serialization capabilities."""
|
|
||||||
|
|
||||||
model: Optional[nn.Module] = field(
|
|
||||||
default=None,
|
|
||||||
metadata={"help": "Transformer model."}
|
|
||||||
)
|
|
||||||
tokenizer: BpeTokenizer = field(
|
|
||||||
default_factory=BpeTokenizer,
|
|
||||||
metadata={"help": "Tokenizer for the model."}
|
|
||||||
)
|
|
||||||
config: ModelConfig = field(
|
|
||||||
default_factory=ModelConfig,
|
|
||||||
metadata={"help": "Transformer model configuration."}
|
|
||||||
)
|
|
||||||
|
|
||||||
def save(self, save_dir: Union[str, Path]):
|
|
||||||
self.save_components(save_dir)
|
|
||||||
|
|
||||||
def load(self, load_dir: Union[str, Path]) -> Self:
|
|
||||||
return self.load_components(load_dir)
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class Checkpoint(BaseModelIO):
|
|
||||||
"""Extended model parameters with training state."""
|
|
||||||
|
|
||||||
model: Optional[nn.Module] = field(
|
|
||||||
default=None,
|
|
||||||
metadata={"help": "Transformer model."}
|
|
||||||
)
|
|
||||||
tokenizer: BpeTokenizer = field(
|
|
||||||
default_factory=BpeTokenizer,
|
|
||||||
metadata={"help": "Tokenizer for the model."}
|
|
||||||
)
|
|
||||||
config: ModelConfig = field(
|
|
||||||
default_factory=ModelConfig,
|
|
||||||
metadata={"help": "Transformer model configuration."}
|
|
||||||
)
|
|
||||||
optimizer_state: Dict[str, Any] = field(
|
|
||||||
default=None,
|
|
||||||
metadata={"help": "Optimizer state."}
|
|
||||||
)
|
|
||||||
scheduler_state: Dict[str, Any] = field(
|
|
||||||
default=None,
|
|
||||||
metadata={"help": "Sampler state."}
|
|
||||||
)
|
|
||||||
loss_list: List[float] = field(
|
|
||||||
default_factory=list,
|
|
||||||
metadata={"help": "List of training losses."}
|
|
||||||
)
|
|
||||||
epoch: int = field(
|
|
||||||
default=0,
|
|
||||||
metadata={"help": "Current epoch."}
|
|
||||||
)
|
|
||||||
batch_iter: int = field(
|
|
||||||
default=0,
|
|
||||||
metadata={"help": "Current iteration."}
|
|
||||||
)
|
|
||||||
|
|
||||||
def _get_training_paths(self, directory: Union[str, Path]) -> dict[str, Path]:
|
|
||||||
paths = self._get_file_paths(directory)
|
|
||||||
paths.update({
|
|
||||||
"loss_list": paths["model"].parent / "loss.pkl",
|
|
||||||
"loss_plot": paths["model"].parent / "loss.png",
|
|
||||||
"optimizer_state": paths["model"].parent / "optimizer_state.pkl",
|
|
||||||
"sampler_state": paths["model"].parent / "sampler_state.pkl"
|
|
||||||
})
|
|
||||||
return paths
|
|
||||||
|
|
||||||
def save_training_state(self, save_dir: Union[str, Path]):
|
|
||||||
paths = self._get_training_paths(save_dir)
|
|
||||||
|
|
||||||
# Save loss plot
|
|
||||||
self._plot_loss(str(paths["loss_plot"]))
|
|
||||||
|
|
||||||
# Save loss list
|
|
||||||
with open(str(paths["loss_list"]), "wb") as f:
|
|
||||||
pkl.dump(self.loss_list, f)
|
|
||||||
|
|
||||||
# Save optimizer state
|
|
||||||
with open(str(paths["optimizer_state"]), "wb") as f:
|
|
||||||
pkl.dump(self.optimizer_state, f)
|
|
||||||
|
|
||||||
# Save sampler state
|
|
||||||
with open(str(paths["sampler_state"]), "wb") as f:
|
|
||||||
pkl.dump(self.scheduler_state, f)
|
|
||||||
|
|
||||||
def load_training_state(self, load_dir: Union[str, Path]) -> Self:
|
|
||||||
paths = self._get_training_paths(load_dir)
|
|
||||||
|
|
||||||
# Load loss list
|
|
||||||
if paths["loss_list"].exists():
|
|
||||||
with open(str(paths["loss_list"]), "rb") as f:
|
|
||||||
self.loss_list = pkl.load(f)
|
|
||||||
|
|
||||||
# Load optimizer state
|
|
||||||
if paths["optimizer_state"].exists():
|
|
||||||
with open(str(paths["optimizer_state"]), "rb") as f:
|
|
||||||
self.optimizer_state = pkl.load(f)
|
|
||||||
|
|
||||||
# Load sampler state
|
|
||||||
if paths["sampler_state"].exists():
|
|
||||||
with open(str(paths["sampler_state"]), "rb") as f:
|
|
||||||
self.scheduler_state = pkl.load(f)
|
|
||||||
|
|
||||||
return self
|
|
||||||
|
|
||||||
def _plot_loss(self, save_path: str):
|
|
||||||
"""Plot and save loss curve."""
|
|
||||||
if not self.loss_list:
|
|
||||||
return
|
|
||||||
|
|
||||||
batch_iter = len(self.loss_list)
|
|
||||||
|
|
||||||
plt.figure(figsize=(10, 6))
|
|
||||||
plt.plot(self.loss_list)
|
|
||||||
plt.title(f"Training Loss - Iteration {batch_iter}")
|
|
||||||
plt.xlabel("Batch")
|
|
||||||
plt.ylabel("Loss")
|
|
||||||
plt.grid(True)
|
|
||||||
plt.savefig(save_path, dpi=300, bbox_inches="tight")
|
|
||||||
plt.close()
|
|
||||||
|
|
||||||
def save(self, save_dir: Union[str, Path]):
|
|
||||||
"""Save complete checkpoint."""
|
|
||||||
self.save_components(save_dir)
|
|
||||||
self.save_training_state(save_dir)
|
|
||||||
|
|
||||||
def load(self, load_dir: Union[str, Path]) -> Self:
|
|
||||||
"""Load complete checkpoint."""
|
|
||||||
self.load_components(load_dir)
|
|
||||||
self.load_training_state(load_dir)
|
|
||||||
return self
|
|
||||||
|
|
||||||
|
|
||||||
class ParameterLoader:
|
|
||||||
"""Factory class for loading model parameters or checkpoints."""
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def load(load_dir: Union[str, Path]) -> Union[ModelParameter, Checkpoint]:
|
|
||||||
"""Load either ModelParameter or Checkpoint based on directory contents."""
|
|
||||||
load_dir = Path(load_dir)
|
|
||||||
|
|
||||||
# Check for training-specific files
|
|
||||||
loss_file = load_dir / "loss.pkl"
|
|
||||||
has_training_data = loss_file.exists()
|
|
||||||
|
|
||||||
# Create appropriate instance
|
|
||||||
if has_training_data:
|
|
||||||
checkpoint = Checkpoint()
|
|
||||||
checkpoint.load(str(load_dir))
|
|
||||||
return checkpoint
|
|
||||||
else:
|
|
||||||
params = ModelParameter()
|
|
||||||
params.load(str(load_dir))
|
|
||||||
return params
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def create_checkpoint(
|
|
||||||
model: nn.Module,
|
|
||||||
tokenizer: BpeTokenizer,
|
|
||||||
config: ModelConfig,
|
|
||||||
loss_list: Optional[list[float]] = None,
|
|
||||||
optimizer: Optional[optim.Optimizer] = None,
|
|
||||||
) -> Checkpoint:
|
|
||||||
"""Convenience method to create a training checkpoint."""
|
|
||||||
return Checkpoint(
|
|
||||||
model=model,
|
|
||||||
tokenizer=tokenizer,
|
|
||||||
config=config,
|
|
||||||
loss_list=loss_list or [],
|
|
||||||
optimizer_state=optimizer
|
|
||||||
)
|
|
||||||
@@ -1,87 +0,0 @@
|
|||||||
from typing import Any, Literal, Dict
|
|
||||||
from abc import ABC, abstractmethod
|
|
||||||
from dataclasses import dataclass, field
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class ScheduleConfig(ABC):
|
|
||||||
schedule_type: str = field(
|
|
||||||
default="cosine",
|
|
||||||
metadata={
|
|
||||||
"help": "Type of learning rate schedule.",
|
|
||||||
"choices": ["cosine", "sgdr"]
|
|
||||||
}
|
|
||||||
)
|
|
||||||
warmup_steps: int = field(
|
|
||||||
default=1000,
|
|
||||||
metadata={"help": "Number of warmup steps."}
|
|
||||||
)
|
|
||||||
min_rate: float = field(
|
|
||||||
default=0.05,
|
|
||||||
metadata={"help": "Minimum learning rate multiplier."}
|
|
||||||
)
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def get_kwargs(self) -> Dict[str, Any]:
|
|
||||||
raise NotImplementedError
|
|
||||||
|
|
||||||
def validate(self) -> None:
|
|
||||||
"""Validate configuration parameters."""
|
|
||||||
if self.warmup_steps < 0:
|
|
||||||
raise ValueError(f"warmup_steps must be non-negative, got {self.warmup_steps}")
|
|
||||||
if not 0 <= self.min_rate <= 1:
|
|
||||||
raise ValueError(f"min_rate must be between 0 and 1, got {self.min_rate}")
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class CosineScheduleConfig(ScheduleConfig):
|
|
||||||
total_steps: int = field(
|
|
||||||
default=None,
|
|
||||||
metadata={"help": "Total training steps for cosine schedule."}
|
|
||||||
)
|
|
||||||
schedule_type: Literal["cosine"] = "cosine"
|
|
||||||
|
|
||||||
def get_kwargs(self) -> Dict[str, Any]:
|
|
||||||
if self.total_steps is None:
|
|
||||||
raise ValueError("total_steps must be specified for cosine schedule")
|
|
||||||
|
|
||||||
return {
|
|
||||||
"schedule_type": self.schedule_type,
|
|
||||||
"warmup_steps": self.warmup_steps,
|
|
||||||
"lr_decay_steps": self.total_steps - self.warmup_steps,
|
|
||||||
"min_rate": self.min_rate
|
|
||||||
}
|
|
||||||
|
|
||||||
def validate(self) -> None:
|
|
||||||
super().validate()
|
|
||||||
if self.total_steps is not None and self.total_steps <= self.warmup_steps:
|
|
||||||
raise ValueError(f"total_steps ({self.total_steps}) must be greater than warmup_steps ({self.warmup_steps})")
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class SGDRScheduleConfig(ScheduleConfig):
|
|
||||||
cycle_length: int = field(
|
|
||||||
default=1000,
|
|
||||||
metadata={"help": "Length of the first cycle in steps."}
|
|
||||||
)
|
|
||||||
t_mult: int = field(
|
|
||||||
default=2,
|
|
||||||
metadata={"help": "Multiplier for cycle length growth."}
|
|
||||||
)
|
|
||||||
schedule_type: Literal["sgdr"] = "sgdr"
|
|
||||||
|
|
||||||
def get_kwargs(self) -> Dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"schedule_type": self.schedule_type,
|
|
||||||
"warmup_steps": self.warmup_steps,
|
|
||||||
"cycle_length": self.cycle_length,
|
|
||||||
"min_rate": self.min_rate,
|
|
||||||
"t_mult": self.t_mult
|
|
||||||
}
|
|
||||||
|
|
||||||
def validate(self) -> None:
|
|
||||||
super().validate()
|
|
||||||
if self.cycle_length <= 0:
|
|
||||||
raise ValueError(f"cycle_length must be positive, got {self.cycle_length}")
|
|
||||||
if self.t_mult < 1:
|
|
||||||
raise ValueError(f"t_mult must be >= 1, got {self.t_mult}")
|
|
||||||
@@ -1,72 +0,0 @@
|
|||||||
from dataclasses import dataclass, field
|
|
||||||
from typing import Optional, TYPE_CHECKING
|
|
||||||
from torch.utils.data import Dataset
|
|
||||||
from torch.optim import Optimizer
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from khaosz.trainer.strategy import BaseStrategy
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class TrainConfig:
|
|
||||||
|
|
||||||
strategy: "BaseStrategy" = field(
|
|
||||||
default=None,
|
|
||||||
metadata={"help": "Training strategy."}
|
|
||||||
)
|
|
||||||
dataset: Dataset = field(
|
|
||||||
default=None,
|
|
||||||
metadata={"help": "Dataset for training."}
|
|
||||||
)
|
|
||||||
optimizer: Optimizer = field(
|
|
||||||
default=None,
|
|
||||||
metadata={"help": "Optimizer for training."}
|
|
||||||
)
|
|
||||||
checkpoint_dir: str = field(
|
|
||||||
default="./checkpoint",
|
|
||||||
metadata={"help": "Checkpoint directory."}
|
|
||||||
)
|
|
||||||
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."}
|
|
||||||
)
|
|
||||||
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."}
|
|
||||||
)
|
|
||||||
checkpoint_interval: int = field(
|
|
||||||
default=5000,
|
|
||||||
metadata={"help": "Number of iterations between checkpoints."}
|
|
||||||
)
|
|
||||||
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."}
|
|
||||||
)
|
|
||||||
random_seed: int = field(
|
|
||||||
default=3407,
|
|
||||||
metadata={"help": "Random seed."}
|
|
||||||
)
|
|
||||||
num_workers: int = field(
|
|
||||||
default=0,
|
|
||||||
metadata={"help": "Number of workers for dataloader."}
|
|
||||||
)
|
|
||||||
prefetch_factor: Optional[int] = field(
|
|
||||||
default=None,
|
|
||||||
metadata={"help": "Prefetch factor for dataloader."}
|
|
||||||
)
|
|
||||||
pin_memory: bool = field(
|
|
||||||
default=False,
|
|
||||||
metadata={"help": "Pin memory for dataloader."}
|
|
||||||
)
|
|
||||||
@@ -1,26 +0,0 @@
|
|||||||
from khaosz.data.data_util import (
|
|
||||||
BaseDataset,
|
|
||||||
SeqDataset,
|
|
||||||
DpoDataset,
|
|
||||||
SftDataset,
|
|
||||||
PpoDataset,
|
|
||||||
MutiSegmentFetcher,
|
|
||||||
ResumeableRandomSampler,
|
|
||||||
DatasetLoader,
|
|
||||||
load_pkl_files,
|
|
||||||
)
|
|
||||||
|
|
||||||
from khaosz.data.tokenizer import BpeTokenizer
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
"BaseDataset",
|
|
||||||
"SeqDataset",
|
|
||||||
"DpoDataset",
|
|
||||||
"SftDataset",
|
|
||||||
"PpoDataset",
|
|
||||||
"MutiSegmentFetcher",
|
|
||||||
"ResumeableRandomSampler",
|
|
||||||
"DatasetLoader",
|
|
||||||
"load_pkl_files",
|
|
||||||
"BpeTokenizer"
|
|
||||||
]
|
|
||||||
@@ -1,256 +0,0 @@
|
|||||||
import torch
|
|
||||||
import bisect
|
|
||||||
import pickle as pkl
|
|
||||||
from abc import ABC, abstractmethod
|
|
||||||
from torch import Tensor
|
|
||||||
from torch.utils.data import Dataset, Sampler
|
|
||||||
from typing import Callable, List, Dict, Literal, Optional, Union
|
|
||||||
|
|
||||||
MutiSeg = Dict[str, List[Tensor]]
|
|
||||||
Seg = Dict[str, Tensor]
|
|
||||||
|
|
||||||
def load_pkl_files(paths: List[str]):
|
|
||||||
segments: MutiSeg = {}
|
|
||||||
total_samples = 0
|
|
||||||
|
|
||||||
for path in paths:
|
|
||||||
with open(path, "rb") as f:
|
|
||||||
pkl_file: Seg = pkl.load(f)
|
|
||||||
for key, value in pkl_file.items():
|
|
||||||
if key not in segments:
|
|
||||||
segments[key] = []
|
|
||||||
segments[key].append(value)
|
|
||||||
first_key = list(pkl_file.keys())[0]
|
|
||||||
total_samples += pkl_file[first_key].numel()
|
|
||||||
|
|
||||||
return segments, total_samples
|
|
||||||
|
|
||||||
|
|
||||||
class BaseSegmentFetcher:
|
|
||||||
def __init__(self, segments: List[Tensor]):
|
|
||||||
self.segments = segments
|
|
||||||
self.cum_lengths = []
|
|
||||||
total = 0
|
|
||||||
for seg in segments:
|
|
||||||
total += len(seg)
|
|
||||||
self.cum_lengths.append(total)
|
|
||||||
self.total_length = total if segments else 0
|
|
||||||
|
|
||||||
def fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
|
||||||
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)
|
|
||||||
|
|
||||||
# fix the range index bug
|
|
||||||
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 MutiSegmentFetcher:
|
|
||||||
def __init__(self, muti_segments: MutiSeg):
|
|
||||||
self.muti_keys = list(muti_segments.keys())
|
|
||||||
self.muti_fetchers = {
|
|
||||||
key: BaseSegmentFetcher(segments)
|
|
||||||
for key, segments in muti_segments.items()
|
|
||||||
}
|
|
||||||
|
|
||||||
def key_fetch(self, begin_idx: int, end_idx: int, keys: Union[str, List[str]]) -> Union[Tensor, Seg]:
|
|
||||||
fetch_dict = {}
|
|
||||||
keys = [keys] if isinstance(keys, str) else keys
|
|
||||||
|
|
||||||
for key in keys:
|
|
||||||
fetcher = self.muti_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) -> Union[Tensor, Seg]:
|
|
||||||
return self.key_fetch(begin_idx, end_idx, self.muti_keys)
|
|
||||||
|
|
||||||
|
|
||||||
class BaseDataset(Dataset, ABC):
|
|
||||||
def __init__(self, window_size: int, stride: int):
|
|
||||||
super().__init__()
|
|
||||||
self.segments: MutiSeg = {}
|
|
||||||
self.window_size = window_size
|
|
||||||
self.stride = stride
|
|
||||||
self.total_samples = None
|
|
||||||
|
|
||||||
def save(self, save_path: str):
|
|
||||||
keys = list(self.segments.keys())
|
|
||||||
if not keys:
|
|
||||||
return
|
|
||||||
|
|
||||||
first_item = self.segments[keys[0]]
|
|
||||||
segment_size = len(first_item)
|
|
||||||
|
|
||||||
for i in range(segment_size):
|
|
||||||
formated_segment = {key: self.segments[key][i] for key in keys}
|
|
||||||
pkl.dump(formated_segment, open(f"{save_path}_{i}.pkl", "wb"))
|
|
||||||
|
|
||||||
def load(self, load_path: Union[str, List[str]]):
|
|
||||||
paths = [load_path] if isinstance(load_path, str) else load_path
|
|
||||||
self.segments, self.total_samples = load_pkl_files(paths)
|
|
||||||
self.fetcher = MutiSegmentFetcher(self.segments)
|
|
||||||
|
|
||||||
def get_index(self, index: int) -> int:
|
|
||||||
begin_idx = min(index * self.stride, self.total_samples - self.window_size - 1)
|
|
||||||
end_idx = begin_idx + self.window_size
|
|
||||||
|
|
||||||
return begin_idx, end_idx
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
|
||||||
raise NotImplementedError
|
|
||||||
|
|
||||||
def __len__(self) -> int:
|
|
||||||
assert self.total_samples is not None
|
|
||||||
if self.total_samples <= self.window_size:
|
|
||||||
return 0
|
|
||||||
return self.total_samples // self.stride + 1
|
|
||||||
|
|
||||||
|
|
||||||
class SeqDataset(BaseDataset):
|
|
||||||
def __init__(self, window_size: int, stride: int):
|
|
||||||
super().__init__(window_size, stride)
|
|
||||||
self.fetcher = MutiSegmentFetcher(self.segments)
|
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
|
||||||
return self.fetcher.key_fetch(begin_idx, end_idx, "sequence")
|
|
||||||
|
|
||||||
def __getitem__(self, index):
|
|
||||||
# fix the range index bug
|
|
||||||
begin_idx, end_idx = self.get_index(index)
|
|
||||||
|
|
||||||
x = self._fetch_data(begin_idx, end_idx).to(dtype=torch.long)
|
|
||||||
y = self._fetch_data(begin_idx + 1, end_idx + 1).to(dtype=torch.long)
|
|
||||||
|
|
||||||
return {"input_ids": x, "target_ids": y}
|
|
||||||
|
|
||||||
|
|
||||||
class SftDataset(BaseDataset):
|
|
||||||
def __init__(self, window_size: int, stride: int):
|
|
||||||
super().__init__(window_size, stride)
|
|
||||||
self.fetcher = MutiSegmentFetcher(self.segments)
|
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
|
||||||
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
|
||||||
|
|
||||||
def __getitem__(self, index):
|
|
||||||
begin_idx, end_idx = self.get_index(index)
|
|
||||||
|
|
||||||
x = self._fetch_data(begin_idx, end_idx, "sequence").to(dtype=torch.long)
|
|
||||||
y = self._fetch_data(begin_idx + 1, end_idx + 1, "sequence").to(dtype=torch.long)
|
|
||||||
loss_mask = self._fetch_data(begin_idx + 1, end_idx + 1, "loss_mask").to(dtype=torch.bool)
|
|
||||||
|
|
||||||
return {"input_ids": x, "target_ids": y, "loss_mask": loss_mask}
|
|
||||||
|
|
||||||
|
|
||||||
class DpoDataset(BaseDataset):
|
|
||||||
def __init__(self, window_size: int, stride: int):
|
|
||||||
super().__init__(window_size, stride)
|
|
||||||
self.fetcher = MutiSegmentFetcher(self.segments)
|
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
|
||||||
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
|
||||||
|
|
||||||
def __getitem__(self, index: int):
|
|
||||||
begin_idx, end_idx = self.get_index(index)
|
|
||||||
|
|
||||||
chosen = self._fetch_data(begin_idx, end_idx, "chosen").to(dtype=torch.long)
|
|
||||||
rejected = self._fetch_data(begin_idx, end_idx, "rejected").to(dtype=torch.long)
|
|
||||||
chosen_mask = self._fetch_data(begin_idx, end_idx, "chosen_mask").to(dtype=torch.bool)
|
|
||||||
rejected_mask = self._fetch_data(begin_idx, end_idx, "rejected_mask").to(dtype=torch.bool)
|
|
||||||
|
|
||||||
return {"chosen": chosen, "rejected": rejected, "chosen_mask": chosen_mask, "rejected_mask": rejected_mask}
|
|
||||||
|
|
||||||
|
|
||||||
class PpoDataset(BaseDataset):
|
|
||||||
def __init__(self, window_size: int, stride: int):
|
|
||||||
super().__init__(window_size, stride)
|
|
||||||
self.fetcher = MutiSegmentFetcher(self.segments)
|
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
|
||||||
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
|
||||||
|
|
||||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
|
||||||
begin_idx, end_idx = self.get_index(index)
|
|
||||||
|
|
||||||
input_ids = self._fetch_data(begin_idx, end_idx, "input_ids"),
|
|
||||||
actions = self._fetch_data(begin_idx, end_idx, "actions"),
|
|
||||||
logprobs = self._fetch_data(begin_idx, end_idx, "logprobs"),
|
|
||||||
rewards = self._fetch_data(begin_idx, end_idx, "rewards")
|
|
||||||
|
|
||||||
return {"input_ids": input_ids, "actions": actions, "logprobs": logprobs, "rewards": rewards}
|
|
||||||
|
|
||||||
|
|
||||||
class DatasetLoader:
|
|
||||||
@staticmethod
|
|
||||||
def load(
|
|
||||||
train_type: Literal["seq", "sft", "dpo"],
|
|
||||||
load_path: Union[str, List[str]],
|
|
||||||
window_size: int,
|
|
||||||
stride: Optional[int] = None,
|
|
||||||
**kwargs
|
|
||||||
) -> BaseDataset:
|
|
||||||
if stride is None:
|
|
||||||
stride = window_size
|
|
||||||
|
|
||||||
dataset_router: Dict[str, Callable[[int], BaseDataset]] = {
|
|
||||||
"seq": lambda window_size: SeqDataset(window_size, stride),
|
|
||||||
"sft": lambda window_size: SftDataset(window_size, stride),
|
|
||||||
"dpo": lambda window_size: DpoDataset(window_size, stride),
|
|
||||||
}
|
|
||||||
dataset = dataset_router[train_type](window_size)
|
|
||||||
dataset.load(load_path)
|
|
||||||
|
|
||||||
return dataset
|
|
||||||
|
|
||||||
|
|
||||||
class ResumeableRandomSampler(Sampler[int]):
|
|
||||||
def __init__(self, data_source, start_epoch=0, start_iter=0, seed=42):
|
|
||||||
self.num_samples = len(data_source)
|
|
||||||
self.epoch = start_epoch
|
|
||||||
self.iter = start_iter
|
|
||||||
|
|
||||||
generator = torch.Generator()
|
|
||||||
generator.manual_seed(seed)
|
|
||||||
|
|
||||||
# consume previous epochs
|
|
||||||
for _ in range(start_epoch):
|
|
||||||
torch.randperm(self.num_samples, generator=generator)
|
|
||||||
|
|
||||||
self.generator = generator
|
|
||||||
self._indices = None
|
|
||||||
|
|
||||||
def _get_indices(self):
|
|
||||||
current_epoch_indices = torch.randperm(self.num_samples, generator=self.generator).tolist()
|
|
||||||
self._indices = current_epoch_indices[self.iter % self.num_samples:]
|
|
||||||
|
|
||||||
def __iter__(self):
|
|
||||||
if self._indices is None:
|
|
||||||
self._get_indices()
|
|
||||||
|
|
||||||
for i in self._indices:
|
|
||||||
self.iter += 1
|
|
||||||
yield i
|
|
||||||
|
|
||||||
self.epoch += 1
|
|
||||||
self._indices = None
|
|
||||||
|
|
||||||
def __len__(self):
|
|
||||||
if self._indices is None:
|
|
||||||
self._get_indices()
|
|
||||||
return len(self._indices)
|
|
||||||
@@ -1,110 +0,0 @@
|
|||||||
from tokenizers import Tokenizer, Encoding
|
|
||||||
from tokenizers import decoders, processors, normalizers, pre_tokenizers
|
|
||||||
from tokenizers.models import BPE
|
|
||||||
from tokenizers.trainers import BpeTrainer
|
|
||||||
from typing import List, Union
|
|
||||||
|
|
||||||
|
|
||||||
class BpeTokenizer:
|
|
||||||
def __init__(self, path=None):
|
|
||||||
self._control_tokens = ["<bos>", "<eos>", "<pad>"]
|
|
||||||
self._special_tokens = ["<|im_start|>", "<|im_end|>"]
|
|
||||||
model = BPE()
|
|
||||||
tokenizer = Tokenizer(model)
|
|
||||||
tokenizer.normalizer = normalizers.Sequence([
|
|
||||||
normalizers.NFC()
|
|
||||||
])
|
|
||||||
tokenizer.pre_tokenizer = pre_tokenizers.Sequence([
|
|
||||||
pre_tokenizers.Punctuation(behavior="isolated"),
|
|
||||||
pre_tokenizers.Metaspace(prepend_scheme="never"),
|
|
||||||
pre_tokenizers.Split(pattern=r"(\d+|[a-zA-Z]+|(?:'s|'t|'re|'ve|'m|'ll|'d))", behavior="isolated"),
|
|
||||||
pre_tokenizers.ByteLevel(add_prefix_space=False, use_regex=False)
|
|
||||||
])
|
|
||||||
tokenizer.decoder = decoders.Sequence([
|
|
||||||
decoders.ByteLevel(),
|
|
||||||
decoders.Metaspace(prepend_scheme="never")
|
|
||||||
])
|
|
||||||
tokenizer.post_processor = processors.Sequence([
|
|
||||||
processors.ByteLevel(trim_offsets=False)
|
|
||||||
])
|
|
||||||
self._tokenizer = tokenizer
|
|
||||||
|
|
||||||
if path is not None:
|
|
||||||
self._tokenizer = Tokenizer.from_file(path)
|
|
||||||
|
|
||||||
def _prepare_trainer(self, vocab_size: int, min_freq: int, reserved_token_size: int) -> tuple:
|
|
||||||
assert reserved_token_size > len(self._special_tokens)
|
|
||||||
reserved_tokens = [f"<|rsv{i:02d}|>" for i in range(reserved_token_size - len(self._special_tokens))]
|
|
||||||
detail_vocab_size = vocab_size - (len(reserved_tokens) + len(self._special_tokens))
|
|
||||||
|
|
||||||
alphabet = pre_tokenizers.ByteLevel.alphabet()
|
|
||||||
min_size = len(alphabet) + len(self._control_tokens)
|
|
||||||
assert detail_vocab_size > min_size
|
|
||||||
|
|
||||||
trainer = BpeTrainer(
|
|
||||||
vocab_size=detail_vocab_size,
|
|
||||||
min_frequency=min_freq,
|
|
||||||
limit_alphabet=detail_vocab_size // 4,
|
|
||||||
max_token_length=18,
|
|
||||||
special_tokens=self._control_tokens,
|
|
||||||
show_progress=True,
|
|
||||||
initial_alphabet=alphabet,
|
|
||||||
)
|
|
||||||
|
|
||||||
return trainer, detail_vocab_size, reserved_tokens
|
|
||||||
|
|
||||||
def train(self, files, vocab_size, min_freq, reserved_token_size=100):
|
|
||||||
trainer, _, reserved_tokens = self._prepare_trainer(
|
|
||||||
vocab_size=vocab_size,
|
|
||||||
min_freq=min_freq,
|
|
||||||
reserved_token_size=reserved_token_size
|
|
||||||
)
|
|
||||||
self._tokenizer.train(files=files, trainer=trainer)
|
|
||||||
self._tokenizer.add_special_tokens(self._special_tokens + reserved_tokens)
|
|
||||||
|
|
||||||
def train_from_iterator(self, iterator, vocab_size, min_freq, reserved_token_size=100):
|
|
||||||
trainer, _, reserved_tokens = self._prepare_trainer(
|
|
||||||
vocab_size=vocab_size,
|
|
||||||
min_freq=min_freq,
|
|
||||||
reserved_token_size=reserved_token_size
|
|
||||||
)
|
|
||||||
self._tokenizer.train_from_iterator(iterator=iterator, trainer=trainer)
|
|
||||||
self._tokenizer.add_special_tokens(self._special_tokens + reserved_tokens)
|
|
||||||
|
|
||||||
def save(self, path):
|
|
||||||
self._tokenizer.save(path)
|
|
||||||
|
|
||||||
def load(self, path):
|
|
||||||
self._tokenizer = Tokenizer.from_file(path)
|
|
||||||
|
|
||||||
def encode(self, tokens: Union[str, List[str]], out_ids: bool=True, add_special_tokens: bool=False) -> List:
|
|
||||||
if isinstance(tokens, str):
|
|
||||||
encoded: Encoding = self._tokenizer.encode(tokens, add_special_tokens=add_special_tokens)
|
|
||||||
return encoded.ids if out_ids else encoded.tokens
|
|
||||||
elif isinstance(tokens, list):
|
|
||||||
encoded_list: List[Encoding] = self._tokenizer.encode_batch(tokens, add_special_tokens=add_special_tokens)
|
|
||||||
return [encoded.ids if out_ids else encoded.tokens for encoded in encoded_list]
|
|
||||||
|
|
||||||
def decode(self, tokens: List[int], skip_special_tokens: bool=True) -> str:
|
|
||||||
return self._tokenizer.decode(tokens, skip_special_tokens=skip_special_tokens)
|
|
||||||
|
|
||||||
def __len__(self) -> int:
|
|
||||||
return self._tokenizer.get_vocab_size()
|
|
||||||
|
|
||||||
@property
|
|
||||||
def stop_ids(self) -> List[int]:
|
|
||||||
stop_token = self._control_tokens + self._special_tokens
|
|
||||||
stop_ids = [self._tokenizer.token_to_id(token) for token in stop_token]
|
|
||||||
return stop_ids
|
|
||||||
|
|
||||||
@property
|
|
||||||
def bos_id(self) -> int:
|
|
||||||
return self._tokenizer.token_to_id("<bos>")
|
|
||||||
|
|
||||||
@property
|
|
||||||
def eos_id(self) -> int:
|
|
||||||
return self._tokenizer.token_to_id("<eos>")
|
|
||||||
|
|
||||||
@property
|
|
||||||
def pad_id(self) -> int:
|
|
||||||
return self._tokenizer.token_to_id("<pad>")
|
|
||||||
@@ -1,240 +0,0 @@
|
|||||||
import torch
|
|
||||||
from torch import Tensor
|
|
||||||
from typing import Any, Callable, List, Tuple, Union, Optional, Self
|
|
||||||
from khaosz.config import ModelParameter, ModelConfig
|
|
||||||
|
|
||||||
|
|
||||||
def apply_sampling_strategies(
|
|
||||||
logits: Tensor,
|
|
||||||
temperature: float,
|
|
||||||
top_k: int,
|
|
||||||
top_p: float,
|
|
||||||
filter_value: float = -float("inf")
|
|
||||||
) -> Tensor:
|
|
||||||
"""
|
|
||||||
Apply sampling strategies to the logits tensor.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
logits (Tensor): The logits tensor.
|
|
||||||
temperature (float): The temperature parameter.
|
|
||||||
top_k (int): The top-k parameter.
|
|
||||||
top_p (float): The top-p parameter.
|
|
||||||
filter_value (float, optional): The filter value. Defaults to -float("inf").
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tensor: The sampled logits tensor.
|
|
||||||
|
|
||||||
"""
|
|
||||||
|
|
||||||
if temperature != 1.0:
|
|
||||||
logits = logits / temperature
|
|
||||||
|
|
||||||
if top_k > 0:
|
|
||||||
top_k = min(top_k, logits.size(-1))
|
|
||||||
indices_to_remove = logits < torch.topk(logits, top_k, dim=-1)[0][..., -1, None]
|
|
||||||
logits[indices_to_remove] = filter_value
|
|
||||||
|
|
||||||
if top_p < 1.0:
|
|
||||||
sorted_logits, sorted_indices = torch.sort(logits, descending=True, dim=-1)
|
|
||||||
cumulative_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
|
|
||||||
|
|
||||||
sorted_indices_to_remove = cumulative_probs > top_p
|
|
||||||
sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
|
|
||||||
sorted_indices_to_remove[..., 0] = 0
|
|
||||||
|
|
||||||
indices_to_remove = torch.zeros_like(logits, dtype=torch.bool)
|
|
||||||
indices_to_remove.scatter_(
|
|
||||||
dim=1,
|
|
||||||
index=sorted_indices,
|
|
||||||
src=sorted_indices_to_remove
|
|
||||||
)
|
|
||||||
|
|
||||||
logits[indices_to_remove] = filter_value
|
|
||||||
|
|
||||||
return logits
|
|
||||||
|
|
||||||
|
|
||||||
class GeneratorCore:
|
|
||||||
def __init__(self, parameter: ModelParameter):
|
|
||||||
self.model = parameter.model
|
|
||||||
self.tokenizer = parameter.tokenizer
|
|
||||||
self.config = parameter.config
|
|
||||||
|
|
||||||
def generate_iterator(
|
|
||||||
self,
|
|
||||||
input_ids: Tensor,
|
|
||||||
temperature: float,
|
|
||||||
top_k: int,
|
|
||||||
top_p: float,
|
|
||||||
attn_mask: Optional[Tensor] = None,
|
|
||||||
kv_caches: Optional[List[Tuple[Tensor, Tensor]]] = None,
|
|
||||||
start_pos: int = 0
|
|
||||||
)-> Tuple[Tensor, int]:
|
|
||||||
|
|
||||||
with torch.inference_mode():
|
|
||||||
outputs = self.model(input_ids, attn_mask, kv_caches, start_pos)
|
|
||||||
logits = outputs["logits"][:, -1, :]
|
|
||||||
cache_increase = input_ids.size(-1)
|
|
||||||
|
|
||||||
logits = apply_sampling_strategies(logits, temperature, top_k, top_p)
|
|
||||||
probs = torch.softmax(logits, dim=-1)
|
|
||||||
next_token_id = torch.multinomial(probs, num_samples=1)
|
|
||||||
|
|
||||||
return next_token_id, cache_increase
|
|
||||||
|
|
||||||
def to(self, *args, **kargs) -> Self:
|
|
||||||
self.model.to(*args, **kargs)
|
|
||||||
return self
|
|
||||||
|
|
||||||
def generate_loop(
|
|
||||||
self,
|
|
||||||
input_ids: Tensor,
|
|
||||||
ids: List[int],
|
|
||||||
temperature: float,
|
|
||||||
top_k: int,
|
|
||||||
top_p: float,
|
|
||||||
attn_mask: Optional[Tensor] = None,
|
|
||||||
kv_caches: Optional[List[Tuple[Tensor, Tensor]]] = None,
|
|
||||||
start_pos: int = 0,
|
|
||||||
callback: Optional[Callable[..., Any]] = None
|
|
||||||
) -> List[int]:
|
|
||||||
cur_cache_pos = start_pos
|
|
||||||
|
|
||||||
for _ in range(len(ids), self.config.m_len):
|
|
||||||
next_token_id, cache_increase = self.generate_iterator(
|
|
||||||
input_ids, temperature, top_k, top_p, attn_mask, kv_caches, cur_cache_pos)
|
|
||||||
|
|
||||||
input_ids = next_token_id
|
|
||||||
ids.append(next_token_id.item())
|
|
||||||
cur_cache_pos += cache_increase
|
|
||||||
|
|
||||||
if callback:
|
|
||||||
callback(next_token_id.item(), ids.copy())
|
|
||||||
|
|
||||||
if next_token_id.item() in self.tokenizer.stop_ids:
|
|
||||||
break
|
|
||||||
|
|
||||||
return ids
|
|
||||||
|
|
||||||
|
|
||||||
class EmbeddingEncoderCore:
|
|
||||||
def __init__(self, parameter: ModelParameter):
|
|
||||||
self.model = parameter.model
|
|
||||||
self.tokenizer = parameter.tokenizer
|
|
||||||
self.config = parameter.config
|
|
||||||
|
|
||||||
def encode(self, sentence: Union[str, List[str]]) -> Union[Tensor, List[Tensor]]:
|
|
||||||
with_batch = isinstance(sentence, list)
|
|
||||||
ids = self.tokenizer.encode(sentence)
|
|
||||||
batch_ids = ids if with_batch else [ids]
|
|
||||||
max_model_len = self.config.m_len
|
|
||||||
|
|
||||||
all_fragments = []
|
|
||||||
fragment_origin_idx = []
|
|
||||||
|
|
||||||
for i, seq in enumerate(batch_ids):
|
|
||||||
if len(seq) > max_model_len:
|
|
||||||
fragments = [seq[j:j+max_model_len] for j in range(0, len(seq), max_model_len)]
|
|
||||||
all_fragments.extend(fragments)
|
|
||||||
fragment_origin_idx.extend([i] * len(fragments))
|
|
||||||
else:
|
|
||||||
all_fragments.append(seq)
|
|
||||||
fragment_origin_idx.append(i)
|
|
||||||
|
|
||||||
#if empty fragments
|
|
||||||
if not all_fragments or not ids:
|
|
||||||
return [] if with_batch else torch.tensor([])
|
|
||||||
|
|
||||||
device = next(self.model.parameters()).device
|
|
||||||
max_len = min(max(len(seq) for seq in all_fragments), max_model_len)
|
|
||||||
|
|
||||||
padded_ids = []
|
|
||||||
masks = []
|
|
||||||
for seq in all_fragments:
|
|
||||||
pad_len = max_len - len(seq)
|
|
||||||
padded_seq = seq + [self.tokenizer.pad_id] * pad_len
|
|
||||||
mask = [token_id != self.tokenizer.pad_id for token_id in padded_seq]
|
|
||||||
padded_ids.append(padded_seq)
|
|
||||||
masks.append(mask)
|
|
||||||
|
|
||||||
input_tensor = torch.tensor(padded_ids, device=device, dtype=torch.long)
|
|
||||||
seq_mask = torch.tensor(masks, device=device, dtype=torch.bool)
|
|
||||||
|
|
||||||
with torch.inference_mode():
|
|
||||||
outputs = self.model(input_tensor, seq_mask)["hidden_states"]
|
|
||||||
# [num_fragments, seq_len, hidden_size]
|
|
||||||
fragment_embs = torch.mul(outputs, seq_mask.unsqueeze(-1))
|
|
||||||
|
|
||||||
sentence_embs: List[Tensor] = []
|
|
||||||
for i in range(len(batch_ids)):
|
|
||||||
indices = [idx for idx, orig_idx in enumerate(fragment_origin_idx) if orig_idx == i]
|
|
||||||
if indices is not None:
|
|
||||||
sum_frags = torch.sum(fragment_embs[indices, :, :], dim=1) # [frags, hidden_size]
|
|
||||||
length = torch.sum(seq_mask[indices, :], dim=1).unsqueeze(1) # [frags, 1]
|
|
||||||
emb = torch.sum(sum_frags / length, dim=0) # [frags, hidden_size]
|
|
||||||
sentence_embs.append(emb.flatten())
|
|
||||||
|
|
||||||
if with_batch:
|
|
||||||
return [emb.flatten() for emb in sentence_embs]
|
|
||||||
else:
|
|
||||||
return sentence_embs[0].flatten()
|
|
||||||
|
|
||||||
def to(self, *args, **kargs) -> Self:
|
|
||||||
self.model.to(*args, **kargs)
|
|
||||||
return self
|
|
||||||
|
|
||||||
|
|
||||||
class KVCacheManager:
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
config: ModelConfig,
|
|
||||||
batch_size: int,
|
|
||||||
device: torch.device = "cuda",
|
|
||||||
dtype: torch.dtype = torch.bfloat16
|
|
||||||
):
|
|
||||||
self.batch_size = batch_size
|
|
||||||
self.device = device
|
|
||||||
self.dtype = dtype
|
|
||||||
self.num_layers = config.n_layer
|
|
||||||
self.max_len = config.m_len
|
|
||||||
self.num_heads = config.n_kvhead
|
|
||||||
self.head_dim = config.n_dim //config.n_head
|
|
||||||
|
|
||||||
self._kv_cache: Tuple[Tensor, Tensor] = None
|
|
||||||
self._seq_mask: Tensor = None
|
|
||||||
self._initialize()
|
|
||||||
|
|
||||||
def _initialize(self):
|
|
||||||
k_cache = torch.zeros(
|
|
||||||
(self.batch_size, self.num_layers, self.max_len, self.num_heads, self.head_dim),
|
|
||||||
device=self.device, dtype=self.dtype
|
|
||||||
)
|
|
||||||
v_cache = torch.zeros(
|
|
||||||
(self.batch_size, self.num_layers, self.max_len, self.num_heads, self.head_dim),
|
|
||||||
device=self.device, dtype=self.dtype
|
|
||||||
)
|
|
||||||
self._kv_cache = (k_cache, v_cache)
|
|
||||||
self._seq_mask = torch.ones((self.batch_size, self.max_len), device=self.device, dtype=torch.bool)
|
|
||||||
|
|
||||||
def update(self, active_mask: Tensor):
|
|
||||||
k_cache, v_cache = self._kv_cache
|
|
||||||
self._kv_cache = (k_cache[active_mask], v_cache[active_mask])
|
|
||||||
self._seq_mask = self._seq_mask[active_mask]
|
|
||||||
|
|
||||||
def reset(self, full_reset=False):
|
|
||||||
if full_reset:
|
|
||||||
self._kv_cache = None
|
|
||||||
self._seq_mask = None
|
|
||||||
else:
|
|
||||||
self._initialize()
|
|
||||||
|
|
||||||
def set_seq_mask(self, input_ids: Tensor, pad_id: int):
|
|
||||||
batch_size, seq_len = input_ids.shape
|
|
||||||
bool_mask = (input_ids != pad_id)
|
|
||||||
self._seq_mask[: batch_size, : seq_len] = bool_mask
|
|
||||||
|
|
||||||
def get_kvcache(self) -> Tuple[Tensor, Tensor]:
|
|
||||||
return self._kv_cache
|
|
||||||
|
|
||||||
def get_seq_mask(self) -> Tensor:
|
|
||||||
return self._seq_mask
|
|
||||||
@@ -1,296 +0,0 @@
|
|||||||
import torch
|
|
||||||
from torch import Tensor
|
|
||||||
from typing import List, Tuple, Union, Optional, Generator
|
|
||||||
from khaosz.inference.core import GeneratorCore, EmbeddingEncoderCore, KVCacheManager
|
|
||||||
from khaosz.config.param_config import ModelParameter
|
|
||||||
|
|
||||||
|
|
||||||
def build_prompt(
|
|
||||||
query: str,
|
|
||||||
init_prompt: Optional[str] = None,
|
|
||||||
history: Optional[List[Tuple[str, str]]] = None
|
|
||||||
) -> str:
|
|
||||||
"""
|
|
||||||
Build prompt in ChatML format for query and history
|
|
||||||
|
|
||||||
Args:
|
|
||||||
query(str): query string
|
|
||||||
history(Optional[List[Tuple[str, str]]]): history list of query and response
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
str: prompt string in ChatML format
|
|
||||||
|
|
||||||
"""
|
|
||||||
prompt = f"<|im_start|>system\n{init_prompt}<|im_end|>\n" if init_prompt else ""
|
|
||||||
|
|
||||||
# (convert tuple format to ChatML)
|
|
||||||
if history:
|
|
||||||
for user_msg, assistant_msg in history:
|
|
||||||
prompt += f"<|im_start|>user\n{user_msg}<|im_end|>\n"
|
|
||||||
prompt += f"<|im_start|>assistant\n{assistant_msg}<|im_end|>\n"
|
|
||||||
|
|
||||||
prompt += f"<|im_start|>user\n{query}<|im_end|>\n"
|
|
||||||
prompt += "<|im_start|>assistant\n"
|
|
||||||
|
|
||||||
return prompt
|
|
||||||
|
|
||||||
def pad_sequence(ids_list: List[List[int]], max_ids_len: int, pad_id: int) -> List[List[int]]:
|
|
||||||
"""
|
|
||||||
Pad a list of sequences to a fixed length.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
ids_list (List[List[int]]): A list of sequences.
|
|
||||||
max_ids_len (int): The maximum length of sequences.
|
|
||||||
pad_id (int): The id to pad sequences.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
List[List[int]]: A list of padded sequences.
|
|
||||||
|
|
||||||
"""
|
|
||||||
new_ids_list = []
|
|
||||||
for ids in ids_list:
|
|
||||||
pad_len = max_ids_len - len(ids)
|
|
||||||
padded_seq = [pad_id] * pad_len + ids
|
|
||||||
new_ids_list.append(padded_seq)
|
|
||||||
|
|
||||||
return new_ids_list
|
|
||||||
|
|
||||||
|
|
||||||
class TextGenerator(GeneratorCore):
|
|
||||||
def __init__(self, parameter: ModelParameter):
|
|
||||||
super().__init__(parameter)
|
|
||||||
|
|
||||||
def generate(
|
|
||||||
self,
|
|
||||||
query: str,
|
|
||||||
temperature: float,
|
|
||||||
top_k: int,
|
|
||||||
top_p: float,
|
|
||||||
) -> str:
|
|
||||||
assert temperature >= 0.0
|
|
||||||
assert top_k >= 0
|
|
||||||
assert top_p >= 0.0 and top_p <= 1.0
|
|
||||||
|
|
||||||
device = next(self.model.parameters()).device
|
|
||||||
cache_manager = KVCacheManager(self.config, 1, device=device)
|
|
||||||
|
|
||||||
ids = self.tokenizer.encode(query)
|
|
||||||
input_ids = torch.tensor([ids], device=device, dtype=torch.long)
|
|
||||||
|
|
||||||
start_cache_pos = len(ids)
|
|
||||||
cur_cache_pos = 0
|
|
||||||
self.model.eval()
|
|
||||||
kv_caches = cache_manager.get_kvcache()
|
|
||||||
|
|
||||||
ids = self.generate_loop(
|
|
||||||
input_ids, ids, temperature, top_k, top_p,
|
|
||||||
kv_caches=kv_caches,
|
|
||||||
start_pos=cur_cache_pos
|
|
||||||
)
|
|
||||||
|
|
||||||
response = self.tokenizer.decode(ids[start_cache_pos:])
|
|
||||||
|
|
||||||
return response
|
|
||||||
|
|
||||||
|
|
||||||
class ChatGenerator(GeneratorCore):
|
|
||||||
def __init__(self, parameter: ModelParameter):
|
|
||||||
super().__init__(parameter)
|
|
||||||
|
|
||||||
def generate(
|
|
||||||
self,
|
|
||||||
query: str,
|
|
||||||
history: List[Tuple[str, str]],
|
|
||||||
temperature: float,
|
|
||||||
top_k: int,
|
|
||||||
top_p: float,
|
|
||||||
) -> str:
|
|
||||||
|
|
||||||
assert temperature >= 0.0
|
|
||||||
assert top_k >= 0
|
|
||||||
assert top_p >= 0.0 and top_p <= 1.0
|
|
||||||
|
|
||||||
if history is None:
|
|
||||||
history = []
|
|
||||||
|
|
||||||
device = next(self.model.parameters()).device
|
|
||||||
cache_manager = KVCacheManager(self.config, 1, device=device)
|
|
||||||
|
|
||||||
ids = self.tokenizer.encode(build_prompt(query, history))
|
|
||||||
input_ids = torch.tensor([ids], device=device, dtype=torch.long)
|
|
||||||
|
|
||||||
start_cache_pos = len(ids)
|
|
||||||
cur_cache_pos = 0
|
|
||||||
self.model.eval()
|
|
||||||
kv_caches = cache_manager.get_kvcache()
|
|
||||||
|
|
||||||
ids = self.generate_loop(
|
|
||||||
input_ids, ids, temperature, top_k, top_p,
|
|
||||||
kv_caches=kv_caches,
|
|
||||||
start_pos=cur_cache_pos
|
|
||||||
)
|
|
||||||
|
|
||||||
response = self.tokenizer.decode(ids[start_cache_pos:])
|
|
||||||
|
|
||||||
return response
|
|
||||||
|
|
||||||
|
|
||||||
class StreamGenerator(GeneratorCore):
|
|
||||||
def __init__(self, parameter: ModelParameter):
|
|
||||||
super().__init__(parameter)
|
|
||||||
|
|
||||||
def generate(
|
|
||||||
self,
|
|
||||||
query: str,
|
|
||||||
history: List[Tuple[str, str]],
|
|
||||||
temperature: float,
|
|
||||||
top_k: int,
|
|
||||||
top_p: float,
|
|
||||||
) -> Generator[Tuple[str, List[Tuple[str, str]]], None, None]:
|
|
||||||
|
|
||||||
assert temperature >= 0.0
|
|
||||||
assert top_k >= 0
|
|
||||||
assert top_p >= 0.0 and top_p <= 1.0
|
|
||||||
|
|
||||||
if history is None:
|
|
||||||
history = []
|
|
||||||
|
|
||||||
device = next(self.model.parameters()).device
|
|
||||||
cache_manager = KVCacheManager(self.config, 1, device=device)
|
|
||||||
|
|
||||||
ids = self.tokenizer.encode(build_prompt(query, history))
|
|
||||||
input_ids = torch.tensor([ids], device=device, dtype=torch.long)
|
|
||||||
cpy_history = history.copy()
|
|
||||||
|
|
||||||
start_cache_pos = len(ids)
|
|
||||||
cur_cache_pos = 0
|
|
||||||
self.model.eval()
|
|
||||||
kv_caches = cache_manager.get_kvcache()
|
|
||||||
|
|
||||||
for _ in range(len(ids), self.config.m_len):
|
|
||||||
next_token_id, cache_increase = self.generate_iterator(
|
|
||||||
input_ids, temperature, top_k, top_p, kv_caches=kv_caches, start_pos=cur_cache_pos)
|
|
||||||
|
|
||||||
input_ids = next_token_id
|
|
||||||
ids.append(next_token_id.item())
|
|
||||||
cur_cache_pos += cache_increase
|
|
||||||
|
|
||||||
response = self.tokenizer.decode(ids[start_cache_pos:])
|
|
||||||
yield response, cpy_history + [(query, response)]
|
|
||||||
|
|
||||||
if next_token_id.item() in self.tokenizer.stop_ids:
|
|
||||||
yield response + "\n", cpy_history + [(query, response)]
|
|
||||||
break
|
|
||||||
|
|
||||||
|
|
||||||
class BatchGenerator(GeneratorCore):
|
|
||||||
def __init__(self, parameter: ModelParameter):
|
|
||||||
super().__init__(parameter)
|
|
||||||
|
|
||||||
def generate(
|
|
||||||
self,
|
|
||||||
queries: List[str],
|
|
||||||
histories: List[List[Tuple[str, str]]],
|
|
||||||
temperature: float,
|
|
||||||
top_k: int,
|
|
||||||
top_p: float
|
|
||||||
) -> List[str]:
|
|
||||||
|
|
||||||
assert temperature >= 0.0
|
|
||||||
assert top_k >= 0
|
|
||||||
assert top_p >= 0.0 and top_p <= 1.0
|
|
||||||
|
|
||||||
batch_size = len(queries)
|
|
||||||
if histories is None:
|
|
||||||
histories = [[] for _ in range(batch_size)]
|
|
||||||
|
|
||||||
prompts = [build_prompt(query, history) for query, history in zip(queries, histories)]
|
|
||||||
ids_list = [self.tokenizer.encode(prompt) for prompt in prompts]
|
|
||||||
max_ids_len = max(len(ids) for ids in ids_list)
|
|
||||||
ids_list = pad_sequence(ids_list, max_ids_len, self.tokenizer.pad_id)
|
|
||||||
|
|
||||||
device = next(self.model.parameters()).device
|
|
||||||
cache_manager = KVCacheManager(self.config, batch_size, device=device)
|
|
||||||
|
|
||||||
input_tensor = torch.tensor(ids_list, device=device, dtype=torch.long)
|
|
||||||
cache_manager.set_seq_mask(input_tensor, self.tokenizer.pad_id)
|
|
||||||
activate_task_mask = [True] * batch_size
|
|
||||||
|
|
||||||
start_cache_pos = max_ids_len
|
|
||||||
cur_cache_pos = 0
|
|
||||||
|
|
||||||
while max_ids_len < self.config.m_len and sum(activate_task_mask) != 0:
|
|
||||||
kv_caches = cache_manager.get_kvcache()
|
|
||||||
attn_mask =cache_manager.get_seq_mask()
|
|
||||||
|
|
||||||
next_token_id, cache_increase = self.generate_iterator(
|
|
||||||
input_tensor, temperature, top_k, top_p, attn_mask=attn_mask, kv_caches=kv_caches, start_pos=cur_cache_pos)
|
|
||||||
|
|
||||||
cur_cache_pos += cache_increase
|
|
||||||
active_mask = []
|
|
||||||
c_ids = 0
|
|
||||||
|
|
||||||
for i in range(batch_size):
|
|
||||||
if activate_task_mask[i]:
|
|
||||||
token = next_token_id[c_ids, :].item()
|
|
||||||
ids_list[i].append(token)
|
|
||||||
c_ids += 1
|
|
||||||
|
|
||||||
is_active = not token in self.tokenizer.stop_ids
|
|
||||||
activate_task_mask[i] = is_active
|
|
||||||
active_mask.append(is_active)
|
|
||||||
|
|
||||||
active_mask = torch.tensor(active_mask, device=device, dtype=torch.bool)
|
|
||||||
cache_manager.update(active_mask)
|
|
||||||
input_tensor = next_token_id[active_mask, :]
|
|
||||||
|
|
||||||
max_ids_len += 1
|
|
||||||
|
|
||||||
|
|
||||||
responses = [str()] * batch_size
|
|
||||||
for i in range(batch_size):
|
|
||||||
responses[i] = self.tokenizer.decode(ids_list[i][start_cache_pos:])
|
|
||||||
histories[i].append((queries[i], responses[i]))
|
|
||||||
|
|
||||||
return responses
|
|
||||||
|
|
||||||
|
|
||||||
class RetrievalGenerator(GeneratorCore):
|
|
||||||
def __init__(self, retriever_parameter: ModelParameter):
|
|
||||||
super().__init__(retriever_parameter)
|
|
||||||
|
|
||||||
def generate(
|
|
||||||
self,
|
|
||||||
retrieved: List[str],
|
|
||||||
query: str,
|
|
||||||
history: List[Tuple[str, str]],
|
|
||||||
temperature: float,
|
|
||||||
top_k: int,
|
|
||||||
top_p: float,
|
|
||||||
) -> str:
|
|
||||||
assert temperature >= 0.0
|
|
||||||
assert top_k >= 0
|
|
||||||
assert top_p >= 0.0 and top_p <= 1.0
|
|
||||||
|
|
||||||
if history is None:
|
|
||||||
history = []
|
|
||||||
|
|
||||||
retrieved = "\n".join([f"{idx + 1}. {key}" for idx, key in enumerate(retrieved)]) if retrieved else ""
|
|
||||||
retrieved_query = f"{retrieved}\n\n{query}" if retrieved else query
|
|
||||||
parameter = ModelParameter(self.model, self.tokenizer, self.config)
|
|
||||||
|
|
||||||
return ChatGenerator(parameter).generate(
|
|
||||||
retrieved_query,
|
|
||||||
history,
|
|
||||||
temperature=temperature,
|
|
||||||
top_k=top_k,
|
|
||||||
top_p=top_p,
|
|
||||||
)
|
|
||||||
|
|
||||||
class EmbeddingEncoder(EmbeddingEncoderCore):
|
|
||||||
def __init__(self, parameter: ModelParameter):
|
|
||||||
super().__init__(parameter)
|
|
||||||
|
|
||||||
def encode(self, sentence: Union[str, List[str]]) -> Union[Tensor, List[Tensor]]:
|
|
||||||
return super().encode(sentence)
|
|
||||||
|
|
||||||
@@ -1,17 +0,0 @@
|
|||||||
from khaosz.model.module import (
|
|
||||||
Linear,
|
|
||||||
RMSNorm,
|
|
||||||
MLP,
|
|
||||||
GQA,
|
|
||||||
DecoderBlock,
|
|
||||||
)
|
|
||||||
from khaosz.model.transformer import Transformer
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
"Linear",
|
|
||||||
"RMSNorm",
|
|
||||||
"MLP",
|
|
||||||
"GQA",
|
|
||||||
"DecoderBlock",
|
|
||||||
"Transformer"
|
|
||||||
]
|
|
||||||
@@ -1,250 +0,0 @@
|
|||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
import torch.nn.functional as F
|
|
||||||
|
|
||||||
from torch import Tensor
|
|
||||||
from typing import Optional, Tuple
|
|
||||||
|
|
||||||
|
|
||||||
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
|
||||||
"""
|
|
||||||
Repeat k times along the dimension for attention heads.
|
|
||||||
Args:
|
|
||||||
x (Tensor): The input tensor.
|
|
||||||
n_rep (int): The number of repetitions.
|
|
||||||
Returns:
|
|
||||||
Tensor: The repeated tensor.
|
|
||||||
"""
|
|
||||||
|
|
||||||
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,
|
|
||||||
) -> Tuple[Tensor, Tensor]:
|
|
||||||
"""
|
|
||||||
Get the rotary embedding for the given dimension and maximum length.
|
|
||||||
Args:
|
|
||||||
dim (int): The dimension of the input.
|
|
||||||
max_len (int): The maximum length of the input.
|
|
||||||
base (float, optional): The base for the frequency. Defaults to 10000.
|
|
||||||
Returns:
|
|
||||||
Tensor: The rotary embedding tensor.
|
|
||||||
"""
|
|
||||||
|
|
||||||
theta = base ** (-torch.arange(0, dim, 2, dtype=torch.float64) / dim)
|
|
||||||
t = torch.arange(0, max_len, dtype=torch.float64)
|
|
||||||
freqs = torch.outer(t, theta)
|
|
||||||
|
|
||||||
return torch.cos(freqs).float(), torch.sin(freqs).float()
|
|
||||||
|
|
||||||
def apply_rotary_emb(x: torch.Tensor, rotary_emb: Tuple[Tensor, Tensor]) -> Tensor:
|
|
||||||
"""
|
|
||||||
Apply rotary embedding to the input tensor using cos/sin form.
|
|
||||||
Args:
|
|
||||||
x (Tensor): The input tensor (shape [..., seq_len, dim]).
|
|
||||||
rotary_emb (Tuple[Tensor, Tensor]): The rotary embedding (shape [seq_len, dim//2]).
|
|
||||||
Returns:
|
|
||||||
Tensor: The output tensor (rotated, same shape as input).
|
|
||||||
"""
|
|
||||||
|
|
||||||
dtype = x.dtype
|
|
||||||
cos, sin = rotary_emb
|
|
||||||
|
|
||||||
cos = cos.unsqueeze(0).unsqueeze(2) # [1, seq_len, 1, dim//2]
|
|
||||||
sin = sin.unsqueeze(0).unsqueeze(2) # [1, seq_len, 1, dim//2]
|
|
||||||
|
|
||||||
x_real = x[..., 0::2] # [batch, seq_len, dim//2]
|
|
||||||
x_imag = x[..., 1::2] # [batch, seq_len, dim//2]
|
|
||||||
|
|
||||||
x_real_rot = x_real * cos - x_imag * sin
|
|
||||||
x_imag_rot = x_real * sin + x_imag * cos
|
|
||||||
|
|
||||||
x_out = torch.stack([x_real_rot, x_imag_rot], dim=-1) # [batch, seq_len, dim//2, 2]
|
|
||||||
x_out = x_out.view(*x_out.shape[:-2], -1) # [batch, seq_len, dim]
|
|
||||||
|
|
||||||
return x_out.to(dtype)
|
|
||||||
|
|
||||||
|
|
||||||
class RotaryEmbedding(nn.Module):
|
|
||||||
def __init__(self, dim: int, max_len: int, base: int=10000):
|
|
||||||
super().__init__()
|
|
||||||
self.dim = dim
|
|
||||||
self.max_len = max_len
|
|
||||||
self.base = base
|
|
||||||
self.max_len_cached = None
|
|
||||||
self._set_rotary_buffer(self.max_len)
|
|
||||||
|
|
||||||
def _set_rotary_buffer(self, max_len: int):
|
|
||||||
cos_cached, sin_cached = get_rotary_emb(self.dim, max_len, self.base)
|
|
||||||
self.register_buffer("cos_cached", cos_cached, persistent=False)
|
|
||||||
self.register_buffer("sin_cached", sin_cached, persistent=False)
|
|
||||||
self.max_len_cached = max_len
|
|
||||||
|
|
||||||
def forward(self, x: Tensor, start_pos: int=0) -> Tuple[Tensor, Tensor]:
|
|
||||||
seq_len = x.size(1)
|
|
||||||
|
|
||||||
if self.max_len_cached < seq_len + start_pos:
|
|
||||||
self._set_rotary_buffer(seq_len)
|
|
||||||
|
|
||||||
cos = self.cos_cached[start_pos : start_pos + seq_len]
|
|
||||||
sin = self.sin_cached[start_pos : start_pos + seq_len]
|
|
||||||
|
|
||||||
return (cos, sin)
|
|
||||||
|
|
||||||
|
|
||||||
class Linear(nn.Module):
|
|
||||||
def __init__(self, in_dim: int, out_dim: int, bias: bool = False, weight_param=None, bias_param=None):
|
|
||||||
super().__init__()
|
|
||||||
weight_param = torch.empty((out_dim, in_dim)) if weight_param is None else weight_param
|
|
||||||
bias_param = torch.zeros(out_dim) if bias_param is None else bias_param
|
|
||||||
|
|
||||||
self.weight = nn.Parameter(weight_param)
|
|
||||||
self.bias = nn.Parameter(bias_param) 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, n_dim, norm_eps):
|
|
||||||
super().__init__()
|
|
||||||
self.weight = nn.Parameter(torch.ones(n_dim))
|
|
||||||
self.norm_eps = norm_eps
|
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
|
||||||
dtype = x.dtype
|
|
||||||
x = x.float()
|
|
||||||
mean_square = torch.mean(torch.pow(x, 2), dim=-1, keepdim=True)
|
|
||||||
norm = x * torch.rsqrt(mean_square + self.norm_eps)
|
|
||||||
norm = norm.to(dtype)
|
|
||||||
out = norm * self.weight
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
class MLP(nn.Module):
|
|
||||||
def __init__(self, n_dim: int, d_ffn: int):
|
|
||||||
super().__init__()
|
|
||||||
self.up = Linear(n_dim, d_ffn)
|
|
||||||
self.gate = Linear(n_dim, d_ffn)
|
|
||||||
self.down = Linear(d_ffn, n_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,
|
|
||||||
n_dim: int,
|
|
||||||
n_head: int,
|
|
||||||
n_kvhead: int,
|
|
||||||
layer_id: int
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
assert n_dim % n_head == 0
|
|
||||||
assert n_head % n_kvhead == 0
|
|
||||||
|
|
||||||
self.head_dim = n_dim // n_head
|
|
||||||
self.layer_id = layer_id
|
|
||||||
self.n_dim = n_dim
|
|
||||||
self.n_heads = n_head
|
|
||||||
self.n_kvheads = n_kvhead
|
|
||||||
self.n_rep = n_head // n_kvhead
|
|
||||||
|
|
||||||
self.q_proj = Linear(n_dim, n_head * self.head_dim)
|
|
||||||
self.k_proj = Linear(n_dim, n_kvhead * self.head_dim)
|
|
||||||
self.v_proj = Linear(n_dim, n_kvhead * self.head_dim)
|
|
||||||
self.o_proj = Linear(n_dim, n_dim)
|
|
||||||
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
x: Tensor,
|
|
||||||
rotary_emb: Tuple[Tensor, Tensor],
|
|
||||||
mask: Tensor = None,
|
|
||||||
kv_cache: Optional[Tuple[Tensor, Tensor]] = None,
|
|
||||||
start_pos: int = 0
|
|
||||||
) -> Tensor:
|
|
||||||
bsz, seq_len, _ = x.size()
|
|
||||||
# x(bsz, seq_len, n_heads * head_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_kvheads)
|
|
||||||
v = self._split_heads(self.v_proj(x), self.n_kvheads)
|
|
||||||
q, k = apply_rotary_emb(q, rotary_emb), apply_rotary_emb(k, rotary_emb)
|
|
||||||
|
|
||||||
if kv_cache is not None:
|
|
||||||
k_cache, v_cache = kv_cache
|
|
||||||
|
|
||||||
# copy to cache
|
|
||||||
k_cache[:bsz, self.layer_id, start_pos:start_pos + seq_len] = k
|
|
||||||
v_cache[:bsz, self.layer_id, start_pos:start_pos + seq_len] = v
|
|
||||||
|
|
||||||
# get cache
|
|
||||||
k = k_cache[:bsz, self.layer_id, :start_pos + seq_len]
|
|
||||||
v = v_cache[:bsz, self.layer_id, :start_pos + seq_len]
|
|
||||||
|
|
||||||
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)
|
|
||||||
|
|
||||||
# (bsz, seq_len, n_heads, head_dim) -> (bsz, n_heads, seq_len, head_dim)
|
|
||||||
q, k, v = q.permute(0, 2, 1, 3), k.permute(0, 2, 1, 3), v.permute(0, 2, 1, 3)
|
|
||||||
sdqa_out = F.scaled_dot_product_attention(q, k, v, mask, is_causal=(mask == None)).permute(0, 2, 1, 3)
|
|
||||||
out = self.o_proj(sdqa_out.contiguous().view(bsz, seq_len, -1))
|
|
||||||
|
|
||||||
return out
|
|
||||||
|
|
||||||
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
|
|
||||||
|
|
||||||
|
|
||||||
class DecoderBlock(nn.Module):
|
|
||||||
def __init__(self, n_dim, n_head, d_ffn, n_kvhead, norm_eps, layer_id):
|
|
||||||
super().__init__()
|
|
||||||
self.attention = GQA(n_dim, n_head, n_kvhead, layer_id)
|
|
||||||
self.norm_attn = RMSNorm(n_dim, norm_eps)
|
|
||||||
self.ffn = MLP(n_dim, d_ffn)
|
|
||||||
self.norm_ffn = RMSNorm(n_dim, norm_eps)
|
|
||||||
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
x: Tensor,
|
|
||||||
rotary_emb: Tuple[Tensor, Tensor],
|
|
||||||
attention_mask: Optional[Tensor] = None,
|
|
||||||
kv_cache: Optional[Tuple[Tensor, Tensor]] = None,
|
|
||||||
start_pos: int = 0
|
|
||||||
) -> Tensor:
|
|
||||||
# attention
|
|
||||||
attn_output = self.attention(
|
|
||||||
self.norm_attn(x),
|
|
||||||
rotary_emb,
|
|
||||||
attention_mask,
|
|
||||||
kv_cache,
|
|
||||||
start_pos
|
|
||||||
)
|
|
||||||
x = attn_output + x
|
|
||||||
|
|
||||||
# feed forward
|
|
||||||
x = self.ffn(self.norm_ffn(x)) + x
|
|
||||||
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
class Embedding(nn.Module):
|
|
||||||
def __init__(self, vocab_size: int, embedding_dim: int, weight_param=None):
|
|
||||||
super().__init__()
|
|
||||||
weight_param = torch.empty((vocab_size, embedding_dim)) if weight_param is None else weight_param
|
|
||||||
self.weight = nn.Parameter(weight_param)
|
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
|
||||||
return F.embedding(x, self.weight)
|
|
||||||
@@ -1,132 +0,0 @@
|
|||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
|
|
||||||
from torch import Tensor
|
|
||||||
from typing import Any, Mapping, Optional, Tuple
|
|
||||||
from khaosz.config.model_config import ModelConfig
|
|
||||||
from khaosz.model.module import Embedding, DecoderBlock, Linear, RMSNorm, RotaryEmbedding
|
|
||||||
|
|
||||||
|
|
||||||
def process_attention_mask(
|
|
||||||
seq_mask: Tensor,
|
|
||||||
input_tensor: Tensor,
|
|
||||||
start_pos: int = 0,
|
|
||||||
is_causal: bool = False,
|
|
||||||
) -> Tensor:
|
|
||||||
"""
|
|
||||||
Create attention mask for GQA
|
|
||||||
Args:
|
|
||||||
seq_mask (Tensor): A tensor indicating whether each position is valid or not.
|
|
||||||
input_tensor (Tensor): The input tensor.
|
|
||||||
start_pos (int): The starting position of the sequence.
|
|
||||||
is_causal (bool): Whether the attention is causal or not.
|
|
||||||
Returns:
|
|
||||||
Tensor: The attention mask tensor.
|
|
||||||
"""
|
|
||||||
device = input_tensor.device
|
|
||||||
dtype = input_tensor.dtype
|
|
||||||
seq_len = input_tensor.size(1)
|
|
||||||
|
|
||||||
if seq_mask is None:
|
|
||||||
if start_pos != 0:
|
|
||||||
# for single prompt chat
|
|
||||||
seq_mask = torch.ones((1, seq_len), dtype=torch.bool, device=device)
|
|
||||||
else:
|
|
||||||
return None
|
|
||||||
|
|
||||||
if seq_mask.dim() > 2:
|
|
||||||
# shape (bsz, seq_len) or (bsz,n_heads, seq_len, seq_len + start_pos)
|
|
||||||
# if ndim > 2, it's 4D tensor
|
|
||||||
return seq_mask
|
|
||||||
|
|
||||||
batch_size = seq_mask.size(0)
|
|
||||||
seq_mask = seq_mask[:, :start_pos + seq_len].to(device=device, dtype=torch.bool)
|
|
||||||
# (bsz, start_pos + seq_len)
|
|
||||||
expanded_mask = seq_mask.unsqueeze(1).expand(batch_size, seq_len, start_pos + seq_len)
|
|
||||||
# (bsz, seq_len, start_pos + seq_len)
|
|
||||||
|
|
||||||
if is_causal:
|
|
||||||
expanded_mask = torch.tril(expanded_mask, diagonal=start_pos)
|
|
||||||
|
|
||||||
attention_mask = torch.zeros_like(expanded_mask, dtype=dtype, device=device)
|
|
||||||
attention_mask = attention_mask.masked_fill_(~expanded_mask, -torch.finfo(dtype).max / 2).unsqueeze(1)
|
|
||||||
# (bsz, 1, seq_len, seq_len + start_pos)
|
|
||||||
|
|
||||||
return attention_mask
|
|
||||||
|
|
||||||
|
|
||||||
class Transformer(nn.Module):
|
|
||||||
def __init__(self, config: ModelConfig):
|
|
||||||
super().__init__()
|
|
||||||
self.config = config
|
|
||||||
self.rotary_embeding = RotaryEmbedding(config.n_dim // config.n_head, config.m_len)
|
|
||||||
self.embed_tokens = Embedding(config.vocab_size, config.n_dim)
|
|
||||||
|
|
||||||
self.layers = nn.ModuleList([
|
|
||||||
DecoderBlock(config.n_dim, config.n_head, config.d_ffn, config.n_kvhead, config.norm_eps, layer_id)
|
|
||||||
for layer_id in range(config.n_layer)
|
|
||||||
])
|
|
||||||
|
|
||||||
self.norm = RMSNorm(config.n_dim, config.norm_eps)
|
|
||||||
self.lm_head = Linear(config.n_dim, config.vocab_size)
|
|
||||||
|
|
||||||
if self.config.tie_weight == True:
|
|
||||||
self.lm_head.weight = self.embed_tokens.weight
|
|
||||||
|
|
||||||
self._init_parameters()
|
|
||||||
|
|
||||||
def load_state_dict(self, state_dict: Mapping[str, Any], strict=True, assign=False):
|
|
||||||
lm_head_key = 'lm_head.weight'
|
|
||||||
embed_key = 'embed_tokens.weight'
|
|
||||||
|
|
||||||
if self.config.tie_weight == True:
|
|
||||||
# same tensor
|
|
||||||
state_dict[lm_head_key] = state_dict[embed_key]
|
|
||||||
else:
|
|
||||||
# use clone to avoid sharing the same tensor
|
|
||||||
state_dict[lm_head_key] = torch.clone(state_dict[embed_key])
|
|
||||||
|
|
||||||
return super().load_state_dict(state_dict, strict, assign)
|
|
||||||
|
|
||||||
def state_dict(self, destination=None, prefix='', keep_vars=False):
|
|
||||||
state_dict = super().state_dict(destination=destination, prefix=prefix, keep_vars=keep_vars)
|
|
||||||
|
|
||||||
if self.config.tie_weight == True:
|
|
||||||
lm_head_key = prefix + 'lm_head.weight'
|
|
||||||
if lm_head_key in state_dict:
|
|
||||||
del state_dict[lm_head_key]
|
|
||||||
|
|
||||||
return state_dict
|
|
||||||
|
|
||||||
def _init_parameters(self):
|
|
||||||
for param in self.parameters():
|
|
||||||
if param.dim() > 1:
|
|
||||||
nn.init.normal_(param, mean=0.0, std=0.006)
|
|
||||||
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
input_ids: Tensor,
|
|
||||||
input_mask: Optional[Tensor]=None,
|
|
||||||
persistent_key_values: Optional[Tuple[Tensor, Tensor]]=None,
|
|
||||||
start_pos: int = 0
|
|
||||||
) -> Tensor:
|
|
||||||
assert input_ids.ndim == 2
|
|
||||||
|
|
||||||
x = self.embed_tokens(input_ids)
|
|
||||||
rotary_emb = self.rotary_embeding(x, start_pos)
|
|
||||||
|
|
||||||
attn_mask = process_attention_mask(
|
|
||||||
input_mask, x, start_pos, is_causal=True
|
|
||||||
)
|
|
||||||
|
|
||||||
for layer in self.layers:
|
|
||||||
x = layer(x, rotary_emb, attn_mask, persistent_key_values, start_pos)
|
|
||||||
|
|
||||||
hidden_states = self.norm(x)
|
|
||||||
logits = self.lm_head(hidden_states)
|
|
||||||
|
|
||||||
return {
|
|
||||||
"logits": logits,
|
|
||||||
"hidden_states": hidden_states
|
|
||||||
}
|
|
||||||
|
|
||||||
@@ -1,29 +0,0 @@
|
|||||||
from khaosz.trainer.trainer import Trainer
|
|
||||||
from khaosz.trainer.strategy import StrategyFactory
|
|
||||||
from khaosz.trainer.schedule import SchedulerFactory
|
|
||||||
|
|
||||||
from khaosz.trainer.train_callback import (
|
|
||||||
TrainCallback,
|
|
||||||
ProgressBarCallback,
|
|
||||||
CheckpointCallback,
|
|
||||||
TrainCallback,
|
|
||||||
SchedulerCallback,
|
|
||||||
StepMonitorCallback
|
|
||||||
)
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
# trainer
|
|
||||||
"Trainer",
|
|
||||||
|
|
||||||
# factory
|
|
||||||
"StrategyFactory",
|
|
||||||
"SchedulerFactory",
|
|
||||||
|
|
||||||
# callback
|
|
||||||
"TrainCallback",
|
|
||||||
"ProgressBarCallback",
|
|
||||||
"CheckpointCallback",
|
|
||||||
"TrainCallback",
|
|
||||||
"SchedulerCallback",
|
|
||||||
"StepMonitorCallback"
|
|
||||||
]
|
|
||||||
@@ -1,164 +0,0 @@
|
|||||||
import math
|
|
||||||
from abc import abstractmethod, ABC
|
|
||||||
from typing import Any, Dict, List
|
|
||||||
from torch.optim.lr_scheduler import LRScheduler
|
|
||||||
from khaosz.config.schedule_config import ScheduleConfig
|
|
||||||
|
|
||||||
|
|
||||||
class BaseScheduler(LRScheduler, ABC):
|
|
||||||
"""
|
|
||||||
Base scheduler class for all other schedulers.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, optimizer, last_epoch: int = -1):
|
|
||||||
super().__init__(optimizer, last_epoch)
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def get_lr(self) -> List[float]:
|
|
||||||
raise NotImplementedError
|
|
||||||
|
|
||||||
def state_dict(self) -> Dict[str, Any]:
|
|
||||||
return super().state_dict()
|
|
||||||
|
|
||||||
def load_state_dict(self, state_dict: Dict[str, Any]):
|
|
||||||
super().load_state_dict(state_dict)
|
|
||||||
|
|
||||||
|
|
||||||
class CosineScheduler(BaseScheduler):
|
|
||||||
"""
|
|
||||||
Cosine decay scheduler with warmup, implemented as PyTorch LRScheduler.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
optimizer,
|
|
||||||
warmup_steps: int,
|
|
||||||
lr_decay_steps: int,
|
|
||||||
min_rate: float = 0.05,
|
|
||||||
last_epoch: int = -1
|
|
||||||
):
|
|
||||||
self.warmup_steps = warmup_steps
|
|
||||||
self.lr_decay_steps = lr_decay_steps
|
|
||||||
self.min_rate = min_rate
|
|
||||||
self.total_steps = warmup_steps + lr_decay_steps
|
|
||||||
super().__init__(optimizer, last_epoch)
|
|
||||||
|
|
||||||
|
|
||||||
def get_lr(self) -> List[float]:
|
|
||||||
# warmup
|
|
||||||
if self.last_epoch < self.warmup_steps:
|
|
||||||
warmup_factor = max(self.min_rate, self.last_epoch / self.warmup_steps)
|
|
||||||
return [base_lr * warmup_factor for base_lr in self.base_lrs]
|
|
||||||
|
|
||||||
# cosine decay
|
|
||||||
decay_progress = (self.last_epoch - self.warmup_steps) / self.lr_decay_steps
|
|
||||||
decay_progress = min(decay_progress, 1.0)
|
|
||||||
cosine_decay = 0.5 * (1.0 + math.cos(math.pi * decay_progress))
|
|
||||||
decay_factor = max(self.min_rate, cosine_decay)
|
|
||||||
return [base_lr * decay_factor for base_lr in self.base_lrs]
|
|
||||||
|
|
||||||
def state_dict(self):
|
|
||||||
state = super().state_dict()
|
|
||||||
state.update({
|
|
||||||
'warmup_steps': self.warmup_steps,
|
|
||||||
'lr_decay_steps': self.lr_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.lr_decay_steps = state_dict.pop('lr_decay_steps')
|
|
||||||
self.min_rate = state_dict.pop('min_rate')
|
|
||||||
self.total_steps = state_dict.pop('total_steps')
|
|
||||||
super().load_state_dict(state_dict)
|
|
||||||
|
|
||||||
|
|
||||||
class SGDRScheduler(BaseScheduler):
|
|
||||||
"""
|
|
||||||
SGDR (Stochastic Gradient Descent with Warm Restarts) scheduler,
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
optimizer,
|
|
||||||
warmup_steps: int,
|
|
||||||
cycle_length: int,
|
|
||||||
min_rate: float = 0.05,
|
|
||||||
t_mult: int = 2,
|
|
||||||
last_epoch: int = -1,
|
|
||||||
):
|
|
||||||
self.warmup_steps = warmup_steps
|
|
||||||
self.cycle_length = cycle_length
|
|
||||||
self.min_rate = min_rate
|
|
||||||
self.t_mult = t_mult
|
|
||||||
|
|
||||||
super().__init__(optimizer, last_epoch)
|
|
||||||
|
|
||||||
|
|
||||||
def get_lr(self):
|
|
||||||
# warmup
|
|
||||||
if self.last_epoch < self.warmup_steps:
|
|
||||||
warmup_factor = max(self.min_rate, self.last_epoch / self.warmup_steps)
|
|
||||||
return [base_lr * warmup_factor for base_lr in self.base_lrs]
|
|
||||||
|
|
||||||
# SGDR
|
|
||||||
steps_since_warmup = self.last_epoch - self.warmup_steps
|
|
||||||
|
|
||||||
# 1. Calculate current cycle and position within cycle
|
|
||||||
current_cycle_length = self.cycle_length
|
|
||||||
total_cycles_length = 0
|
|
||||||
cycle_num = 0
|
|
||||||
|
|
||||||
while total_cycles_length + current_cycle_length <= steps_since_warmup:
|
|
||||||
total_cycles_length += current_cycle_length
|
|
||||||
current_cycle_length *= self.t_mult
|
|
||||||
cycle_num += 1
|
|
||||||
|
|
||||||
steps_in_cycle = steps_since_warmup - total_cycles_length
|
|
||||||
|
|
||||||
# 2. Cosine annealing within the current cycle
|
|
||||||
cosine_factor = 0.5 * (1 + math.cos(math.pi * steps_in_cycle / current_cycle_length))
|
|
||||||
learning_rate_factor = self.min_rate + (1 - self.min_rate) * cosine_factor
|
|
||||||
|
|
||||||
return [base_lr * learning_rate_factor for base_lr in self.base_lrs]
|
|
||||||
|
|
||||||
def state_dict(self):
|
|
||||||
"""Returns the state of the scheduler as a dict."""
|
|
||||||
state = super().state_dict()
|
|
||||||
state.update({
|
|
||||||
'warmup_steps': self.warmup_steps,
|
|
||||||
'cycle_length': self.cycle_length,
|
|
||||||
'min_rate': self.min_rate,
|
|
||||||
't_mult': self.t_mult
|
|
||||||
})
|
|
||||||
return state
|
|
||||||
|
|
||||||
def load_state_dict(self, state_dict):
|
|
||||||
"""Loads the scheduler's state."""
|
|
||||||
self.warmup_steps = state_dict.pop('warmup_steps')
|
|
||||||
self.cycle_length = state_dict.pop('cycle_length')
|
|
||||||
self.min_rate = state_dict.pop('min_rate')
|
|
||||||
self.t_mult = state_dict.pop('t_mult')
|
|
||||||
super().load_state_dict(state_dict)
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class SchedulerFactory:
|
|
||||||
"""
|
|
||||||
Factory class for creating learning rate schedulers.
|
|
||||||
"""
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def load_scheduler(optimizer, scedule_config: ScheduleConfig) -> BaseScheduler:
|
|
||||||
kwargs = scedule_config.get_kwargs()
|
|
||||||
schedule_type = kwargs.pop("schedule_type")
|
|
||||||
|
|
||||||
if schedule_type == "cosine":
|
|
||||||
return CosineScheduler(optimizer, **kwargs)
|
|
||||||
elif schedule_type == "sgdr":
|
|
||||||
return SGDRScheduler(optimizer, **kwargs)
|
|
||||||
else:
|
|
||||||
raise ValueError(f"Unsupported schedule type: {schedule_type}")
|
|
||||||
|
|
||||||
@@ -1,167 +0,0 @@
|
|||||||
import copy
|
|
||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
import torch.nn.functional as F
|
|
||||||
|
|
||||||
from torch import Tensor
|
|
||||||
from typing import Any, Tuple, Callable, Dict, Union
|
|
||||||
from abc import ABC, abstractmethod
|
|
||||||
|
|
||||||
|
|
||||||
def get_logprobs(model:nn.Module, input_ids: Tensor, mask: Tensor, pad_token_id: int):
|
|
||||||
input_mask = input_ids.ne(pad_token_id)
|
|
||||||
logits = model(input_ids, input_mask)["logits"]
|
|
||||||
log_probs = torch.log_softmax(logits, dim=-1)
|
|
||||||
|
|
||||||
shifted_log_probs = log_probs[:, :-1, :]
|
|
||||||
shifted_input_ids = input_ids[:, 1:]
|
|
||||||
shifted_response_mask = mask[:, 1:]
|
|
||||||
|
|
||||||
token_logprobs = torch.gather(
|
|
||||||
shifted_log_probs,
|
|
||||||
dim=-1,
|
|
||||||
index=shifted_input_ids.unsqueeze(-1)
|
|
||||||
).squeeze(-1)
|
|
||||||
|
|
||||||
prompt_mask = input_mask[:, 1:]
|
|
||||||
valid_mask = (prompt_mask & shifted_response_mask).float()
|
|
||||||
|
|
||||||
return (token_logprobs * valid_mask).sum(dim=-1)
|
|
||||||
|
|
||||||
def move_to_device(batch:Dict[str, Tensor], device: str) -> Any:
|
|
||||||
return {key: value.to(device, non_blocking=True) for key, value in batch.items()}
|
|
||||||
|
|
||||||
|
|
||||||
class BaseStrategy(ABC):
|
|
||||||
def __init__(self, model: Union[nn.Module, Callable[..., Dict[str, Tensor]]], device: str):
|
|
||||||
self.model = model
|
|
||||||
self.device = device
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
|
||||||
raise NotImplementedError
|
|
||||||
|
|
||||||
def __call__(self, batch: Tuple[Tensor, ...]) -> Tensor:
|
|
||||||
return self.compute_loss(batch)
|
|
||||||
|
|
||||||
|
|
||||||
class SeqStrategy(BaseStrategy):
|
|
||||||
def __init__(self, model, device):
|
|
||||||
super().__init__(model, device)
|
|
||||||
|
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
|
||||||
batch = move_to_device(batch, self.device)
|
|
||||||
input_ids, target_ids = batch["input_ids"], batch["target_ids"]
|
|
||||||
logits = self.model(input_ids=input_ids)["logits"]
|
|
||||||
|
|
||||||
loss = F.cross_entropy(
|
|
||||||
input=logits.flatten(0, 1),
|
|
||||||
target=target_ids.flatten()
|
|
||||||
)
|
|
||||||
|
|
||||||
return loss
|
|
||||||
|
|
||||||
|
|
||||||
class SftStrategy(BaseStrategy):
|
|
||||||
def __init__(self, model, device):
|
|
||||||
super().__init__(model, device)
|
|
||||||
|
|
||||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
|
||||||
batch = move_to_device(batch, self.device)
|
|
||||||
input_ids, target_ids, loss_mask = batch["input_ids"], batch["target_ids"], batch["loss_mask"]
|
|
||||||
|
|
||||||
ignore_index = -100
|
|
||||||
logits = self.model(input_ids=input_ids)["logits"]
|
|
||||||
target_ids = target_ids.masked_fill(loss_mask == 0, ignore_index)
|
|
||||||
|
|
||||||
loss = F.cross_entropy(
|
|
||||||
input=logits.flatten(0, 1),
|
|
||||||
target=target_ids.flatten(),
|
|
||||||
ignore_index=ignore_index
|
|
||||||
)
|
|
||||||
|
|
||||||
return loss
|
|
||||||
|
|
||||||
|
|
||||||
class DpoStrategy(BaseStrategy):
|
|
||||||
def __init__(self, model, device, pad_token_id, beta):
|
|
||||||
super().__init__(model, device)
|
|
||||||
ref_model = copy.deepcopy(self.model)
|
|
||||||
ref_model.requires_grad_(False)
|
|
||||||
ref_model.eval()
|
|
||||||
|
|
||||||
self.ref_model = ref_model
|
|
||||||
self.pad_token_id = pad_token_id
|
|
||||||
self.beta = beta
|
|
||||||
|
|
||||||
def compute_loss(self, batch: Tuple[Tensor, ...]) -> Tensor:
|
|
||||||
batch = move_to_device(batch, self.device)
|
|
||||||
good_ids, bad_ids = batch["chosen"], batch["rejected"]
|
|
||||||
good_mask, bad_mask = batch["chosen_mask"], batch["rejected_mask"]
|
|
||||||
|
|
||||||
log_pi_good = get_logprobs(self.model, good_ids, good_mask, self.pad_token_id)
|
|
||||||
log_pi_bad = get_logprobs(self.model, bad_ids, bad_mask, self.pad_token_id)
|
|
||||||
|
|
||||||
with torch.no_grad():
|
|
||||||
log_ref_good = get_logprobs(self.ref_model, good_ids, good_mask, self.pad_token_id)
|
|
||||||
log_ref_bad = get_logprobs(self.ref_model, bad_ids, bad_mask, self.pad_token_id)
|
|
||||||
|
|
||||||
pi_log_ratio = log_pi_good - log_pi_bad
|
|
||||||
ref_log_ratio = log_ref_good - log_ref_bad
|
|
||||||
|
|
||||||
ratio_diff = pi_log_ratio - ref_log_ratio
|
|
||||||
|
|
||||||
dpo_loss = -F.logsigmoid(self.beta * ratio_diff).mean()
|
|
||||||
return dpo_loss
|
|
||||||
|
|
||||||
|
|
||||||
class PpoStrategy(BaseStrategy):
|
|
||||||
def __init__(self, model, pad_token_id, epsilon):
|
|
||||||
super().__init__(model)
|
|
||||||
ref_model = copy.deepcopy(self.model)
|
|
||||||
ref_model.requires_grad_(False)
|
|
||||||
ref_model.eval()
|
|
||||||
|
|
||||||
self.ref_model = ref_model
|
|
||||||
self.pad_token_id = pad_token_id
|
|
||||||
self.epsilon = epsilon
|
|
||||||
|
|
||||||
def ppo_clip_loss_masked(
|
|
||||||
self,
|
|
||||||
log_probs: Tensor,
|
|
||||||
old_log_probs: Tensor,
|
|
||||||
advantages: Tensor,
|
|
||||||
values: Tensor,
|
|
||||||
returns: Tensor,
|
|
||||||
mask: Tensor,
|
|
||||||
clip_eps: float=0.2,
|
|
||||||
):
|
|
||||||
ratio = torch.exp(log_probs - old_log_probs)
|
|
||||||
surr1 = ratio * advantages
|
|
||||||
surr2 = torch.clamp(ratio, 1 - clip_eps, 1 + clip_eps) * advantages
|
|
||||||
policy_loss = -torch.min(surr1, surr2).masked_select(mask).mean()
|
|
||||||
|
|
||||||
value_loss = F.mse_loss(values.masked_select(mask),
|
|
||||||
returns.masked_select(mask))
|
|
||||||
|
|
||||||
entropy = -(log_probs.exp() * log_probs).masked_select(mask).mean()
|
|
||||||
entropy_loss = -entropy
|
|
||||||
return policy_loss, value_loss, entropy_loss
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class StrategyFactory:
|
|
||||||
|
|
||||||
def load(model, train_type, device, **kwargs):
|
|
||||||
train_strategy: Dict[str, Callable[[], BaseStrategy]] = {
|
|
||||||
"seq": lambda: SeqStrategy(model, device),
|
|
||||||
"sft": lambda: SftStrategy(model, device),
|
|
||||||
"dpo": lambda: DpoStrategy(
|
|
||||||
model,
|
|
||||||
device,
|
|
||||||
kwargs.get("pad_token_id"),
|
|
||||||
kwargs.get("dpo_beta")
|
|
||||||
)
|
|
||||||
}
|
|
||||||
strategy = train_strategy[train_type]()
|
|
||||||
return strategy
|
|
||||||
@@ -1,226 +0,0 @@
|
|||||||
import os
|
|
||||||
import json
|
|
||||||
import time
|
|
||||||
|
|
||||||
from pathlib import Path
|
|
||||||
from tqdm import tqdm
|
|
||||||
from torch.nn.utils import clip_grad_norm_
|
|
||||||
from torch.optim.lr_scheduler import LambdaLR
|
|
||||||
from typing import List, Optional, Protocol, TYPE_CHECKING
|
|
||||||
|
|
||||||
from khaosz.config import ScheduleConfig
|
|
||||||
from khaosz.trainer.metric_util import (
|
|
||||||
grad_max,
|
|
||||||
grad_min,
|
|
||||||
grad_norm,
|
|
||||||
grad_mean,
|
|
||||||
grad_std,
|
|
||||||
grad_nan_num
|
|
||||||
)
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from khaosz.trainer.trainer import Trainer
|
|
||||||
from khaosz.trainer.train_context import TrainContext
|
|
||||||
|
|
||||||
|
|
||||||
class TrainCallback(Protocol):
|
|
||||||
"""
|
|
||||||
Callback interface for trainer.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def on_train_begin(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
""" Called at the beginning of training. """
|
|
||||||
|
|
||||||
def on_train_end(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
""" Called at the end of training. """
|
|
||||||
|
|
||||||
def on_epoch_begin(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
""" Called at the beginning of each epoch. """
|
|
||||||
|
|
||||||
def on_epoch_end(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
""" Called at the end of each epoch. """
|
|
||||||
|
|
||||||
def on_step_begin(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
""" Called at the beginning of each step. """
|
|
||||||
|
|
||||||
def on_step_end(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
""" Called at the end of each step."""
|
|
||||||
|
|
||||||
def on_batch_begin(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
""" Called at the beginning of each batch. """
|
|
||||||
|
|
||||||
def on_batch_end(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
""" Called at the end of each batch. """
|
|
||||||
|
|
||||||
def on_error(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
""" Called when an error occurs during training. """
|
|
||||||
|
|
||||||
|
|
||||||
class GradientClippingCallback(TrainCallback):
|
|
||||||
"""
|
|
||||||
Gradient clipping callback for trainer.
|
|
||||||
"""
|
|
||||||
def __init__(self, max_grad_norm: float):
|
|
||||||
self.max_grad_norm = max_grad_norm
|
|
||||||
|
|
||||||
def on_step_begin(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
_ = context
|
|
||||||
clip_grad_norm_(trainer.parameter.model.parameters(), self.max_grad_norm)
|
|
||||||
|
|
||||||
|
|
||||||
class SchedulerCallback(TrainCallback):
|
|
||||||
"""
|
|
||||||
Scheduler callback for trainer.
|
|
||||||
"""
|
|
||||||
def __init__(self, schedule_config: ScheduleConfig):
|
|
||||||
self.schedule_config = schedule_config
|
|
||||||
self.scheduler: Optional[LambdaLR] = None
|
|
||||||
|
|
||||||
def on_train_begin(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
|
|
||||||
for group in trainer.train_config.optimizer.param_groups:
|
|
||||||
if "initial_lr" not in group:
|
|
||||||
group["initial_lr"] = group["lr"]
|
|
||||||
|
|
||||||
self.scheduler = context.scheduler
|
|
||||||
|
|
||||||
def on_batch_end(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
_ = trainer, context
|
|
||||||
if self.scheduler:
|
|
||||||
self.scheduler.step()
|
|
||||||
|
|
||||||
|
|
||||||
class CheckpointCallback(TrainCallback):
|
|
||||||
"""
|
|
||||||
Checkpoint callback for trainer.
|
|
||||||
"""
|
|
||||||
def __init__(self, checkpoint_interval: int):
|
|
||||||
self.checkpoint_interval = checkpoint_interval
|
|
||||||
self.last_ckpt_iter = 0
|
|
||||||
|
|
||||||
def _save_checkpoint(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
save_path = os.path.join(trainer.train_config.checkpoint_dir, f"iter_{context.batch_iter}")
|
|
||||||
context.checkpoint.optimizer_state = context.optimizer.state_dict()
|
|
||||||
context.checkpoint.scheduler_state = context.scheduler.state_dict()
|
|
||||||
context.checkpoint.epoch = context.epoch
|
|
||||||
context.checkpoint.batch_iter = context.batch_iter
|
|
||||||
context.checkpoint.save(save_path)
|
|
||||||
self.last_ckpt_iter = context.batch_iter
|
|
||||||
|
|
||||||
def on_batch_end(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
context.checkpoint.loss_list.append(context.loss)
|
|
||||||
|
|
||||||
if context.batch_iter - self.last_ckpt_iter >= self.checkpoint_interval:
|
|
||||||
self._save_checkpoint(trainer, context)
|
|
||||||
|
|
||||||
def on_train_end(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
if context.batch_iter != self.last_ckpt_iter:
|
|
||||||
self._save_checkpoint(trainer, context)
|
|
||||||
|
|
||||||
|
|
||||||
class ProgressBarCallback(TrainCallback):
|
|
||||||
"""
|
|
||||||
Progress bar callback for trainer.
|
|
||||||
"""
|
|
||||||
def __init__(self):
|
|
||||||
self.progress_bar: tqdm = None
|
|
||||||
|
|
||||||
def on_epoch_begin(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
self.progress_bar = tqdm(
|
|
||||||
context.dataloader,
|
|
||||||
desc=f"Epoch {context.epoch+1}/{trainer.train_config.n_epoch}",
|
|
||||||
dynamic_ncols=True
|
|
||||||
)
|
|
||||||
|
|
||||||
def on_batch_end(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
_ = trainer
|
|
||||||
self.progress_bar.set_postfix({
|
|
||||||
"loss": f"{context.loss:.4f}",
|
|
||||||
"lr": f"{context.optimizer.param_groups[-1]['lr']:.2e}"
|
|
||||||
})
|
|
||||||
self.progress_bar.update(1)
|
|
||||||
|
|
||||||
def on_epoch_end(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
_ = trainer, context
|
|
||||||
if self.progress_bar:
|
|
||||||
self.progress_bar.close()
|
|
||||||
|
|
||||||
|
|
||||||
class StepMonitorCallback(TrainCallback):
|
|
||||||
"""
|
|
||||||
Customizable logger callback for trainer.
|
|
||||||
|
|
||||||
This callback provides flexible logging capabilities for training metrics,
|
|
||||||
supporting multiple log formats and custom log handlers.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
log_dir: Optional[str] = None,
|
|
||||||
log_interval: int = 100,
|
|
||||||
metrics: Optional[List[str]] = None
|
|
||||||
):
|
|
||||||
"""
|
|
||||||
Args:
|
|
||||||
log_dir: Directory to save log files. If None, logs won't be saved to file.
|
|
||||||
log_interval: Log every N steps
|
|
||||||
metrics: List of metrics to log. Supported: ['loss', 'lr', 'grad_norm', 'grad_std',
|
|
||||||
grad_max', 'grad_min', 'grad_mean', 'grad_nan_num']
|
|
||||||
custom_handlers: List of custom log handler functions
|
|
||||||
json_log: Whether to save logs in JSON format
|
|
||||||
"""
|
|
||||||
|
|
||||||
self.log_dir = Path(log_dir) if log_dir else Path(os.getcwd()) / "logs"
|
|
||||||
self.log_interval = log_interval
|
|
||||||
self.metrics = metrics or ['loss', 'lr']
|
|
||||||
self.step_num = 0
|
|
||||||
|
|
||||||
self.log_dir.mkdir(parents=True, exist_ok=True)
|
|
||||||
|
|
||||||
def _handle_info(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
""" Logs training information to console and file. """
|
|
||||||
|
|
||||||
log_data = {
|
|
||||||
"timestamp": time.strftime('%Y-%m-%d %H:%M:%S'),
|
|
||||||
"epoch": context.epoch,
|
|
||||||
"iter": context.batch_iter,
|
|
||||||
"metrics": self.metrics,
|
|
||||||
}
|
|
||||||
|
|
||||||
for metric in self.metrics:
|
|
||||||
if metric == 'loss':
|
|
||||||
log_data[metric] = context.loss
|
|
||||||
elif metric == 'lr':
|
|
||||||
log_data[metric] = context.optimizer.param_groups[-1]['lr']
|
|
||||||
elif metric == 'grad_norm':
|
|
||||||
log_data[metric] = grad_norm(trainer.parameter.model)
|
|
||||||
elif metric == 'grad_std':
|
|
||||||
log_data[metric] = grad_std(trainer.parameter.model)
|
|
||||||
elif metric == 'grad_max':
|
|
||||||
log_data[metric] = grad_max(trainer.parameter.model)
|
|
||||||
elif metric == 'grad_min':
|
|
||||||
log_data[metric] = grad_min(trainer.parameter.model)
|
|
||||||
elif metric == 'grad_mean':
|
|
||||||
log_data[metric] = grad_mean(trainer.parameter.model)
|
|
||||||
elif metric == 'grad_nan_num':
|
|
||||||
log_data[metric] = grad_nan_num(trainer.parameter.model)
|
|
||||||
else:
|
|
||||||
raise ValueError(f"Invalid metric: {metric}")
|
|
||||||
|
|
||||||
return log_data
|
|
||||||
|
|
||||||
def _handle_log(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
""" Logs training information to console and file. """
|
|
||||||
log_data = self._handle_info(trainer, context)
|
|
||||||
try:
|
|
||||||
log_file = self.log_dir / f"log_epoch_{context.epoch}_iter_{context.batch_iter}.json"
|
|
||||||
with open(log_file, 'a') as f:
|
|
||||||
json.dump(log_data, f, indent=4)
|
|
||||||
except Exception:
|
|
||||||
raise
|
|
||||||
|
|
||||||
def on_step_end(self, trainer: 'Trainer', context: 'TrainContext'):
|
|
||||||
if self.step_num % self.log_interval == 0:
|
|
||||||
self._handle_log(trainer, context)
|
|
||||||
|
|
||||||
self.step_num += 1
|
|
||||||
@@ -1,105 +0,0 @@
|
|||||||
from dataclasses import dataclass, field, fields
|
|
||||||
from typing import Optional, Self, TYPE_CHECKING
|
|
||||||
from torch.optim import Optimizer
|
|
||||||
from torch.utils.data import DataLoader
|
|
||||||
from khaosz.config import Checkpoint
|
|
||||||
from khaosz.data import ResumeableRandomSampler
|
|
||||||
from khaosz.trainer.schedule import BaseScheduler, SchedulerFactory
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from khaosz.trainer.trainer import Trainer
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class TrainContext:
|
|
||||||
dataloader: DataLoader = field(default=None)
|
|
||||||
optimizer: Optimizer = field(default=None)
|
|
||||||
scheduler: BaseScheduler = field(default=None)
|
|
||||||
checkpoint: Checkpoint = field(default=None)
|
|
||||||
epoch: int = field(default=0)
|
|
||||||
batch_iter: int = field(default=0)
|
|
||||||
loss: float = field(default=0.0)
|
|
||||||
|
|
||||||
def asdict(self) -> dict:
|
|
||||||
return {field.name: getattr(self, field.name)
|
|
||||||
for field in fields(self)}
|
|
||||||
|
|
||||||
|
|
||||||
class TrainContextBuilder:
|
|
||||||
def __init__(self, trainer: 'Trainer'):
|
|
||||||
self.trainer = trainer
|
|
||||||
self._context: TrainContext = None
|
|
||||||
|
|
||||||
def with_checkpoint(self, checkpoint: Optional[Checkpoint]) -> Self:
|
|
||||||
self._context = TrainContext()
|
|
||||||
if checkpoint is None:
|
|
||||||
checkpoint = Checkpoint(
|
|
||||||
model=self.trainer.parameter.model,
|
|
||||||
tokenizer=self.trainer.parameter.tokenizer,
|
|
||||||
config=self.trainer.parameter.config,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
# resume from the assigned checkpoint or assigned iteration
|
|
||||||
self._context.epoch = max(checkpoint.epoch, self.trainer.train_config.start_epoch)
|
|
||||||
self._context.batch_iter = max(checkpoint.batch_iter, self.trainer.train_config.start_batch)
|
|
||||||
|
|
||||||
self._context.checkpoint = checkpoint
|
|
||||||
return self
|
|
||||||
|
|
||||||
def with_optimizer(self) -> Self:
|
|
||||||
if self._context is None:
|
|
||||||
raise RuntimeError("Must call with_checkpoint() before with_optimizer()")
|
|
||||||
|
|
||||||
optimizer = self.trainer.train_config.optimizer
|
|
||||||
|
|
||||||
if self._context.checkpoint and self._context.checkpoint.optimizer_state:
|
|
||||||
optimizer.load_state_dict(self._context.checkpoint.optimizer_state)
|
|
||||||
|
|
||||||
self._context.optimizer = optimizer
|
|
||||||
|
|
||||||
if self._context.checkpoint:
|
|
||||||
self._context.checkpoint.optimizer_state = optimizer.state_dict()
|
|
||||||
|
|
||||||
return self
|
|
||||||
|
|
||||||
def with_scheduler(self) -> Self:
|
|
||||||
if not hasattr(self._context, 'optimizer') or self._context.optimizer is None:
|
|
||||||
raise RuntimeError("Must call with_optimizer() before with_scheduler()")
|
|
||||||
|
|
||||||
optimizer = self.trainer.train_config.optimizer
|
|
||||||
schedule_config = self.trainer.schedule_config
|
|
||||||
scheduler = SchedulerFactory.load_scheduler(optimizer, schedule_config)
|
|
||||||
|
|
||||||
if self._context.checkpoint and self._context.checkpoint.scheduler_state:
|
|
||||||
scheduler.load_state_dict(self._context.checkpoint.scheduler_state)
|
|
||||||
|
|
||||||
self._context.scheduler = scheduler
|
|
||||||
|
|
||||||
if self._context.checkpoint:
|
|
||||||
self._context.checkpoint.scheduler_state = scheduler.state_dict()
|
|
||||||
|
|
||||||
return self
|
|
||||||
|
|
||||||
def with_dataloader(self) -> Self:
|
|
||||||
# fix: change batch level batch_iter to sample level offset
|
|
||||||
sampler_offset = self._context.batch_iter * self.trainer.train_config.batch_size
|
|
||||||
resumeable_sampler = ResumeableRandomSampler(
|
|
||||||
data_source=self.trainer.train_config.dataset,
|
|
||||||
start_epoch=self._context.epoch,
|
|
||||||
start_iter=sampler_offset,
|
|
||||||
seed=self.trainer.train_config.random_seed
|
|
||||||
)
|
|
||||||
|
|
||||||
dataloader = DataLoader(
|
|
||||||
self.trainer.train_config.dataset,
|
|
||||||
batch_size=self.trainer.train_config.batch_size,
|
|
||||||
sampler=resumeable_sampler,
|
|
||||||
num_workers=self.trainer.train_config.num_workers,
|
|
||||||
pin_memory=self.trainer.train_config.pin_memory,
|
|
||||||
prefetch_factor=self.trainer.train_config.prefetch_factor
|
|
||||||
)
|
|
||||||
self._context.dataloader = dataloader
|
|
||||||
return self
|
|
||||||
|
|
||||||
def build(self) -> TrainContext:
|
|
||||||
return self._context
|
|
||||||
@@ -1,95 +0,0 @@
|
|||||||
import logging
|
|
||||||
from typing import Optional, List
|
|
||||||
from khaosz.config import (
|
|
||||||
ModelParameter,
|
|
||||||
Checkpoint,
|
|
||||||
ScheduleConfig,
|
|
||||||
TrainConfig
|
|
||||||
)
|
|
||||||
from khaosz.trainer.train_callback import (
|
|
||||||
TrainCallback,
|
|
||||||
ProgressBarCallback,
|
|
||||||
CheckpointCallback,
|
|
||||||
GradientClippingCallback,
|
|
||||||
SchedulerCallback
|
|
||||||
)
|
|
||||||
from khaosz.trainer.train_context import TrainContext, TrainContextBuilder
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
|
|
||||||
class Trainer:
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
parameter: ModelParameter,
|
|
||||||
train_config: TrainConfig,
|
|
||||||
schedule_config: ScheduleConfig,
|
|
||||||
callbacks: Optional[List[TrainCallback]] = None
|
|
||||||
):
|
|
||||||
self.parameter = parameter
|
|
||||||
self.train_config = train_config
|
|
||||||
self.schedule_config = schedule_config
|
|
||||||
self.callbacks = callbacks or self._get_default_callbacks()
|
|
||||||
|
|
||||||
def _get_default_callbacks(self) -> List[TrainCallback]:
|
|
||||||
return [
|
|
||||||
ProgressBarCallback(),
|
|
||||||
CheckpointCallback(self.train_config.checkpoint_interval),
|
|
||||||
GradientClippingCallback(self.train_config.max_grad_norm),
|
|
||||||
SchedulerCallback(self.schedule_config),
|
|
||||||
]
|
|
||||||
|
|
||||||
def _build_context(self, checkpoint: Optional[Checkpoint]) -> TrainContext:
|
|
||||||
return (TrainContextBuilder(self)
|
|
||||||
.with_checkpoint(checkpoint)
|
|
||||||
.with_optimizer()
|
|
||||||
.with_scheduler()
|
|
||||||
.with_dataloader()
|
|
||||||
.build())
|
|
||||||
|
|
||||||
def _call_callbacks(self, method_name: str, context: TrainContext):
|
|
||||||
for callback in self.callbacks:
|
|
||||||
method = getattr(callback, method_name, None)
|
|
||||||
if method:
|
|
||||||
method(self, context)
|
|
||||||
|
|
||||||
def train(self, checkpoint: Optional[Checkpoint] = None) -> Checkpoint:
|
|
||||||
context = self._build_context(checkpoint)
|
|
||||||
self._call_callbacks('on_train_begin', context)
|
|
||||||
|
|
||||||
try:
|
|
||||||
self.parameter.model.train()
|
|
||||||
# 1.epoch
|
|
||||||
for epoch in range(context.epoch, self.train_config.n_epoch):
|
|
||||||
context.epoch = epoch
|
|
||||||
self._call_callbacks('on_epoch_begin', context)
|
|
||||||
|
|
||||||
for batch in context.dataloader:
|
|
||||||
if context.batch_iter % self.train_config.accumulation_steps == 0:
|
|
||||||
# 2. step
|
|
||||||
self._call_callbacks('on_step_begin', context)
|
|
||||||
self.train_config.optimizer.step()
|
|
||||||
self.train_config.optimizer.zero_grad()
|
|
||||||
self._call_callbacks('on_step_end', context)
|
|
||||||
|
|
||||||
# 3. batch
|
|
||||||
self._call_callbacks('on_batch_begin', context)
|
|
||||||
loss = self.train_config.strategy(batch)
|
|
||||||
context.loss = loss.item()
|
|
||||||
context.batch_iter += 1
|
|
||||||
|
|
||||||
# to make the loss normalized by accumulation steps
|
|
||||||
normalized_loss = loss / self.train_config.accumulation_steps
|
|
||||||
normalized_loss.backward()
|
|
||||||
|
|
||||||
self._call_callbacks('on_batch_end', context)
|
|
||||||
|
|
||||||
self._call_callbacks('on_epoch_end', context)
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Training failed: {str(e)}", exc_info=True)
|
|
||||||
self._call_callbacks('on_error', context)
|
|
||||||
raise
|
|
||||||
finally:
|
|
||||||
self._call_callbacks('on_train_end', context)
|
|
||||||
return context.checkpoint
|
|
||||||
@@ -1,88 +0,0 @@
|
|||||||
import torch
|
|
||||||
import sqlite3
|
|
||||||
import numpy as np
|
|
||||||
from torch import Tensor
|
|
||||||
from typing import Dict, List, Tuple
|
|
||||||
|
|
||||||
|
|
||||||
class Retriever:
|
|
||||||
def __init__(self, db_path=None):
|
|
||||||
self.data: Dict[str, Tensor] = {}
|
|
||||||
self.embedding_cache: Tensor = None
|
|
||||||
self.is_caculated: bool = False
|
|
||||||
|
|
||||||
if db_path is not None:
|
|
||||||
self.load(db_path)
|
|
||||||
|
|
||||||
def retrieve(self, query: Tensor, top_k: int) -> List[Tuple[str, float]]:
|
|
||||||
if not self.data:
|
|
||||||
return []
|
|
||||||
|
|
||||||
query = query.flatten().unsqueeze(1) # [dim, 1]
|
|
||||||
norm_embeddings = self._embeddings.to(
|
|
||||||
device=query.device,
|
|
||||||
dtype=query.dtype
|
|
||||||
) # [n_vectors, dim]
|
|
||||||
sim_scores = torch.matmul(norm_embeddings, query).squeeze() # [n_vectors]
|
|
||||||
|
|
||||||
top_k = min(top_k, len(self.data))
|
|
||||||
indices = sim_scores.topk(top_k).indices
|
|
||||||
keys = list(self.data.keys())
|
|
||||||
|
|
||||||
return [(keys[i], sim_scores[i].item()) for i in indices]
|
|
||||||
|
|
||||||
def add_vector(self, key: str, vector_data: Tensor):
|
|
||||||
self.is_caculated = False
|
|
||||||
self.data[key] = vector_data.flatten().float().cpu()
|
|
||||||
|
|
||||||
def delete_vector(self, key: str):
|
|
||||||
self.is_caculated = False
|
|
||||||
self.data.pop(key, None)
|
|
||||||
|
|
||||||
def save(self, db_path):
|
|
||||||
conn = sqlite3.connect(db_path)
|
|
||||||
cursor = conn.cursor()
|
|
||||||
self._init_db(cursor)
|
|
||||||
cursor.execute('DELETE FROM vectors')
|
|
||||||
|
|
||||||
for item, vec in self.data.items():
|
|
||||||
vec_bytes = vec.numpy().tobytes()
|
|
||||||
cursor.execute('INSERT OR REPLACE INTO vectors (key, vector) VALUES (?, ?)',
|
|
||||||
(item, vec_bytes))
|
|
||||||
|
|
||||||
conn.commit()
|
|
||||||
conn.close()
|
|
||||||
|
|
||||||
def load(self, db_path):
|
|
||||||
conn = sqlite3.connect(db_path)
|
|
||||||
cursor = conn.cursor()
|
|
||||||
self._init_db(cursor)
|
|
||||||
cursor.execute('SELECT key, vector FROM vectors')
|
|
||||||
rows = cursor.fetchall()
|
|
||||||
self.data = {}
|
|
||||||
|
|
||||||
for row in rows:
|
|
||||||
key, vec_bytes = row
|
|
||||||
vec_numpy = np.frombuffer(vec_bytes, dtype=np.float32).copy()
|
|
||||||
vec = torch.from_numpy(vec_numpy)
|
|
||||||
self.data[key] = vec
|
|
||||||
|
|
||||||
conn.close()
|
|
||||||
|
|
||||||
def _init_db(self,cursor: sqlite3.Cursor):
|
|
||||||
# Create table if not exists (in case loading from a new database)
|
|
||||||
cursor.execute('''
|
|
||||||
CREATE TABLE IF NOT EXISTS vectors (
|
|
||||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
||||||
key TEXT UNIQUE NOT NULL,
|
|
||||||
vector BLOB NOT NULL
|
|
||||||
)''')
|
|
||||||
|
|
||||||
@property
|
|
||||||
def _embeddings(self) -> Tensor:
|
|
||||||
if not self.is_caculated:
|
|
||||||
embeddings = torch.stack(list(self.data.values()))
|
|
||||||
norm_embeddings = embeddings / torch.norm(embeddings, dim=-1, keepdim=True)
|
|
||||||
self.embedding_cache = norm_embeddings
|
|
||||||
|
|
||||||
return self.embedding_cache
|
|
||||||
@@ -1,127 +0,0 @@
|
|||||||
import re
|
|
||||||
import torch
|
|
||||||
import torch.nn.functional as F
|
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
|
||||||
from torch import Tensor
|
|
||||||
from typing import List, Callable, Optional
|
|
||||||
|
|
||||||
|
|
||||||
class BaseTextSplitter(ABC):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
max_len: int = 512,
|
|
||||||
chunk_overlap: int = 0,
|
|
||||||
):
|
|
||||||
if max_len <= 0:
|
|
||||||
raise ValueError("max_len must be > 0")
|
|
||||||
if chunk_overlap < 0:
|
|
||||||
raise ValueError("chunk_overlap must be >= 0")
|
|
||||||
|
|
||||||
self.max_len = max_len
|
|
||||||
self.chunk_overlap = chunk_overlap
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def split(self, text: str, **kwargs) -> List[str]:
|
|
||||||
raise NotImplementedError
|
|
||||||
|
|
||||||
def preprocess(self, text: str) -> str:
|
|
||||||
return text.strip()
|
|
||||||
|
|
||||||
def postprocess(self, chunks: List[str]) -> List[str]:
|
|
||||||
return [chunk.strip() for chunk in chunks if chunk.strip()]
|
|
||||||
|
|
||||||
|
|
||||||
class PriorityTextSplitter(BaseTextSplitter):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
separators: List[str],
|
|
||||||
max_len: int = 512,
|
|
||||||
chunk_overlap: int = 0,
|
|
||||||
):
|
|
||||||
super().__init__(max_len=max_len, chunk_overlap=chunk_overlap)
|
|
||||||
if not separators:
|
|
||||||
raise ValueError("separators must be a non-empty list")
|
|
||||||
self.separators = separators
|
|
||||||
|
|
||||||
def split(self, text: str) -> List[str]:
|
|
||||||
text = self.preprocess(text)
|
|
||||||
for sep in self.separators:
|
|
||||||
parts = text.split(sep)
|
|
||||||
|
|
||||||
valid_parts = [p.strip() for p in parts if p.strip()]
|
|
||||||
if len(valid_parts) > 1:
|
|
||||||
return self.postprocess(valid_parts)
|
|
||||||
return [text]
|
|
||||||
|
|
||||||
|
|
||||||
class SemanticTextSplitter(BaseTextSplitter):
|
|
||||||
|
|
||||||
DEFAULT_PATTERN = r'(?<=[。!?!?])(?=(?:[^"\'‘’“”]*["\'‘’“”][^"\'‘’“”]*["\'‘’“”])*[^"\'‘’“”]*$)'
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
embedding_func: Callable[[List[str]], List[Tensor]],
|
|
||||||
pattern: Optional[str] = None,
|
|
||||||
max_len: int = 512,
|
|
||||||
chunk_overlap: int = 0,
|
|
||||||
):
|
|
||||||
super().__init__(max_len=max_len, chunk_overlap=chunk_overlap)
|
|
||||||
if not callable(embedding_func):
|
|
||||||
raise TypeError("embedding_func must be callable")
|
|
||||||
self.embedding_func = embedding_func
|
|
||||||
self.pattern = pattern or SemanticTextSplitter.DEFAULT_PATTERN
|
|
||||||
|
|
||||||
def split(
|
|
||||||
self,
|
|
||||||
text: str,
|
|
||||||
threshold: float = 0.5,
|
|
||||||
window_size: int = 1,
|
|
||||||
) -> List[str]:
|
|
||||||
text = self.preprocess(text)
|
|
||||||
sentences = [s.strip() for s in re.split(self.pattern, text) if s.strip()]
|
|
||||||
|
|
||||||
if len(sentences) <= 1:
|
|
||||||
return self.postprocess(sentences)
|
|
||||||
|
|
||||||
try:
|
|
||||||
sentence_embs = self.embedding_func(sentences)
|
|
||||||
except Exception as e:
|
|
||||||
raise RuntimeError(f"Embedding generation failed: {e}")
|
|
||||||
|
|
||||||
if len(sentence_embs) != len(sentences):
|
|
||||||
raise ValueError("Embedding function must return one vector per sentence")
|
|
||||||
|
|
||||||
chunks = []
|
|
||||||
emb_tensor = torch.stack(sentence_embs) # shape: [N, D]
|
|
||||||
current_chunk: List[str] = [sentences[0]]
|
|
||||||
|
|
||||||
for i in range(1, len(sentences)):
|
|
||||||
start_prev = max(0, i - window_size)
|
|
||||||
end_prev = i
|
|
||||||
start_next = i
|
|
||||||
end_next = min(len(sentences), i + window_size)
|
|
||||||
|
|
||||||
prev_window_emb = emb_tensor[start_prev:end_prev].mean(dim=0)
|
|
||||||
next_window_emb = emb_tensor[start_next:end_next].mean(dim=0)
|
|
||||||
|
|
||||||
similarity = F.cosine_similarity(
|
|
||||||
prev_window_emb.unsqueeze(0),
|
|
||||||
next_window_emb.unsqueeze(0),
|
|
||||||
dim=1
|
|
||||||
).item()
|
|
||||||
|
|
||||||
dynamic_threshold = max(threshold * (1 - 0.03 * (end_next - start_prev)), 0.2)
|
|
||||||
|
|
||||||
if similarity < dynamic_threshold:
|
|
||||||
chunks.append(" ".join(current_chunk))
|
|
||||||
overlap_start = max(0, len(current_chunk) - self.chunk_overlap)
|
|
||||||
current_chunk = current_chunk[overlap_start:]
|
|
||||||
current_chunk.append(sentences[i])
|
|
||||||
else:
|
|
||||||
current_chunk.append(sentences[i])
|
|
||||||
|
|
||||||
if current_chunk:
|
|
||||||
chunks.append(" ".join(current_chunk))
|
|
||||||
|
|
||||||
return self.postprocess(chunks)
|
|
||||||
@@ -0,0 +1,52 @@
|
|||||||
|
[build-system]
|
||||||
|
requires = ["setuptools>=64", "wheel"]
|
||||||
|
build-backend = "setuptools.build_meta"
|
||||||
|
|
||||||
|
[project]
|
||||||
|
dynamic = ["version"]
|
||||||
|
name = "astrai"
|
||||||
|
readme = "README.md"
|
||||||
|
requires-python = ">=3.12"
|
||||||
|
dependencies = [
|
||||||
|
"h5py==3.15.1",
|
||||||
|
"numpy==2.3.2",
|
||||||
|
"torch==2.7.1",
|
||||||
|
"tokenizers==0.21.4",
|
||||||
|
"tqdm==4.67.1",
|
||||||
|
"safetensors==0.5.3",
|
||||||
|
"huggingface-hub==0.34.3",
|
||||||
|
"jinja2>=3.0.0",
|
||||||
|
"fastapi",
|
||||||
|
"uvicorn[standard]",
|
||||||
|
"httpx",
|
||||||
|
"requests",
|
||||||
|
]
|
||||||
|
keywords = ["nlp", "datasets", "language-models", "machine-learning"]
|
||||||
|
license = { text = "GPL-3.0" }
|
||||||
|
classifiers = [
|
||||||
|
"Programming Language :: Python :: 3",
|
||||||
|
"License :: OSI Approved :: GPL-3.0",
|
||||||
|
"Operating System :: OS Independent",
|
||||||
|
]
|
||||||
|
urls = { Homepage = "https://github.com/ViperEkura/AstrAI" }
|
||||||
|
|
||||||
|
[project.optional-dependencies]
|
||||||
|
dev = ["pytest==9.0.2", "ruff"]
|
||||||
|
|
||||||
|
[tool.setuptools.packages.find]
|
||||||
|
where = ["."]
|
||||||
|
|
||||||
|
[tool.pip]
|
||||||
|
extra-index-url = "https://download.pytorch.org/whl/cu126"
|
||||||
|
|
||||||
|
[tool.setuptools.dynamic]
|
||||||
|
version = { attr = "astrai.__version__" }
|
||||||
|
|
||||||
|
[tool.ruff]
|
||||||
|
target-version = "py312"
|
||||||
|
|
||||||
|
[tool.ruff.format]
|
||||||
|
quote-style = "double"
|
||||||
|
indent-style = "space"
|
||||||
|
skip-magic-trailing-comma = false
|
||||||
|
line-ending = "auto"
|
||||||
@@ -1,36 +0,0 @@
|
|||||||
# python=3.12
|
|
||||||
--extra-index-url https://download.pytorch.org/whl/cu126
|
|
||||||
|
|
||||||
certifi==2025.8.3
|
|
||||||
charset-normalizer==3.4.2
|
|
||||||
colorama==0.4.6
|
|
||||||
contourpy==1.3.3
|
|
||||||
cycler==0.12.1
|
|
||||||
filelock==3.13.1
|
|
||||||
fonttools==4.59.0
|
|
||||||
fsspec==2024.6.1
|
|
||||||
huggingface-hub==0.34.3
|
|
||||||
idna==3.10
|
|
||||||
Jinja2==3.1.6
|
|
||||||
kiwisolver==1.4.8
|
|
||||||
MarkupSafe==2.1.5
|
|
||||||
matplotlib==3.10.5
|
|
||||||
mpmath==1.3.0
|
|
||||||
networkx==3.3
|
|
||||||
numpy==2.3.2
|
|
||||||
packaging==25.0
|
|
||||||
pillow==11.3.0
|
|
||||||
pyparsing==3.2.3
|
|
||||||
python-dateutil==2.9.0.post0
|
|
||||||
PyYAML==6.0.2
|
|
||||||
requests==2.32.4
|
|
||||||
safetensors==0.5.3
|
|
||||||
setuptools==78.1.1
|
|
||||||
six==1.17.0
|
|
||||||
sympy==1.13.3
|
|
||||||
tokenizers==0.21.4
|
|
||||||
torch==2.7.1+cu126
|
|
||||||
tqdm==4.67.1
|
|
||||||
typing_extensions==4.12.2
|
|
||||||
urllib3==2.5.0
|
|
||||||
wheel==0.45.1
|
|
||||||
@@ -0,0 +1,41 @@
|
|||||||
|
import argparse
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from huggingface_hub import snapshot_download
|
||||||
|
|
||||||
|
PROJECT_ROOT = Path(__file__).resolve().parents[2]
|
||||||
|
DEFAULT_LOCAL_DIR = Path(PROJECT_ROOT, "params")
|
||||||
|
DEFAULT_REPO_ID = "ViperEk/KHAOSZ"
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
description="Download model parameters from HuggingFace"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--repo-id",
|
||||||
|
type=str,
|
||||||
|
default=DEFAULT_REPO_ID,
|
||||||
|
help=f"HuggingFace repo ID (default: {DEFAULT_REPO_ID})",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--local-dir",
|
||||||
|
type=Path,
|
||||||
|
default=DEFAULT_LOCAL_DIR,
|
||||||
|
help=f"Local directory to save model (default: {DEFAULT_LOCAL_DIR})",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--force",
|
||||||
|
action="store_true",
|
||||||
|
help="Force download even if files exist",
|
||||||
|
)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
print(f"Downloading model from {args.repo_id} to {args.local_dir}")
|
||||||
|
|
||||||
|
snapshot_download(
|
||||||
|
repo_id=args.repo_id,
|
||||||
|
local_dir=args.local_dir,
|
||||||
|
force_download=args.force,
|
||||||
|
)
|
||||||
|
|
||||||
|
print("Download complete!")
|
||||||
@@ -0,0 +1,38 @@
|
|||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.inference import InferenceEngine
|
||||||
|
from astrai.model import AutoModel
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
PROJECT_ROOT = Path(__file__).resolve().parents[2]
|
||||||
|
PARAMETER_ROOT = Path(PROJECT_ROOT, "params")
|
||||||
|
|
||||||
|
|
||||||
|
def generate_text():
|
||||||
|
# Load model from pretrained
|
||||||
|
model = AutoModel.from_pretrained(PARAMETER_ROOT)
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(PARAMETER_ROOT)
|
||||||
|
model.to(device="cuda", dtype=torch.bfloat16)
|
||||||
|
|
||||||
|
query = input(">> ")
|
||||||
|
|
||||||
|
engine = InferenceEngine(
|
||||||
|
model=model,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
)
|
||||||
|
response = engine.generate(
|
||||||
|
prompt=query,
|
||||||
|
stream=False,
|
||||||
|
max_tokens=2048,
|
||||||
|
temperature=0.8,
|
||||||
|
top_p=0.95,
|
||||||
|
top_k=50,
|
||||||
|
)
|
||||||
|
|
||||||
|
print(response)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
generate_text()
|
||||||
@@ -0,0 +1,45 @@
|
|||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.inference import InferenceEngine
|
||||||
|
from astrai.model import AutoModel
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
PROJECT_ROOT = Path(__file__).resolve().parents[2]
|
||||||
|
PARAMETER_ROOT = Path(PROJECT_ROOT, "params")
|
||||||
|
|
||||||
|
|
||||||
|
def batch_generate():
|
||||||
|
# Load model using AutoModel
|
||||||
|
model = AutoModel.from_pretrained(PARAMETER_ROOT)
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(PARAMETER_ROOT)
|
||||||
|
model.to(device="cuda", dtype=torch.bfloat16)
|
||||||
|
|
||||||
|
inputs = [
|
||||||
|
"你好",
|
||||||
|
"请问什么是人工智能",
|
||||||
|
"今天天气如何",
|
||||||
|
"我感到焦虑, 请问我应该怎么办",
|
||||||
|
"请问什么是显卡",
|
||||||
|
]
|
||||||
|
|
||||||
|
engine = InferenceEngine(
|
||||||
|
model=model,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
)
|
||||||
|
responses = engine.generate(
|
||||||
|
prompt=inputs,
|
||||||
|
stream=False,
|
||||||
|
max_tokens=2048,
|
||||||
|
temperature=0.8,
|
||||||
|
top_p=0.95,
|
||||||
|
top_k=50,
|
||||||
|
)
|
||||||
|
|
||||||
|
for q, r in zip(inputs, responses):
|
||||||
|
print((q, r))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
batch_generate()
|
||||||
@@ -0,0 +1,50 @@
|
|||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.inference import InferenceEngine
|
||||||
|
from astrai.model import AutoModel
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
PROJECT_ROOT = Path(__file__).resolve().parents[2]
|
||||||
|
PARAMETER_ROOT = Path(PROJECT_ROOT, "params")
|
||||||
|
|
||||||
|
|
||||||
|
def chat():
|
||||||
|
model = AutoModel.from_pretrained(PARAMETER_ROOT)
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(PARAMETER_ROOT)
|
||||||
|
model.to(device="cuda", dtype=torch.bfloat16)
|
||||||
|
|
||||||
|
messages = [{"role": "system", "content": "You are a helpful assistant."}]
|
||||||
|
engine = InferenceEngine(model=model, tokenizer=tokenizer)
|
||||||
|
|
||||||
|
while True:
|
||||||
|
query = input(">> ")
|
||||||
|
if query == "!exit":
|
||||||
|
break
|
||||||
|
|
||||||
|
# Add user message
|
||||||
|
messages.append({"role": "user", "content": query})
|
||||||
|
|
||||||
|
# Generate response
|
||||||
|
full_response = ""
|
||||||
|
prompt = tokenizer.apply_chat_template(messages, tokenize=False)
|
||||||
|
|
||||||
|
for token in engine.generate(
|
||||||
|
prompt=prompt,
|
||||||
|
stream=True,
|
||||||
|
max_tokens=2048,
|
||||||
|
temperature=0.8,
|
||||||
|
top_p=0.95,
|
||||||
|
top_k=50,
|
||||||
|
):
|
||||||
|
print(token, end="", flush=True)
|
||||||
|
full_response += token
|
||||||
|
|
||||||
|
print()
|
||||||
|
# Add assistant response to messages
|
||||||
|
messages.append({"role": "assistant", "content": full_response.strip()})
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
chat()
|
||||||
Executable
+253
@@ -0,0 +1,253 @@
|
|||||||
|
#!/bin/bash
|
||||||
|
|
||||||
|
# AstrAI Docker Script
|
||||||
|
# Build and manage Docker images
|
||||||
|
|
||||||
|
set -e
|
||||||
|
|
||||||
|
# Colors
|
||||||
|
RED='\033[0;31m'
|
||||||
|
GREEN='\033[0;32m'
|
||||||
|
YELLOW='\033[1;33m'
|
||||||
|
BLUE='\033[0;34m'
|
||||||
|
NC='\033[0m' # No Color
|
||||||
|
|
||||||
|
# Default values
|
||||||
|
IMAGE_NAME="astrai"
|
||||||
|
IMAGE_TAG="latest"
|
||||||
|
REGISTRY=""
|
||||||
|
|
||||||
|
# Print colored messages
|
||||||
|
print_info() {
|
||||||
|
echo -e "${BLUE}[INFO]${NC} $1"
|
||||||
|
}
|
||||||
|
|
||||||
|
print_success() {
|
||||||
|
echo -e "${GREEN}[SUCCESS]${NC} $1"
|
||||||
|
}
|
||||||
|
|
||||||
|
print_error() {
|
||||||
|
echo -e "${RED}[ERROR]${NC} $1"
|
||||||
|
}
|
||||||
|
|
||||||
|
print_warning() {
|
||||||
|
echo -e "${YELLOW}[WARNING]${NC} $1"
|
||||||
|
}
|
||||||
|
|
||||||
|
# Check if Docker is installed
|
||||||
|
check_docker() {
|
||||||
|
if ! command -v docker &> /dev/null; then
|
||||||
|
print_error "Docker is not installed"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
print_success "Docker version: $(docker --version)"
|
||||||
|
}
|
||||||
|
|
||||||
|
# Build Docker image
|
||||||
|
build_image() {
|
||||||
|
local dockerfile="${1:-Dockerfile}"
|
||||||
|
local context="${2:-.}"
|
||||||
|
|
||||||
|
if [ ! -f "$dockerfile" ]; then
|
||||||
|
print_error "Dockerfile not found: $dockerfile"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
|
||||||
|
print_info "Building Docker image: ${IMAGE_NAME}:${IMAGE_TAG}"
|
||||||
|
docker build -t "${IMAGE_NAME}:${IMAGE_TAG}" -f "$dockerfile" "$context"
|
||||||
|
print_success "Image built successfully"
|
||||||
|
}
|
||||||
|
|
||||||
|
# Run container
|
||||||
|
run_container() {
|
||||||
|
local port="${1:-8000}"
|
||||||
|
local gpu="${2:-false}"
|
||||||
|
|
||||||
|
print_info "Running container on port $port..."
|
||||||
|
|
||||||
|
if [ "$gpu" = true ]; then
|
||||||
|
docker run --gpus all -p "${port}:8000" "${IMAGE_NAME}:${IMAGE_TAG}"
|
||||||
|
else
|
||||||
|
docker run -p "${port}:8000" "${IMAGE_NAME}:${IMAGE_TAG}"
|
||||||
|
fi
|
||||||
|
}
|
||||||
|
|
||||||
|
# Push image to registry
|
||||||
|
push_image() {
|
||||||
|
if [ -z "$REGISTRY" ]; then
|
||||||
|
print_error "Registry not set. Use --registry option"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
|
||||||
|
local full_tag="${REGISTRY}/${IMAGE_NAME}:${IMAGE_TAG}"
|
||||||
|
print_info "Tagging image: ${full_tag}"
|
||||||
|
docker tag "${IMAGE_NAME}:${IMAGE_TAG}" "$full_tag"
|
||||||
|
|
||||||
|
print_info "Pushing image to registry..."
|
||||||
|
docker push "$full_tag"
|
||||||
|
print_success "Image pushed successfully"
|
||||||
|
}
|
||||||
|
|
||||||
|
# Remove image
|
||||||
|
remove_image() {
|
||||||
|
print_info "Removing image: ${IMAGE_NAME}:${IMAGE_TAG}"
|
||||||
|
docker rmi "${IMAGE_NAME}:${IMAGE_TAG}" 2>/dev/null || print_warning "Image not found"
|
||||||
|
print_success "Image removed"
|
||||||
|
}
|
||||||
|
|
||||||
|
# Show image info
|
||||||
|
show_info() {
|
||||||
|
print_info "Image information:"
|
||||||
|
docker images "${IMAGE_NAME}"
|
||||||
|
}
|
||||||
|
|
||||||
|
# Show logs
|
||||||
|
show_logs() {
|
||||||
|
local container_id="$1"
|
||||||
|
if [ -z "$container_id" ]; then
|
||||||
|
print_error "Container ID required"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
docker logs "$container_id"
|
||||||
|
}
|
||||||
|
|
||||||
|
# Main function
|
||||||
|
main() {
|
||||||
|
echo "========================================"
|
||||||
|
echo " AstrAI Docker Management"
|
||||||
|
echo "========================================"
|
||||||
|
echo ""
|
||||||
|
|
||||||
|
COMMAND=""
|
||||||
|
DOCKERFILE="Dockerfile"
|
||||||
|
CONTEXT="."
|
||||||
|
PORT="8000"
|
||||||
|
GPU=false
|
||||||
|
|
||||||
|
# Parse arguments
|
||||||
|
while [[ $# -gt 0 ]]; do
|
||||||
|
case $1 in
|
||||||
|
build)
|
||||||
|
COMMAND="build"
|
||||||
|
shift
|
||||||
|
;;
|
||||||
|
run)
|
||||||
|
COMMAND="run"
|
||||||
|
shift
|
||||||
|
;;
|
||||||
|
push)
|
||||||
|
COMMAND="push"
|
||||||
|
shift
|
||||||
|
;;
|
||||||
|
remove|rm)
|
||||||
|
COMMAND="remove"
|
||||||
|
shift
|
||||||
|
;;
|
||||||
|
info)
|
||||||
|
COMMAND="info"
|
||||||
|
shift
|
||||||
|
;;
|
||||||
|
logs)
|
||||||
|
COMMAND="logs"
|
||||||
|
shift
|
||||||
|
;;
|
||||||
|
--image)
|
||||||
|
IMAGE_NAME="$2"
|
||||||
|
shift 2
|
||||||
|
;;
|
||||||
|
--tag)
|
||||||
|
IMAGE_TAG="$2"
|
||||||
|
shift 2
|
||||||
|
;;
|
||||||
|
--registry)
|
||||||
|
REGISTRY="$2"
|
||||||
|
shift 2
|
||||||
|
;;
|
||||||
|
--dockerfile)
|
||||||
|
DOCKERFILE="$2"
|
||||||
|
shift 2
|
||||||
|
;;
|
||||||
|
--context)
|
||||||
|
CONTEXT="$2"
|
||||||
|
shift 2
|
||||||
|
;;
|
||||||
|
--port)
|
||||||
|
PORT="$2"
|
||||||
|
shift 2
|
||||||
|
;;
|
||||||
|
--gpu)
|
||||||
|
GPU=true
|
||||||
|
shift
|
||||||
|
;;
|
||||||
|
--help)
|
||||||
|
echo "Usage: $0 <command> [options]"
|
||||||
|
echo ""
|
||||||
|
echo "Commands:"
|
||||||
|
echo " build Build Docker image"
|
||||||
|
echo " run Run container"
|
||||||
|
echo " push Push image to registry"
|
||||||
|
echo " remove Remove image"
|
||||||
|
echo " info Show image information"
|
||||||
|
echo " logs Show container logs"
|
||||||
|
echo ""
|
||||||
|
echo "Options:"
|
||||||
|
echo " --image NAME Image name (default: astrai)"
|
||||||
|
echo " --tag TAG Image tag (default: latest)"
|
||||||
|
echo " --registry URL Registry URL for push"
|
||||||
|
echo " --dockerfile FILE Dockerfile path (default: Dockerfile)"
|
||||||
|
echo " --context PATH Build context (default: .)"
|
||||||
|
echo " --port PORT Port for run (default: 8000)"
|
||||||
|
echo " --gpu Enable GPU support"
|
||||||
|
echo " --help Show this help message"
|
||||||
|
echo ""
|
||||||
|
echo "Examples:"
|
||||||
|
echo " $0 build"
|
||||||
|
echo " $0 build --tag v1.0.0"
|
||||||
|
echo " $0 run --port 8080"
|
||||||
|
echo " $0 run --gpu"
|
||||||
|
echo " $0 push --registry ghcr.io/username"
|
||||||
|
exit 0
|
||||||
|
;;
|
||||||
|
*)
|
||||||
|
if [ -z "$COMMAND" ]; then
|
||||||
|
print_error "Unknown command: $1"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
shift
|
||||||
|
;;
|
||||||
|
esac
|
||||||
|
done
|
||||||
|
|
||||||
|
check_docker
|
||||||
|
|
||||||
|
case "$COMMAND" in
|
||||||
|
build)
|
||||||
|
build_image "$DOCKERFILE" "$CONTEXT"
|
||||||
|
;;
|
||||||
|
run)
|
||||||
|
run_container "$PORT" "$GPU"
|
||||||
|
;;
|
||||||
|
push)
|
||||||
|
push_image
|
||||||
|
;;
|
||||||
|
remove)
|
||||||
|
remove_image
|
||||||
|
;;
|
||||||
|
info)
|
||||||
|
show_info
|
||||||
|
;;
|
||||||
|
logs)
|
||||||
|
show_logs "$2"
|
||||||
|
;;
|
||||||
|
"")
|
||||||
|
print_error "No command specified. Use --help for usage"
|
||||||
|
exit 1
|
||||||
|
;;
|
||||||
|
*)
|
||||||
|
print_error "Unknown command: $COMMAND"
|
||||||
|
exit 1
|
||||||
|
;;
|
||||||
|
esac
|
||||||
|
}
|
||||||
|
|
||||||
|
main "$@"
|
||||||
@@ -1,14 +0,0 @@
|
|||||||
import os
|
|
||||||
from huggingface_hub import snapshot_download
|
|
||||||
|
|
||||||
|
|
||||||
PROJECT_ROOT = os.path.dirname(
|
|
||||||
os.path.dirname(os.path.abspath(__file__)))
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
snapshot_download(
|
|
||||||
repo_id="ViperEk/KHAOSZ",
|
|
||||||
local_dir=os.path.join(PROJECT_ROOT, "params"),
|
|
||||||
force_download=True
|
|
||||||
)
|
|
||||||
@@ -1,27 +0,0 @@
|
|||||||
import os
|
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
|
|
||||||
PROJECT_ROOT = os.path.dirname(
|
|
||||||
os.path.dirname(os.path.abspath(__file__)))
|
|
||||||
|
|
||||||
def generate_text():
|
|
||||||
model_dir = os.path.join(PROJECT_ROOT, "params")
|
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
|
|
||||||
query = input(">> ")
|
|
||||||
|
|
||||||
response = model.text_generate(
|
|
||||||
query=query,
|
|
||||||
temperature=0.8,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50
|
|
||||||
)
|
|
||||||
|
|
||||||
print(response)
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
generate_text()
|
|
||||||
@@ -1,25 +0,0 @@
|
|||||||
import os
|
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
|
|
||||||
PROJECT_ROOT = os.path.dirname(
|
|
||||||
os.path.dirname(os.path.abspath(__file__)))
|
|
||||||
|
|
||||||
def batch_generate():
|
|
||||||
model_dir = os.path.join(PROJECT_ROOT, "params")
|
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
inputs = ["你好", "请问什么是人工智能", "今天天气如何", "我感到焦虑, 请问我应该怎么办", "请问什么是显卡"]
|
|
||||||
|
|
||||||
responses = model.batch_generate(
|
|
||||||
queries=inputs,
|
|
||||||
temperature=0.8,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50
|
|
||||||
)
|
|
||||||
|
|
||||||
for q, r in zip(inputs, responses):
|
|
||||||
print((q, r))
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
batch_generate()
|
|
||||||
@@ -1,42 +0,0 @@
|
|||||||
import os
|
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz, SemanticTextSplitter, Retriever
|
|
||||||
|
|
||||||
|
|
||||||
PROJECT_ROOT = os.path.dirname(
|
|
||||||
os.path.dirname(os.path.abspath(__file__)))
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
model_dir = os.path.join(PROJECT_ROOT, "params")
|
|
||||||
context_path = os.path.join(PROJECT_ROOT, "README.md")
|
|
||||||
|
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
spliter = SemanticTextSplitter(model.encode)
|
|
||||||
retriever = Retriever()
|
|
||||||
text = open(context_path, "r", encoding="utf-8").read()
|
|
||||||
|
|
||||||
res = spliter.split(text, threshold=0.8, window_size=1)
|
|
||||||
# print(("\n" + "+"*100 + "\n").join(res))
|
|
||||||
|
|
||||||
res_embs = model.encode(res)
|
|
||||||
for sentence, emb in zip(res, res_embs):
|
|
||||||
retriever.add_vector(sentence, emb)
|
|
||||||
|
|
||||||
retrive_top_k = 5
|
|
||||||
query = "作者设计了一个怎样的模型"
|
|
||||||
emb_query = model.encode(query)
|
|
||||||
retrieved = retriever.retrieve(emb_query, retrive_top_k)
|
|
||||||
|
|
||||||
retrive_response = model.retrieve_generate(
|
|
||||||
retrieved=retrieved,
|
|
||||||
query=query,
|
|
||||||
temperature=0.8,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50
|
|
||||||
)
|
|
||||||
|
|
||||||
print("retrive content:")
|
|
||||||
print("\n".join([f"{idx + 1}. " + text for idx, (text, _) in enumerate(retrieved)]))
|
|
||||||
|
|
||||||
print("\n\nretrive generate:")
|
|
||||||
print(retrive_response)
|
|
||||||
Executable
+159
@@ -0,0 +1,159 @@
|
|||||||
|
#!/bin/bash
|
||||||
|
|
||||||
|
# AstrAI Pre-commit Check Script
|
||||||
|
# Runs code format check and tests before committing
|
||||||
|
|
||||||
|
set -e
|
||||||
|
|
||||||
|
# Colors
|
||||||
|
RED='\033[0;31m'
|
||||||
|
GREEN='\033[0;32m'
|
||||||
|
YELLOW='\033[1;33m'
|
||||||
|
NC='\033[0m' # No Color
|
||||||
|
|
||||||
|
# Print colored messages
|
||||||
|
print_info() {
|
||||||
|
echo -e "${YELLOW}[INFO]${NC} $1"
|
||||||
|
}
|
||||||
|
|
||||||
|
print_success() {
|
||||||
|
echo -e "${GREEN}[SUCCESS]${NC} $1"
|
||||||
|
}
|
||||||
|
|
||||||
|
print_error() {
|
||||||
|
echo -e "${RED}[ERROR]${NC} $1"
|
||||||
|
}
|
||||||
|
|
||||||
|
# Check if in project root directory
|
||||||
|
check_project_root() {
|
||||||
|
if [ ! -f "pyproject.toml" ]; then
|
||||||
|
print_error "Please run this script from the project root directory"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
}
|
||||||
|
|
||||||
|
# Check if Python is installed
|
||||||
|
check_python() {
|
||||||
|
if ! command -v python &> /dev/null; then
|
||||||
|
print_error "Python is not installed"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
print_info "Python version: $(python --version)"
|
||||||
|
}
|
||||||
|
|
||||||
|
# Install development dependencies
|
||||||
|
install_dependencies() {
|
||||||
|
print_info "Installing development dependencies..."
|
||||||
|
pip install --upgrade pip
|
||||||
|
pip install .[dev]
|
||||||
|
print_success "Dependencies installed"
|
||||||
|
}
|
||||||
|
|
||||||
|
# Run code format check
|
||||||
|
run_lint() {
|
||||||
|
print_info "Running code format check (ruff format)..."
|
||||||
|
if ruff format --check .; then
|
||||||
|
print_success "Code format check passed"
|
||||||
|
else
|
||||||
|
print_error "Code format check failed. Please run 'ruff format .' to fix formatting issues"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
}
|
||||||
|
|
||||||
|
# Run code style check (linter - import sorting)
|
||||||
|
run_ruff_lint_import() {
|
||||||
|
print_info "Running import sorting check (ruff check --select I)..."
|
||||||
|
if ruff check . --select I; then
|
||||||
|
print_success "Import sorting check passed"
|
||||||
|
else
|
||||||
|
print_error "Import sorting check failed. Please run 'ruff check --select I --fix .' to fix import issues"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
}
|
||||||
|
|
||||||
|
# Run tests
|
||||||
|
run_tests() {
|
||||||
|
print_info "Running tests..."
|
||||||
|
if python -m pytest tests/ -v; then
|
||||||
|
print_success "All tests passed"
|
||||||
|
else
|
||||||
|
print_error "Tests failed"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
}
|
||||||
|
|
||||||
|
# Main function
|
||||||
|
main() {
|
||||||
|
echo "========================================"
|
||||||
|
echo " AstrAI Pre-commit Check Script"
|
||||||
|
echo "========================================"
|
||||||
|
echo ""
|
||||||
|
|
||||||
|
check_project_root
|
||||||
|
check_python
|
||||||
|
|
||||||
|
# Parse arguments
|
||||||
|
SKIP_DEPS=false
|
||||||
|
SKIP_LINT=false
|
||||||
|
SKIP_TESTS=false
|
||||||
|
|
||||||
|
while [[ $# -gt 0 ]]; do
|
||||||
|
case $1 in
|
||||||
|
--skip-deps)
|
||||||
|
SKIP_DEPS=true
|
||||||
|
shift
|
||||||
|
;;
|
||||||
|
--skip-lint)
|
||||||
|
SKIP_LINT=true
|
||||||
|
shift
|
||||||
|
;;
|
||||||
|
--skip-tests)
|
||||||
|
SKIP_TESTS=true
|
||||||
|
shift
|
||||||
|
;;
|
||||||
|
--help)
|
||||||
|
echo "Usage: $0 [options]"
|
||||||
|
echo ""
|
||||||
|
echo "Options:"
|
||||||
|
echo " --skip-deps Skip dependency installation"
|
||||||
|
echo " --skip-lint Skip code checks"
|
||||||
|
echo " --skip-tests Skip tests"
|
||||||
|
echo " --help Show this help message"
|
||||||
|
exit 0
|
||||||
|
;;
|
||||||
|
*)
|
||||||
|
print_error "Unknown option: $1"
|
||||||
|
exit 1
|
||||||
|
;;
|
||||||
|
esac
|
||||||
|
done
|
||||||
|
|
||||||
|
# Install dependencies
|
||||||
|
if [ "$SKIP_DEPS" = false ]; then
|
||||||
|
install_dependencies
|
||||||
|
else
|
||||||
|
print_info "Skipping dependency installation"
|
||||||
|
fi
|
||||||
|
|
||||||
|
# Run code checks
|
||||||
|
if [ "$SKIP_LINT" = false ]; then
|
||||||
|
run_lint
|
||||||
|
run_ruff_lint_import
|
||||||
|
else
|
||||||
|
print_info "Skipping code checks"
|
||||||
|
fi
|
||||||
|
|
||||||
|
# Run tests
|
||||||
|
if [ "$SKIP_TESTS" = false ]; then
|
||||||
|
run_tests
|
||||||
|
else
|
||||||
|
print_info "Skipping tests"
|
||||||
|
fi
|
||||||
|
|
||||||
|
echo ""
|
||||||
|
echo "========================================"
|
||||||
|
print_success "All checks passed! Ready to commit."
|
||||||
|
echo "========================================"
|
||||||
|
}
|
||||||
|
|
||||||
|
main "$@"
|
||||||
@@ -1,32 +0,0 @@
|
|||||||
import os
|
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
|
|
||||||
PROJECT_ROOT = os.path.dirname(
|
|
||||||
os.path.dirname(os.path.abspath(__file__)))
|
|
||||||
|
|
||||||
def chat():
|
|
||||||
model_dir = os.path.join(PROJECT_ROOT, "params")
|
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
|
|
||||||
history = []
|
|
||||||
while True:
|
|
||||||
query = input(">> ")
|
|
||||||
if query == "!exit":
|
|
||||||
break
|
|
||||||
|
|
||||||
response_size = 0
|
|
||||||
for response, history in model.stream_generate(
|
|
||||||
query=query,
|
|
||||||
history=history,
|
|
||||||
temperature=0.8,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50
|
|
||||||
):
|
|
||||||
print(response[response_size:], end="", flush=True)
|
|
||||||
response_size = len(response)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
chat()
|
|
||||||
@@ -0,0 +1,241 @@
|
|||||||
|
"""Benchmark Transformer with PagedCache (replaces old persistent_key_values)."""
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any, Dict
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.config import ModelConfig
|
||||||
|
from astrai.inference.cache import PagedCache
|
||||||
|
from astrai.model.transformer import Transformer
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class BenchmarkResult:
|
||||||
|
total_tokens: int
|
||||||
|
total_time: float
|
||||||
|
tokens_per_second: float
|
||||||
|
metadata: Dict[str, Any]
|
||||||
|
|
||||||
|
|
||||||
|
class GenerationBenchmark:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: ModelConfig,
|
||||||
|
device: str = "cuda",
|
||||||
|
dtype: torch.dtype = torch.bfloat16,
|
||||||
|
page_size: int = 128,
|
||||||
|
):
|
||||||
|
self.config = config
|
||||||
|
self.device = device
|
||||||
|
self.dtype = dtype
|
||||||
|
self.model = Transformer(config).to(device=device, dtype=dtype)
|
||||||
|
self.model.eval()
|
||||||
|
head_dim = config.dim // config.n_heads
|
||||||
|
n_pages = (config.max_len * 4 + page_size - 1) // page_size
|
||||||
|
self._page_cache = PagedCache(
|
||||||
|
config.n_layers,
|
||||||
|
n_pages,
|
||||||
|
page_size,
|
||||||
|
config.n_kv_heads,
|
||||||
|
head_dim,
|
||||||
|
device,
|
||||||
|
dtype,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _prepare_inputs(self, batch_size: int, prompt_length: int, total_length: int):
|
||||||
|
prompt_ids = torch.randint(
|
||||||
|
low=0,
|
||||||
|
high=self.config.vocab_size,
|
||||||
|
size=(batch_size, prompt_length),
|
||||||
|
device=self.device,
|
||||||
|
dtype=torch.long,
|
||||||
|
)
|
||||||
|
gen_ids = torch.randint(
|
||||||
|
low=0,
|
||||||
|
high=self.config.vocab_size,
|
||||||
|
size=(batch_size, total_length - prompt_length),
|
||||||
|
device=self.device,
|
||||||
|
dtype=torch.long,
|
||||||
|
)
|
||||||
|
return prompt_ids, gen_ids
|
||||||
|
|
||||||
|
def _make_mask(self, batch_size: int, seq_len: int) -> Tensor:
|
||||||
|
return torch.ones(batch_size, seq_len, dtype=torch.bool, device=self.device)
|
||||||
|
|
||||||
|
@torch.inference_mode()
|
||||||
|
def run_prefill_benchmark(
|
||||||
|
self,
|
||||||
|
batch_size: int = 1,
|
||||||
|
prompt_length: int = 512,
|
||||||
|
num_trials: int = 10,
|
||||||
|
) -> BenchmarkResult:
|
||||||
|
for _ in range(3):
|
||||||
|
prompt_ids, _ = self._prepare_inputs(
|
||||||
|
batch_size, prompt_length, prompt_length
|
||||||
|
)
|
||||||
|
_ = self.model(prompt_ids)
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
|
total_time = 0.0
|
||||||
|
total_tokens = batch_size * prompt_length * num_trials
|
||||||
|
|
||||||
|
for trial in range(num_trials):
|
||||||
|
prompt_ids, _ = self._prepare_inputs(
|
||||||
|
batch_size, prompt_length, prompt_length
|
||||||
|
)
|
||||||
|
start = torch.cuda.Event(enable_timing=True)
|
||||||
|
end = torch.cuda.Event(enable_timing=True)
|
||||||
|
|
||||||
|
start.record()
|
||||||
|
_ = self.model(prompt_ids)
|
||||||
|
end.record()
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
|
trial_time = start.elapsed_time(end) / 1000
|
||||||
|
total_time += trial_time
|
||||||
|
|
||||||
|
print(
|
||||||
|
f" Trial {trial + 1}/{num_trials}: {prompt_length} tokens in {trial_time:.3f}s "
|
||||||
|
f"({prompt_length / trial_time:.1f} tok/s)"
|
||||||
|
)
|
||||||
|
|
||||||
|
return BenchmarkResult(
|
||||||
|
total_tokens=total_tokens,
|
||||||
|
total_time=total_time,
|
||||||
|
tokens_per_second=total_tokens / total_time,
|
||||||
|
metadata={
|
||||||
|
"benchmark_type": "prefill",
|
||||||
|
"batch_size": batch_size,
|
||||||
|
"prompt_length": prompt_length,
|
||||||
|
"dtype": str(self.dtype),
|
||||||
|
"device": self.device,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
@torch.inference_mode()
|
||||||
|
def run_decoding_benchmark(
|
||||||
|
self,
|
||||||
|
batch_size: int = 1,
|
||||||
|
prompt_length: int = 512,
|
||||||
|
gen_length: int = 128,
|
||||||
|
num_trials: int = 5,
|
||||||
|
) -> BenchmarkResult:
|
||||||
|
total_time = 0.0
|
||||||
|
total_tokens = batch_size * gen_length * num_trials
|
||||||
|
page_size = self._page_cache.page_size
|
||||||
|
|
||||||
|
for trial in range(num_trials):
|
||||||
|
prompt_ids, gen_ids = self._prepare_inputs(
|
||||||
|
batch_size,
|
||||||
|
prompt_length,
|
||||||
|
prompt_length + gen_length,
|
||||||
|
)
|
||||||
|
|
||||||
|
n_pages = (prompt_length + gen_length + page_size - 1) // page_size
|
||||||
|
pages = self._page_cache.alloc_n(n_pages * batch_size)
|
||||||
|
page_table = torch.tensor(
|
||||||
|
[pages[i * n_pages : (i + 1) * n_pages] for i in range(batch_size)],
|
||||||
|
dtype=torch.long,
|
||||||
|
device=self.device,
|
||||||
|
)
|
||||||
|
|
||||||
|
cv = self._page_cache.bind(page_table, total_len=prompt_length)
|
||||||
|
_ = self.model(
|
||||||
|
prompt_ids,
|
||||||
|
paged_cache=cv,
|
||||||
|
start_pos=0,
|
||||||
|
input_mask=self._make_mask(batch_size, prompt_length),
|
||||||
|
)
|
||||||
|
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
|
start = torch.cuda.Event(enable_timing=True)
|
||||||
|
end = torch.cuda.Event(enable_timing=True)
|
||||||
|
|
||||||
|
start.record()
|
||||||
|
current_pos = prompt_length
|
||||||
|
for i in range(gen_length):
|
||||||
|
input_token = gen_ids[:, i : i + 1]
|
||||||
|
cv = self._page_cache.bind(page_table, total_len=current_pos + 1)
|
||||||
|
_ = self.model(
|
||||||
|
input_token,
|
||||||
|
paged_cache=cv,
|
||||||
|
start_pos=current_pos,
|
||||||
|
input_mask=self._make_mask(batch_size, 1),
|
||||||
|
)
|
||||||
|
current_pos += 1
|
||||||
|
end.record()
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
|
trial_time = start.elapsed_time(end) / 1000
|
||||||
|
total_time += trial_time
|
||||||
|
|
||||||
|
for idx in pages:
|
||||||
|
self._page_cache.free(idx)
|
||||||
|
|
||||||
|
print(
|
||||||
|
f" Trial {trial + 1}/{num_trials}: {gen_length} tokens in {trial_time:.3f}s "
|
||||||
|
f"({gen_length / trial_time:.1f} tok/s)"
|
||||||
|
)
|
||||||
|
|
||||||
|
return BenchmarkResult(
|
||||||
|
total_tokens=total_tokens,
|
||||||
|
total_time=total_time,
|
||||||
|
tokens_per_second=total_tokens / total_time,
|
||||||
|
metadata={
|
||||||
|
"benchmark_type": "decoding",
|
||||||
|
"batch_size": batch_size,
|
||||||
|
"prompt_length": prompt_length,
|
||||||
|
"gen_length": gen_length,
|
||||||
|
"dtype": str(self.dtype),
|
||||||
|
"device": self.device,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def print_benchmark_result(result: BenchmarkResult):
|
||||||
|
btype = result.metadata["benchmark_type"]
|
||||||
|
print(f"\n{' ' + btype.upper() + ' Benchmark ':-^80}")
|
||||||
|
print(f"Total Tokens Processed: {result.total_tokens:,}")
|
||||||
|
print(f"Time Consumed: {result.total_time:.3f}s")
|
||||||
|
print(f"Throughput: {result.tokens_per_second:,.1f} tok/s")
|
||||||
|
for k, v in result.metadata.items():
|
||||||
|
if k != "benchmark_type":
|
||||||
|
print(f"{k.replace('_', ' ').title()}: {v}")
|
||||||
|
print("-" * 80)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
config = ModelConfig(
|
||||||
|
vocab_size=10000,
|
||||||
|
dim=1536,
|
||||||
|
n_heads=24,
|
||||||
|
n_kv_heads=4,
|
||||||
|
dim_ffn=6912,
|
||||||
|
max_len=2048,
|
||||||
|
n_layers=24,
|
||||||
|
norm_eps=1e-5,
|
||||||
|
)
|
||||||
|
|
||||||
|
benchmark = GenerationBenchmark(config)
|
||||||
|
|
||||||
|
print("=" * 80)
|
||||||
|
print("Running Transformer Generation Benchmark (PagedCache)")
|
||||||
|
print("=" * 80)
|
||||||
|
|
||||||
|
prefill_result = benchmark.run_prefill_benchmark(
|
||||||
|
batch_size=4,
|
||||||
|
prompt_length=512,
|
||||||
|
num_trials=5,
|
||||||
|
)
|
||||||
|
print_benchmark_result(prefill_result)
|
||||||
|
|
||||||
|
gen_result = benchmark.run_decoding_benchmark(
|
||||||
|
batch_size=4,
|
||||||
|
prompt_length=512,
|
||||||
|
gen_length=128,
|
||||||
|
num_trials=5,
|
||||||
|
)
|
||||||
|
print_benchmark_result(gen_result)
|
||||||
@@ -0,0 +1,129 @@
|
|||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.inference import InferenceEngine
|
||||||
|
from astrai.model import AutoModel
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
|
def processor(
|
||||||
|
param_path: str,
|
||||||
|
input_json_file: str,
|
||||||
|
output_json_file: str,
|
||||||
|
temperature: float,
|
||||||
|
top_k: int,
|
||||||
|
top_p: float,
|
||||||
|
question_key: str,
|
||||||
|
response_key: str,
|
||||||
|
max_tokens: int,
|
||||||
|
):
|
||||||
|
# Load model and tokenizer
|
||||||
|
model = AutoModel.from_pretrained(param_path)
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||||
|
model.to(device="cuda", dtype=torch.bfloat16)
|
||||||
|
|
||||||
|
# Create inference engine
|
||||||
|
engine = InferenceEngine(model=model, tokenizer=tokenizer)
|
||||||
|
|
||||||
|
with open(input_json_file, "r", encoding="utf-8") as f:
|
||||||
|
input_data = [json.loads(line) for line in f]
|
||||||
|
|
||||||
|
# Check input format: chat messages or raw text
|
||||||
|
if input_data and "messages" in input_data[0]:
|
||||||
|
# Chat format: [{"messages": [...]}]
|
||||||
|
prompts = [
|
||||||
|
tokenizer.apply_chat_template(item["messages"], tokenize=False)
|
||||||
|
for item in input_data
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
# Raw text format: [{"question": "..."}]
|
||||||
|
prompts = [item[question_key] for item in input_data]
|
||||||
|
|
||||||
|
# Use provided max_tokens or default to model config max_len
|
||||||
|
if max_tokens is None:
|
||||||
|
max_tokens = model.config.max_len
|
||||||
|
|
||||||
|
# Generate responses (batch)
|
||||||
|
responses = engine.generate(
|
||||||
|
prompt=prompts,
|
||||||
|
stream=False,
|
||||||
|
max_tokens=max_tokens,
|
||||||
|
temperature=temperature,
|
||||||
|
top_p=top_p,
|
||||||
|
top_k=top_k,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Write results
|
||||||
|
with open(output_json_file, "w", encoding="utf-8") as f:
|
||||||
|
for prompt, response in zip(prompts, responses):
|
||||||
|
if input_data and "messages" in input_data[0]:
|
||||||
|
output_item = {"response": response}
|
||||||
|
else:
|
||||||
|
output_item = {question_key: prompt, response_key: response}
|
||||||
|
f.write(json.dumps(output_item, ensure_ascii=False) + "\n")
|
||||||
|
|
||||||
|
# Cleanup
|
||||||
|
engine.shutdown()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
parser = argparse.ArgumentParser(description="Run generate with a Khaosz model.")
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--param_path", type=str, required=True, help="Path to the model directory."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--input_json_file",
|
||||||
|
type=str,
|
||||||
|
required=True,
|
||||||
|
help="Path to the input JSONL file.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--output_json_file",
|
||||||
|
type=str,
|
||||||
|
required=True,
|
||||||
|
help="Path to the output JSONL file.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--question_key",
|
||||||
|
type=str,
|
||||||
|
default="question",
|
||||||
|
help="Key for the question in the input JSON.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--response_key",
|
||||||
|
type=str,
|
||||||
|
default="response",
|
||||||
|
help="Key for the response in the output JSON.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--temperature",
|
||||||
|
type=float,
|
||||||
|
default=0.60,
|
||||||
|
help="Temperature for generating responses.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--top_k", type=int, default=30, help="Top-k value for generating responses."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--top_p",
|
||||||
|
type=float,
|
||||||
|
default=0.95,
|
||||||
|
help="Top-p value for generating responses.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--batch_size", type=int, default=1, help="Batch size for generating responses."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--max_tokens",
|
||||||
|
type=int,
|
||||||
|
default=2048,
|
||||||
|
help="Maximum tokens to generate (default: model config max_len).",
|
||||||
|
)
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
with torch.inference_mode():
|
||||||
|
processor(**vars(args))
|
||||||
@@ -0,0 +1,111 @@
|
|||||||
|
import argparse
|
||||||
|
import json
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
import tqdm
|
||||||
|
|
||||||
|
from astrai.model import AutoModel
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
|
def process_file(
|
||||||
|
model_dir: str, input_file: str, output_file: str, batch_size: int, text_key: str
|
||||||
|
):
|
||||||
|
# Load model and tokenizer
|
||||||
|
model = AutoModel.from_pretrained(model_dir)
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(model_dir)
|
||||||
|
model.to(device="cuda", dtype=torch.bfloat16)
|
||||||
|
|
||||||
|
with open(input_file, "r", encoding="utf-8") as f:
|
||||||
|
input_data = [json.loads(line) for line in f]
|
||||||
|
|
||||||
|
texts = [item[text_key] for item in input_data]
|
||||||
|
|
||||||
|
# Encode all texts
|
||||||
|
print(f"Encoding {len(texts)} texts...")
|
||||||
|
encoded_texts = [tokenizer.encode(text) for text in texts]
|
||||||
|
|
||||||
|
output_data = []
|
||||||
|
total_batches = (len(encoded_texts) + batch_size - 1) // batch_size
|
||||||
|
|
||||||
|
for i in tqdm.tqdm(
|
||||||
|
range(0, len(encoded_texts), batch_size),
|
||||||
|
total=total_batches,
|
||||||
|
desc="Computing perplexity",
|
||||||
|
):
|
||||||
|
batch_encoded = encoded_texts[i : i + batch_size]
|
||||||
|
batch_texts = texts[i : i + batch_size]
|
||||||
|
|
||||||
|
# Find max length in batch and pad
|
||||||
|
max_len = max(len(seq) for seq in batch_encoded)
|
||||||
|
padded_ids = []
|
||||||
|
masks = []
|
||||||
|
|
||||||
|
for seq in batch_encoded:
|
||||||
|
pad_len = max_len - len(seq)
|
||||||
|
padded_seq = [tokenizer.pad_id] * pad_len + seq
|
||||||
|
mask = [False] * pad_len + [True] * len(seq)
|
||||||
|
padded_ids.append(padded_seq)
|
||||||
|
masks.append(mask)
|
||||||
|
|
||||||
|
# Convert to tensors
|
||||||
|
input_ids = torch.tensor(padded_ids, device="cuda", dtype=torch.long)
|
||||||
|
input_mask = torch.tensor(masks, device="cuda", dtype=torch.bool)
|
||||||
|
|
||||||
|
# Compute perplexity
|
||||||
|
output = model(input_ids, input_mask=input_mask)
|
||||||
|
logits = output["logits"]
|
||||||
|
|
||||||
|
# Shift for causal language modeling
|
||||||
|
shifted_logits = logits[:, :-1, :] # [batch_size, seq_len-1, vocab_size]
|
||||||
|
shifted_input_ids = input_ids[:, 1:] # [batch_size, seq_len-1]
|
||||||
|
shifted_mask = input_mask[:, 1:] # [batch_size, seq_len-1]
|
||||||
|
|
||||||
|
# Compute cross entropy loss
|
||||||
|
loss = F.cross_entropy(
|
||||||
|
shifted_logits.flatten(0, 1),
|
||||||
|
shifted_input_ids.flatten(0, 1),
|
||||||
|
reduction="none",
|
||||||
|
)
|
||||||
|
|
||||||
|
loss = loss.view(shifted_input_ids.shape) # [batch_size, seq_len-1]
|
||||||
|
loss = loss * shifted_mask
|
||||||
|
sentence_loss = loss.sum(dim=1) / shifted_mask.sum(dim=1).clamp(min=1)
|
||||||
|
perplexity = torch.exp(sentence_loss) # [batch_size]
|
||||||
|
|
||||||
|
for text, ppl in zip(batch_texts, perplexity):
|
||||||
|
output_data.append({text_key: text, "ppl": float(ppl.item())})
|
||||||
|
|
||||||
|
# Write results
|
||||||
|
with open(output_file, "w", encoding="utf-8") as f:
|
||||||
|
for item in output_data:
|
||||||
|
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
||||||
|
|
||||||
|
print(f"Perplexity computation complete. Results saved to {output_file}")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
parser = argparse.ArgumentParser(description="Run perplexity with a Khaosz model.")
|
||||||
|
parser.add_argument(
|
||||||
|
"--model_dir", type=str, required=True, help="Path to the model directory."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--input_file", type=str, required=True, help="Path to the input file."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--output_file", type=str, required=True, help="Path to the output file."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--batch_size", type=int, default=4, help="Batch size for evaluation."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--text_key",
|
||||||
|
type=str,
|
||||||
|
default="text",
|
||||||
|
help="Key for the text field in the input data.",
|
||||||
|
)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
with torch.inference_mode():
|
||||||
|
process_file(**vars(args))
|
||||||
@@ -0,0 +1,72 @@
|
|||||||
|
import argparse
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.inference.server import run_server
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
parser = argparse.ArgumentParser(description="Start AstrAI inference HTTP server")
|
||||||
|
parser.add_argument(
|
||||||
|
"--host", default="0.0.0.0", help="Host address (default: 0.0.0.0)"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--port", type=int, default=8000, help="Port number (default: 8000)"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--reload", action="store_true", help="Enable auto-reload for development"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--param-path",
|
||||||
|
type=Path,
|
||||||
|
default=None,
|
||||||
|
help="Path to model parameters (default: project_root/params)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--device",
|
||||||
|
type=str,
|
||||||
|
default="cuda",
|
||||||
|
help="Device to load model on (default: cuda)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--dtype",
|
||||||
|
type=str,
|
||||||
|
default="bfloat16",
|
||||||
|
choices=["bfloat16", "float16", "float32"],
|
||||||
|
help="Data type for model weights (default: bfloat16)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--max_batch_size",
|
||||||
|
type=int,
|
||||||
|
default=16,
|
||||||
|
help="Maximum batch size for continuous batching (default: 16)",
|
||||||
|
)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
# Convert dtype string to torch dtype
|
||||||
|
dtype_map = {
|
||||||
|
"bfloat16": torch.bfloat16,
|
||||||
|
"float16": torch.float16,
|
||||||
|
"float32": torch.float32,
|
||||||
|
}
|
||||||
|
dtype = dtype_map[args.dtype]
|
||||||
|
|
||||||
|
project_root = Path(__file__).parent.parent.parent
|
||||||
|
param_path = args.param_path or (project_root / "params")
|
||||||
|
print(f"Starting AstrAI inference server on http://{args.host}:{args.port}")
|
||||||
|
print(f"Model parameters expected at: {param_path}")
|
||||||
|
print(f"Device: {args.device}, Dtype: {args.dtype}")
|
||||||
|
run_server(
|
||||||
|
host=args.host,
|
||||||
|
port=args.port,
|
||||||
|
reload=args.reload,
|
||||||
|
device=args.device,
|
||||||
|
dtype=dtype,
|
||||||
|
param_path=param_path,
|
||||||
|
max_batch_size=args.max_batch_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -0,0 +1,299 @@
|
|||||||
|
import argparse
|
||||||
|
import os
|
||||||
|
from functools import partial
|
||||||
|
|
||||||
|
import safetensors.torch as st
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.optim as optim
|
||||||
|
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||||
|
|
||||||
|
from astrai.config import ModelConfig, TrainConfig
|
||||||
|
from astrai.dataset import DatasetFactory
|
||||||
|
from astrai.model import Transformer
|
||||||
|
from astrai.parallel import get_rank
|
||||||
|
from astrai.trainer import SchedulerFactory, Trainer
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> argparse.Namespace:
|
||||||
|
|
||||||
|
parser = argparse.ArgumentParser(description="Train the Transformer model.")
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--train_type",
|
||||||
|
type=str,
|
||||||
|
required=True,
|
||||||
|
choices=["seq", "sft", "dpo", "grpo"],
|
||||||
|
help="Train type.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--data_root_path",
|
||||||
|
type=str,
|
||||||
|
required=True,
|
||||||
|
help="Path to the root directory of the dataset.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--param_path",
|
||||||
|
type=str,
|
||||||
|
required=True,
|
||||||
|
help="Path to the model parameters or resume checkpoint.",
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--n_epoch", type=int, default=1, help="Number of epochs to train."
|
||||||
|
)
|
||||||
|
parser.add_argument("--batch_size", type=int, default=1, help="Batch size per GPU.")
|
||||||
|
parser.add_argument(
|
||||||
|
"--accumulation_steps",
|
||||||
|
type=int,
|
||||||
|
default=1,
|
||||||
|
help="Number of iterations between each optimizer step.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--warmup_steps",
|
||||||
|
type=int,
|
||||||
|
default=1000,
|
||||||
|
help="Number of warmup steps for LR scheduler.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--max_lr", type=float, default=3e-4, help="Max learning rate for training."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--max_grad_norm",
|
||||||
|
type=float,
|
||||||
|
default=1.0,
|
||||||
|
help="Max gradient norm for clipping.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--adamw_beta1",
|
||||||
|
type=float,
|
||||||
|
default=0.9,
|
||||||
|
help="Beta values for AdamW optimizer.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--adamw_beta2",
|
||||||
|
type=float,
|
||||||
|
default=0.95,
|
||||||
|
help="Beta values for AdamW optimizer.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--adamw_weight_decay",
|
||||||
|
type=float,
|
||||||
|
default=0.01,
|
||||||
|
help="Weight decay for AdamW optimizer.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--random_seed", type=int, default=3407, help="Random seed for reproducibility."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--num_workers", type=int, default=4, help="Number of workers for data loading."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--no_pin_memory",
|
||||||
|
action="store_false",
|
||||||
|
dest="pin_memory",
|
||||||
|
help="Disable pin memory",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--window_size",
|
||||||
|
type=int,
|
||||||
|
default=None,
|
||||||
|
help="Max length of the input sequence.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--stride", type=int, default=None, help="Step size of the input sequence."
|
||||||
|
)
|
||||||
|
parser.add_argument("--dpo_beta", type=float, default=0.1, help="DPO beta value.")
|
||||||
|
parser.add_argument("--group_size", type=int, default=4, help="GRPO group size.")
|
||||||
|
parser.add_argument(
|
||||||
|
"--grpo_clip_eps", type=float, default=0.2, help="GRPO clipping epsilon."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--grpo_kl_coef", type=float, default=0.01, help="GRPO KL penalty coefficient."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--label_smoothing",
|
||||||
|
type=float,
|
||||||
|
default=0.1,
|
||||||
|
help="cross_entropy function label smoothing parameter",
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--ckpt_interval",
|
||||||
|
type=int,
|
||||||
|
default=5000,
|
||||||
|
help="Number of iters between checkpoints.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--ckpt_dir",
|
||||||
|
type=str,
|
||||||
|
default="checkpoint",
|
||||||
|
help="Directory to save checkpoints.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--grpo_sync_interval",
|
||||||
|
type=int,
|
||||||
|
default=200,
|
||||||
|
help="GRPO ref model sync interval (steps).",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--start_epoch", type=int, default=0, help="Start epoch for training."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--start_batch", type=int, default=0, help="Start batch for training."
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument("--nprocs", type=int, default=1, help="Number of GPUs to use.")
|
||||||
|
parser.add_argument(
|
||||||
|
"--device_type", type=str, default="cuda", help="Device type to use."
|
||||||
|
)
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
return args
|
||||||
|
|
||||||
|
|
||||||
|
def ddp_wrap(model: nn.Module):
|
||||||
|
local_rank = get_rank()
|
||||||
|
model = model.to(dtype=torch.bfloat16)
|
||||||
|
ddp_model = DDP(
|
||||||
|
model,
|
||||||
|
device_ids=[local_rank],
|
||||||
|
output_device=local_rank,
|
||||||
|
find_unused_parameters=False,
|
||||||
|
)
|
||||||
|
return ddp_model
|
||||||
|
|
||||||
|
|
||||||
|
def create_optimizer(model: nn.Module, **kwargs) -> optim.Optimizer:
|
||||||
|
return optim.AdamW(model.parameters(), **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
def create_scheduler(
|
||||||
|
optimizer: optim.Optimizer, **kwargs
|
||||||
|
) -> optim.lr_scheduler.LRScheduler:
|
||||||
|
return SchedulerFactory.create(optimizer, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
def prepare_checkpoint(model: nn.Module) -> dict:
|
||||||
|
return model.module.state_dict()
|
||||||
|
|
||||||
|
|
||||||
|
def train(
|
||||||
|
train_type: str,
|
||||||
|
param_path: str,
|
||||||
|
data_root_path: str,
|
||||||
|
max_lr: float,
|
||||||
|
n_epoch: int,
|
||||||
|
batch_size: int,
|
||||||
|
start_epoch: int,
|
||||||
|
start_batch: int,
|
||||||
|
accumulation_steps: int,
|
||||||
|
warmup_steps: int,
|
||||||
|
ckpt_interval: int,
|
||||||
|
ckpt_dir: str,
|
||||||
|
dpo_beta: float,
|
||||||
|
grpo_clip_eps: float,
|
||||||
|
grpo_kl_coef: float,
|
||||||
|
group_size: int,
|
||||||
|
grpo_sync_interval: int,
|
||||||
|
adamw_beta1: float,
|
||||||
|
adamw_beta2: float,
|
||||||
|
adamw_weight_decay: float,
|
||||||
|
max_grad_norm: float,
|
||||||
|
label_smoothing: float,
|
||||||
|
random_seed: int,
|
||||||
|
num_workers: int,
|
||||||
|
pin_memory: bool,
|
||||||
|
window_size: int,
|
||||||
|
stride: int,
|
||||||
|
nprocs: int,
|
||||||
|
device_type: str,
|
||||||
|
):
|
||||||
|
assert train_type in ["seq", "sft", "dpo", "grpo"]
|
||||||
|
assert os.path.exists(param_path)
|
||||||
|
|
||||||
|
# Load config
|
||||||
|
config = ModelConfig()
|
||||||
|
config_path = os.path.join(param_path, "config.json")
|
||||||
|
if os.path.exists(config_path):
|
||||||
|
config.load(config_path)
|
||||||
|
|
||||||
|
if window_size is None:
|
||||||
|
window_size = config.max_len
|
||||||
|
|
||||||
|
# Create bare Transformer (for training, no tokenizer needed)
|
||||||
|
model = Transformer(config)
|
||||||
|
|
||||||
|
# Load weights if available
|
||||||
|
weights_path = os.path.join(param_path, "model.safetensors")
|
||||||
|
if os.path.exists(weights_path):
|
||||||
|
state_dict = st.load_file(weights_path)
|
||||||
|
model.load_state_dict(state_dict, strict=False)
|
||||||
|
|
||||||
|
strategy_kwargs = {
|
||||||
|
"dpo_beta": dpo_beta,
|
||||||
|
"label_smoothing": label_smoothing,
|
||||||
|
"clip_eps": grpo_clip_eps,
|
||||||
|
"kl_coef": grpo_kl_coef,
|
||||||
|
"group_size": group_size,
|
||||||
|
"sync_interval": grpo_sync_interval,
|
||||||
|
}
|
||||||
|
|
||||||
|
dataset = DatasetFactory.load(
|
||||||
|
train_type=train_type,
|
||||||
|
load_path=data_root_path,
|
||||||
|
window_size=window_size,
|
||||||
|
stride=stride,
|
||||||
|
)
|
||||||
|
|
||||||
|
optimizer_fn = partial(
|
||||||
|
create_optimizer,
|
||||||
|
**{
|
||||||
|
"lr": max_lr,
|
||||||
|
"betas": (adamw_beta1, adamw_beta2),
|
||||||
|
"weight_decay": adamw_weight_decay,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
total_steps = len(dataset) * n_epoch // (batch_size * nprocs)
|
||||||
|
scheduler_fn = partial(
|
||||||
|
create_scheduler,
|
||||||
|
**{
|
||||||
|
"schedule_type": "cosine",
|
||||||
|
"warmup_steps": warmup_steps,
|
||||||
|
"lr_decay_steps": total_steps - warmup_steps,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
train_config = TrainConfig(
|
||||||
|
model=model,
|
||||||
|
strategy=train_type,
|
||||||
|
dataset=dataset,
|
||||||
|
optimizer_fn=optimizer_fn,
|
||||||
|
scheduler_fn=scheduler_fn,
|
||||||
|
ckpt_dir=ckpt_dir,
|
||||||
|
n_epoch=n_epoch,
|
||||||
|
batch_size=batch_size,
|
||||||
|
start_epoch=start_epoch,
|
||||||
|
start_batch=start_batch,
|
||||||
|
ckpt_interval=ckpt_interval,
|
||||||
|
accumulation_steps=accumulation_steps,
|
||||||
|
max_grad_norm=max_grad_norm,
|
||||||
|
random_seed=random_seed,
|
||||||
|
num_workers=num_workers,
|
||||||
|
pin_memory=pin_memory,
|
||||||
|
nprocs=nprocs,
|
||||||
|
parallel_wrapper=ddp_wrap,
|
||||||
|
state_dict_fn=prepare_checkpoint,
|
||||||
|
device_type=device_type,
|
||||||
|
extra_kwargs=strategy_kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
trainer = Trainer(train_config)
|
||||||
|
trainer.train()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
args = parse_args()
|
||||||
|
train(**vars(args))
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user