Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
19532440b4 | ||
|
|
9096e413c3 | ||
|
|
9d5e9fa6c4 | ||
|
|
08dde46778 | ||
|
|
513f1f7826 | ||
|
|
e3382f6bb5 | ||
|
|
f0339022c1 | ||
|
|
d8da2cf17c | ||
|
|
205b40bd28 | ||
|
|
18fe6e9339 | ||
|
|
2196c34c52 | ||
|
|
466c2e1efd | ||
|
|
7e26d848ab | ||
|
|
ed95ef245c | ||
|
|
6d6ef99e66 | ||
|
|
a8e2a1ba45 | ||
|
|
6269bacfc3 | ||
|
|
c0effc9f5b | ||
|
|
df0845e916 | ||
|
|
7440e9c809 | ||
|
|
7d4029c2a4 | ||
|
|
0ca6c9e6eb | ||
|
|
6e49d27057 | ||
|
|
5203b7f53e | ||
|
|
5889179c54 | ||
|
|
38e18fdfd3 | ||
|
|
4753958f92 | ||
|
|
73d6cc0f26 | ||
|
|
317ed90bac | ||
|
|
951df8155c | ||
|
|
a58fab8d6e | ||
|
|
a3c8296135 | ||
|
|
c95ace41aa | ||
|
|
3da428e0e4 | ||
|
|
133a9de98f | ||
|
|
523eacf5fe | ||
|
|
cffedaad5e | ||
|
|
3583c46b66 | ||
|
|
ca4e6b907c | ||
|
|
db99d8b254 | ||
|
|
b98c9cefdc | ||
|
|
283bcaf2ff | ||
|
|
bc7c82977e | ||
|
|
34a511e36e | ||
|
|
d73f52a2f8 | ||
|
|
9d96b0431d | ||
|
|
f81e2b4a73 | ||
|
|
4e324d8f26 | ||
|
|
6ed0506491 | ||
|
|
30cc2d67a4 | ||
|
|
7ddebf2cd9 | ||
|
|
78dc2bd41c | ||
|
|
44d7a4e959 | ||
|
|
c4401512f2 | ||
|
|
a6f5ff3b37 | ||
|
|
ffff05b2c6 | ||
|
|
b89f8436ea | ||
|
|
123f25e339 | ||
|
|
520de3ebe8 | ||
|
|
466c34d7a8 | ||
|
|
6831a15424 | ||
|
|
0f9e5c5049 | ||
|
|
cb0e7f2a80 | ||
|
|
296db909aa | ||
|
|
a2ae742988 | ||
|
|
29beb174a5 | ||
|
|
bbeaff4c60 | ||
|
|
ab5e207f42 | ||
|
|
b0eff02446 | ||
|
|
408f0cb513 | ||
|
|
64b78ecce3 | ||
|
|
f2ffdf60d0 | ||
|
|
ace8f6ee68 | ||
|
|
a57a16430d | ||
|
|
3fee87897d | ||
|
|
3f67e53088 | ||
|
|
bf7adb35b3 | ||
|
|
feaa3fca36 | ||
|
|
39766aa1dc | ||
|
|
9b22b1651e | ||
|
|
e58dbd7c57 | ||
|
|
d2fe8afbd1 | ||
|
|
23ce4bc3ae | ||
|
|
d2b36cc85d | ||
|
|
fc278d17ab | ||
|
|
ff43a2fab8 | ||
|
|
2b26f03bd3 | ||
|
|
861d33b1a1 | ||
|
|
99b821ebf5 | ||
|
|
c94a246c71 | ||
|
|
2dc9545d7f | ||
|
|
9c31d78a22 | ||
|
|
bd9741dc5f | ||
|
|
b531232a9b | ||
|
|
3346c75584 | ||
|
|
aa5e03d7f6 | ||
|
|
073baf105c | ||
|
|
e97536758f | ||
|
|
7861af12e4 | ||
|
|
7f0552013a | ||
|
|
3535de5cc4 | ||
|
|
26989e54aa | ||
|
|
70d52935f0 | ||
|
|
c0e0e6afd9 | ||
|
|
0852b852f8 | ||
|
|
3a7d98a950 | ||
|
|
c5560740b6 | ||
|
|
94c6a015c8 | ||
|
|
8b6509b305 | ||
|
|
912d7c7f54 | ||
|
|
475de51c7d | ||
|
|
9f1561afe7 | ||
|
|
80c0b20877 | ||
|
|
e7721eafc6 | ||
|
|
4ead0a20cf | ||
|
|
b1527d9575 | ||
|
|
2e009cf59a | ||
|
|
780b9e1855 | ||
|
|
aef7615abd | ||
|
|
50488bd659 | ||
|
|
eb57e55fca | ||
|
|
426af2d75f | ||
|
|
345fd2f091 | ||
|
|
e1f9901384 | ||
|
|
0e7fc623b4 | ||
|
|
3e33c14376 | ||
|
|
60f4df95bd | ||
|
|
c01791ff54 | ||
|
|
980299cd54 | ||
|
|
3e8f2eba81 | ||
|
|
361cdeb296 | ||
|
|
50f76cd7c7 | ||
|
|
0f518473af | ||
|
|
a5574f92e2 | ||
|
|
abcedf892e | ||
|
|
abc3a06266 | ||
|
|
62fba9a298 | ||
|
|
e23a5ca426 | ||
|
|
e55b57d771 | ||
|
|
c4feab96fe | ||
|
|
e35cb0d84a | ||
|
|
6d6ef6dbb6 | ||
|
|
493fe4e84b |
@@ -0,0 +1,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
|
||||||
@@ -1,17 +0,0 @@
|
|||||||
name: Spell Check
|
|
||||||
on: [push, pull_request]
|
|
||||||
|
|
||||||
permissions:
|
|
||||||
contents: read
|
|
||||||
|
|
||||||
jobs:
|
|
||||||
spellcheck:
|
|
||||||
runs-on: ubuntu-latest
|
|
||||||
steps:
|
|
||||||
- uses: actions/checkout@v4
|
|
||||||
- name: Check spelling in specific files
|
|
||||||
uses: codespell-project/actions-codespell@v2
|
|
||||||
with:
|
|
||||||
check_filenames: true
|
|
||||||
only_warn: false
|
|
||||||
path: "**/*.{md, py}"
|
|
||||||
@@ -0,0 +1,31 @@
|
|||||||
|
name: Tests
|
||||||
|
|
||||||
|
on:
|
||||||
|
push:
|
||||||
|
branches: [main]
|
||||||
|
pull_request:
|
||||||
|
branches: [main]
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
test:
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
strategy:
|
||||||
|
matrix:
|
||||||
|
python-version: ["3.12"]
|
||||||
|
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
|
||||||
|
- name: Set up Python ${{ matrix.python-version }}
|
||||||
|
uses: actions/setup-python@v5
|
||||||
|
with:
|
||||||
|
python-version: ${{ matrix.python-version }}
|
||||||
|
|
||||||
|
- name: Install dependencies
|
||||||
|
run: |
|
||||||
|
pip install --upgrade pip
|
||||||
|
pip install .[dev]
|
||||||
|
|
||||||
|
- name: Run tests with pytest
|
||||||
|
run: |
|
||||||
|
python -m pytest tests/ -v
|
||||||
+15
-4
@@ -6,7 +6,18 @@
|
|||||||
|
|
||||||
# Allow specific file types and root files
|
# Allow specific file types and root files
|
||||||
!*.py
|
!*.py
|
||||||
!*.md
|
!*.sh
|
||||||
!*.png
|
|
||||||
!LICENSE
|
# Allow GitHub files
|
||||||
!pyproject.toml
|
!/.github/**
|
||||||
|
|
||||||
|
# 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 (requires conda environment `nlp`):
|
||||||
|
```bash
|
||||||
|
conda run -n nlp ruff format .
|
||||||
|
conda run -n nlp ruff check --fix .
|
||||||
|
```
|
||||||
|
- The project uses **double quotes** for strings and **4‑space indentation** (as configured in `pyproject.toml`).
|
||||||
|
|
||||||
|
## Testing
|
||||||
|
|
||||||
|
If you add or modify functionality, please include appropriate tests.
|
||||||
|
|
||||||
|
- Run the test suite with:
|
||||||
|
```bash
|
||||||
|
conda run -n nlp python -u -m 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,286 +1,240 @@
|
|||||||

|
<div align="center">
|
||||||
|
|
||||||
<div style="display: flex; flex-direction: column; align-items: center; justify-content: center; text-align: center; font-size: 16px; font-weight: bold; margin-top: 50px;">
|
<img src="assets/images/logo.png" width="auto" alt="Logo">
|
||||||
|
<p>
|
||||||
<div>
|
<strong>A lightweight Transformer training & inference framework</strong>
|
||||||
<a href="#english" style="text-decoration: none; margin: 0 10px; color: blue;">English</a> |
|
</p>
|
||||||
<a href="#chinese" style="text-decoration: none; margin: 0 10px; color: blue;">中文</a>
|
|
||||||
</div>
|
|
||||||
|
|
||||||
<h1 style="margin: 20px 0 0 0; font-size: 2.5em; font-weight: bold;">KHAOSZ </h1>
|
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<h2 id="english">English Version</h2>
|
<div align="center">
|
||||||
|
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
||||||
|
<img src="https://img.shields.io/badge/license-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>
|
||||||
|
|
||||||
A training and inference framework for autoregressive Transformer language models.
|
<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) and check **Files and versions**
|
## 📖 Table of Contents
|
||||||
2. Run `scripts/download.py` to download model parameters
|
|
||||||
|
|
||||||
**Demo Video:** [bilibili](https://www.bilibili.com/video/BV1z5RPYHEkd)
|
- [Features](#features)
|
||||||
|
- [Quick Start](#quick-start)
|
||||||
|
- [Documentation](#documentation)
|
||||||
|
- [Contributing](#contributing)
|
||||||
|
- [Community](#community)
|
||||||
|
- [License](#license)
|
||||||
|
|
||||||
For training data sources, please refer to the **Model Card** section on the HuggingFace download page.
|
---
|
||||||
|
|
||||||
**License:** The code follows the GPL-3.0 license. Please provide attribution when using it.
|
<a id="english"></a>
|
||||||
|
## English
|
||||||
|
|
||||||
- **📊 Device Selection:** Uses CUDA for training by default
|
### Features
|
||||||
- **🌐 Performance Optimization:** Enable `dtype=torch.bfloat16` to accelerate training and reduce memory usage. Ensure your hardware supports this feature
|
|
||||||
- **🤖 Language Support:** The model supports training in Chinese and English. Since the BBPE tokenizer hasn't been trained on multilingual text, OOV (Out-of-Vocabulary) issues are minimal for Chinese and English, but may exist for other languages
|
|
||||||
|
|
||||||
|
- 🚀 **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.
|
||||||
|
|
||||||
### 📌 Training Guide
|
### Quick Start
|
||||||
|
|
||||||
To train this Transformer model, follow these steps:
|
#### Installation
|
||||||
|
|
||||||
**(1). Prepare the Dataset:**
|
|
||||||
|
|
||||||
Place the dataset in the specified root directory. This system uses the BBPE tokenizer for tokenization and requires training with pre-tokenized segments (stored as *.h5 format files).
|
|
||||||
|
|
||||||
**(2). Install Dependencies:**
|
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
git clone https://github.com/ViperEkura/AstrAI.git
|
||||||
|
cd AstrAI
|
||||||
pip install -e .
|
pip install -e .
|
||||||
```
|
```
|
||||||
|
|
||||||
**(3). Run the 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
|
|
||||||
```
|
```
|
||||||
|
|
||||||
**Parameter Explanation:**
|
#### Download Pre-trained Model
|
||||||
- `--train_type`: Training type (seq, sft, dpo)
|
|
||||||
- `--data_root_path`: Dataset root directory
|
|
||||||
- `--param_path`: Path to model training parameters
|
|
||||||
- `--n_epoch`: Total number of training epochs
|
|
||||||
- `--batch_size`: Batch size
|
|
||||||
- `--accumulation_steps`: Number of batches per training step
|
|
||||||
- `--warmup_steps`: Warmup steps
|
|
||||||
- `--max_lr`: Maximum learning rate (using warmup + cosine decay)
|
|
||||||
- `--checkpoint_interval`: Checkpoint saving interval
|
|
||||||
- `--checkpoint_dir`: Checkpoint saving directory
|
|
||||||
- `--resume_dir`: Resume training from specified path
|
|
||||||
|
|
||||||
|
Download pre-trained model weights (1B bilingual checkpoint) to `params/`:
|
||||||
|
|
||||||
### 👉 Usage Guide
|
|
||||||
|
|
||||||
**(1). Chat with the Model:**
|
|
||||||
|
|
||||||
Open `chat.py` or use the 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)
|
|
||||||
```
|
|
||||||
|
|
||||||
<h2 id="chinese">中文版本</h2>
|
|
||||||
这是一个支持基于自回归模式的 Transfomer 语言模型训练以及推理框架
|
|
||||||
|
|
||||||
**模型下载选项(任选其一):**
|
|
||||||
|
|
||||||
1. 访问 [HuggingFace](https://huggingface.co/ViperEk/KHAOSZ) 查看 **Files and versions**
|
|
||||||
2. 运行 `scripts/download.py` 下载模型参数
|
|
||||||
|
|
||||||
**演示视频:** [bilibili](https://www.bilibili.com/video/BV1z5RPYHEkd)
|
|
||||||
|
|
||||||
训练数据来源请参见 HuggingFace 下载页面中的 **Model Card** 部分。
|
|
||||||
|
|
||||||
**许可证:** 代码遵循 GPL-3.0 协议,使用时请注明出处。
|
|
||||||
|
|
||||||
- **📊 设备选择:** 默认使用 CUDA 进行训练
|
|
||||||
- **🌐 性能优化:** 启用 `dtype=torch.bfloat16` 以加速训练并减少内存占用,请确保硬件支持该特性
|
|
||||||
- **🤖 语言支持:** 模型支持中文和英文训练。由于 BBPE 分词器未使用多语言文本训练,因此中英文的 OOV(未登录词)问题较少,其他语言可能存在 OOV 问题
|
|
||||||
|
|
||||||
|
|
||||||
### 📌 训练指南
|
|
||||||
|
|
||||||
要训练该 Transformer 模型,请按照以下步骤操作:
|
|
||||||
|
|
||||||
**(1). 准备数据集:**
|
|
||||||
|
|
||||||
将数据集放置在指定的根目录下, 本系统采用 BBPE 分词器进行分词,并且要求使用已经经过分词的 token 分段训练(分段存储为 *.h5 格式)
|
|
||||||
|
|
||||||
**(2). 安装依赖:**
|
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
pip install -e .
|
python scripts/demo/download.py
|
||||||
```
|
```
|
||||||
|
|
||||||
**(3). 运行训练脚本:**
|
Or download manually from [HuggingFace](https://huggingface.co/ViperEk/KHAOSZ) into `params/`.
|
||||||
|
|
||||||
|
#### Train a Model
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python train.py \
|
CUDA_VISIBLE_DEVICES=0,1,2,3 python scripts/tools/train.py \
|
||||||
--train_type=train_type[seq, sft, dpo] \
|
--train_type seq \
|
||||||
--data_root_path=/path/to/dataset \
|
--data_root_path /path/to/dataset \
|
||||||
--param_path=/path/to/param_path \
|
--param_path /path/to/model \
|
||||||
--n_epoch=5 \
|
--batch_size 4 \
|
||||||
--batch_size=8 \
|
--accumulation_steps 8 \
|
||||||
--max_lr=2e-4 \
|
--max_lr 3e-4 \
|
||||||
--checkpoint_interval=10000 \
|
--warmup_steps 1000 \
|
||||||
--checkpoint_dir=checkpoints
|
--n_epoch 1
|
||||||
```
|
```
|
||||||
|
|
||||||
**参数说明:**
|
Full reference at [Parameter Guide](assets/docs/params.md).
|
||||||
- `--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`: 从指定路径恢复训练
|
|
||||||
|
|
||||||
|
#### Generate Text
|
||||||
|
|
||||||
|
```bash
|
||||||
### 👉 使用指南
|
python scripts/tools/generate.py \
|
||||||
|
--param_path /path/to/model \
|
||||||
**(1). 与模型对话:**
|
--input_json_file /path/to/input.json \
|
||||||
|
--output_json_file /path/to/output.json
|
||||||
打开 `chat.py` 或使用流式/非流式接口:
|
|
||||||
|
|
||||||
**流式输出:**
|
|
||||||
```python
|
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
model_dir = "your_model_parameter_dir"
|
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
history = []
|
|
||||||
|
|
||||||
while True:
|
|
||||||
query = input(">> ")
|
|
||||||
if query == "!exit":
|
|
||||||
break
|
|
||||||
|
|
||||||
response_size = 0
|
|
||||||
for response, history in model.stream_generate(
|
|
||||||
query=query,
|
|
||||||
history=history,
|
|
||||||
temperature=0.85,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50
|
|
||||||
):
|
|
||||||
print(response[response_size:], end="")
|
|
||||||
response_size = len(response)
|
|
||||||
```
|
```
|
||||||
|
|
||||||
**非流式输出:**
|
#### Docker
|
||||||
```python
|
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
model_dir = "your_model_parameter_dir"
|
Build and run with Docker (recommended for GPU environments):
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
history = []
|
|
||||||
|
|
||||||
while True:
|
```bash
|
||||||
query = input(">> ")
|
# Build image
|
||||||
if query == "!exit":
|
docker build -t astrai:latest .
|
||||||
break
|
|
||||||
|
|
||||||
response = model.generate(
|
# Run with GPU support
|
||||||
query=query,
|
docker run --gpus all -it astrai:latest
|
||||||
history=history,
|
|
||||||
temperature=0.85,
|
# Run with specific GPUs
|
||||||
top_p=0.95,
|
docker run --gpus '"device=0,1"' -it astrai:latest
|
||||||
top_k=50
|
|
||||||
)
|
# Run inference server
|
||||||
print(response)
|
docker run --gpus all -p 8000:8000 astrai:latest \
|
||||||
|
python -m scripts.tools.server --port 8000 --device cuda
|
||||||
|
|
||||||
|
# Run with volume mount for data
|
||||||
|
docker run --gpus all -v /path/to/data:/data -it astrai:latest
|
||||||
|
|
||||||
|
# Docker Compose (GPU, default)
|
||||||
|
docker compose up -d
|
||||||
|
|
||||||
|
# Docker Compose (CPU only)
|
||||||
|
docker compose --profile cpu up -d
|
||||||
```
|
```
|
||||||
|
|
||||||
**(2). 基于检索的生成(RAG):**
|
> **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`.
|
||||||
|
|
||||||
```python
|
#### Start HTTP Server
|
||||||
import torch
|
|
||||||
from khaosz import Khaosz
|
|
||||||
|
|
||||||
model_dir = "your_model_parameter_dir"
|
Start the inference server with OpenAI and Anthropic-compatible HTTP API:
|
||||||
model = Khaosz(model_dir).to(device='cuda', dtype=torch.bfloat16)
|
|
||||||
|
|
||||||
retrieved_content = model.retrieve_generate(
|
```bash
|
||||||
query=query,
|
python -m scripts.tools.server --port 8000 --device cuda
|
||||||
retrieve_top_k=5,
|
|
||||||
temperature=0.6,
|
|
||||||
top_k=30,
|
|
||||||
top_p=0.95
|
|
||||||
)
|
|
||||||
print(retrieved_content)
|
|
||||||
```
|
```
|
||||||
|
|
||||||
|
Make requests:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# OpenAI-compatible
|
||||||
|
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"messages": [{"role": "user", "content": "Hello"}],
|
||||||
|
"max_tokens": 512
|
||||||
|
}'
|
||||||
|
|
||||||
|
# OpenAI-compatible streaming
|
||||||
|
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"messages": [{"role": "user", "content": "Tell a story"}],
|
||||||
|
"stream": true,
|
||||||
|
"max_tokens": 500
|
||||||
|
}'
|
||||||
|
|
||||||
|
# Anthropic-compatible
|
||||||
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"model": "astrai",
|
||||||
|
"system": "You are a helpful assistant.",
|
||||||
|
"messages": [{"role": "user", "content": "Hello"}],
|
||||||
|
"max_tokens": 512
|
||||||
|
}'
|
||||||
|
|
||||||
|
# Anthropic-compatible streaming with stop sequences
|
||||||
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"model": "astrai",
|
||||||
|
"messages": [{"role": "user", "content": "Write a story"}],
|
||||||
|
"max_tokens": 500,
|
||||||
|
"stream": true,
|
||||||
|
"stop_sequences": ["The end"]
|
||||||
|
}'
|
||||||
|
|
||||||
|
# Health check
|
||||||
|
curl http://localhost:8000/health
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Demo
|
||||||
|
|
||||||
|
Check out the demos in the `scripts/demo/` folder:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Download pre‑processed data (required before running demos)
|
||||||
|
python scripts/demo/download.py
|
||||||
|
|
||||||
|
# Interactive streaming chat
|
||||||
|
python scripts/demo/stream_chat.py
|
||||||
|
|
||||||
|
# Batch generation
|
||||||
|
python scripts/demo/generate_batch.py
|
||||||
|
|
||||||
|
# Auto‑regressive generation
|
||||||
|
python scripts/demo/generate_ar.py
|
||||||
|
```
|
||||||
|
|
||||||
|
Watch a video walkthrough on [bilibili](https://www.bilibili.com/video/BV1z5RPYHEkd).
|
||||||
|
|
||||||
|
### Documentation
|
||||||
|
|
||||||
|
| Document | Description |
|
||||||
|
|----------|-------------|
|
||||||
|
| [Parameter Guide](./assets/docs/params.md) | Training & inference parameters |
|
||||||
|
| [Design Document](./assets/docs/design.md) | Framework architecture & module design |
|
||||||
|
| [Data Flow](./assets/docs/dataflow.md) | Data processing pipeline details |
|
||||||
|
| [Model Introduction](./assets/docs/introduction.md) | Model architecture & technical details |
|
||||||
|
|
||||||
|
### 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,246 @@
|
|||||||
|
<div align="center">
|
||||||
|
|
||||||
|
<img src="../images/logo.png" width="auto" alt="Logo">
|
||||||
|
|
||||||
|
<div>
|
||||||
|
<a href="../../README.md">English</a> •
|
||||||
|
<a href="#chinese">中文</a>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<p>
|
||||||
|
<strong>轻量级 Transformer 训练与推理框架</strong>
|
||||||
|
</p>
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<div align="center">
|
||||||
|
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
||||||
|
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license">
|
||||||
|
<img src="https://img.shields.io/github/v/release/ViperEkura/AstrAI?color=76bad9" alt="release">
|
||||||
|
<img src="https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fapi.github.com%2Frepos%2FViperEkura%2FAstrAI&query=%24.stargazers_count&label=stars&suffix=%20stars&color=76bad9" alt="stars">
|
||||||
|
<img src="https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fapi.github.com%2Frepos%2FViperEkura%2FAstrAI&query=%24.forks_count&label=forks&suffix=%20forks&color=76bad9" alt="forks">
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<br>
|
||||||
|
|
||||||
|
<div align="center">
|
||||||
|
<a href="../../README.md">English</a> •
|
||||||
|
<a href="#chinese">中文</a> •
|
||||||
|
<a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> •
|
||||||
|
<a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> •
|
||||||
|
<a href="https://huggingface.co/ViperEk">HuggingFace</a>
|
||||||
|
</div>
|
||||||
|
<br>
|
||||||
|
|
||||||
|
## 📖 目录
|
||||||
|
|
||||||
|
- [特性](#特性)
|
||||||
|
- [快速开始](#快速开始)
|
||||||
|
- [文档](#文档)
|
||||||
|
- [贡献](#贡献)
|
||||||
|
- [社区](#社区)
|
||||||
|
- [许可证](#许可证)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
<a id="chinese"></a>
|
||||||
|
## 中文
|
||||||
|
|
||||||
|
### 特性
|
||||||
|
|
||||||
|
- 🚀 **高性能**: 训练与推理双向优化,高效并行。
|
||||||
|
- 🔧 **灵活**: 支持 seq/sft/dpo/grpo 多种训练方式,可定制模型架构。
|
||||||
|
- 💡 **易用**: 简洁的 API 与丰富的示例、演示。
|
||||||
|
- 📦 **轻量**: 依赖少,部署简单。
|
||||||
|
- 🔬 **研究友好**: 模块化设计,便于实验新想法。
|
||||||
|
- 🤗 **HuggingFace 风格 API**: 类 HuggingFace 的 AutoModel/AutoTokenizer 接口,方便加载模型和分词器。
|
||||||
|
- 🔌 **双 API 兼容**: 同时支持 OpenAI 和 Anthropic 聊天补全 API,开箱即用。
|
||||||
|
|
||||||
|
### 快速开始
|
||||||
|
|
||||||
|
#### 安装
|
||||||
|
|
||||||
|
```bash
|
||||||
|
git clone https://github.com/ViperEkura/AstrAI.git
|
||||||
|
cd AstrAI
|
||||||
|
pip install -e .
|
||||||
|
```
|
||||||
|
|
||||||
|
安装开发依赖:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
pip install -e ".[dev]"
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 下载预训练模型
|
||||||
|
|
||||||
|
下载预训练模型权重(1B 双语检查点)到 `params/` 目录:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/demo/download.py
|
||||||
|
```
|
||||||
|
|
||||||
|
或从 [HuggingFace](https://huggingface.co/ViperEk/KHAOSZ) 手动下载放入 `params/`。
|
||||||
|
|
||||||
|
#### 训练模型
|
||||||
|
|
||||||
|
```bash
|
||||||
|
CUDA_VISIBLE_DEVICES=0,1,2,3 python scripts/tools/train.py \
|
||||||
|
--train_type seq \
|
||||||
|
--data_root_path /path/to/dataset \
|
||||||
|
--param_path /path/to/model \
|
||||||
|
--batch_size 4 \
|
||||||
|
--accumulation_steps 8 \
|
||||||
|
--max_lr 3e-4 \
|
||||||
|
--warmup_steps 1000 \
|
||||||
|
--n_epoch 1
|
||||||
|
```
|
||||||
|
|
||||||
|
完整参数列表见[参数说明](./params.md)。
|
||||||
|
|
||||||
|
#### 文本生成
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/tools/generate.py \
|
||||||
|
--param_path /path/to/model \
|
||||||
|
--input_json_file /path/to/input.json \
|
||||||
|
--output_json_file /path/to/output.json
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Docker
|
||||||
|
|
||||||
|
使用 Docker 构建和运行(推荐用于 GPU 环境):
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 构建镜像
|
||||||
|
docker build -t astrai:latest .
|
||||||
|
|
||||||
|
# 启用 GPU 运行
|
||||||
|
docker run --gpus all -it astrai:latest
|
||||||
|
|
||||||
|
# 指定特定 GPU
|
||||||
|
docker run --gpus '"device=0,1"' -it astrai:latest
|
||||||
|
|
||||||
|
# 运行推理服务
|
||||||
|
docker run --gpus all -p 8000:8000 astrai:latest \
|
||||||
|
python -m scripts.tools.server --port 8000 --device cuda
|
||||||
|
|
||||||
|
# 挂载数据卷
|
||||||
|
docker run --gpus all -v /path/to/data:/data -it astrai:latest
|
||||||
|
|
||||||
|
# Docker Compose(GPU,默认)
|
||||||
|
docker compose up -d
|
||||||
|
|
||||||
|
# Docker Compose(仅 CPU)
|
||||||
|
docker compose --profile cpu up -d
|
||||||
|
```
|
||||||
|
|
||||||
|
> **注意**: 必须使用 `--gpus all` 才能启用 CUDA 支持,否则 `torch.cuda.is_available()` 将返回 `False`。
|
||||||
|
|
||||||
|
#### 启动 HTTP 服务
|
||||||
|
|
||||||
|
启动推理服务器,支持 OpenAI 和 Anthropic 兼容的 HTTP API:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python -m scripts.tools.server --port 8000 --device cuda
|
||||||
|
```
|
||||||
|
|
||||||
|
发起请求:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# OpenAI 兼容
|
||||||
|
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"messages": [{"role": "user", "content": "你好"}],
|
||||||
|
"max_tokens": 512
|
||||||
|
}'
|
||||||
|
|
||||||
|
# OpenAI 兼容流式
|
||||||
|
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"messages": [{"role": "user", "content": "讲个故事"}],
|
||||||
|
"stream": true,
|
||||||
|
"max_tokens": 500
|
||||||
|
}'
|
||||||
|
|
||||||
|
# Anthropic 兼容
|
||||||
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"model": "astrai",
|
||||||
|
"system": "你是一个乐于助人的助手。",
|
||||||
|
"messages": [{"role": "user", "content": "你好"}],
|
||||||
|
"max_tokens": 512
|
||||||
|
}'
|
||||||
|
|
||||||
|
# Anthropic 兼容流式并设置停止序列
|
||||||
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"model": "astrai",
|
||||||
|
"messages": [{"role": "user", "content": "写个故事"}],
|
||||||
|
"max_tokens": 500,
|
||||||
|
"stream": true,
|
||||||
|
"stop_sequences": ["结束"]
|
||||||
|
}'
|
||||||
|
|
||||||
|
# 健康检查
|
||||||
|
curl http://localhost:8000/health
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 演示
|
||||||
|
|
||||||
|
查看 `scripts/demo/` 文件夹中的演示:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 下载预处理数据(运行演示前必需)
|
||||||
|
python scripts/demo/download.py
|
||||||
|
|
||||||
|
# 交互式流式聊天
|
||||||
|
python scripts/demo/stream_chat.py
|
||||||
|
|
||||||
|
# 批量生成
|
||||||
|
python scripts/demo/generate_batch.py
|
||||||
|
|
||||||
|
# 自回归生成
|
||||||
|
python scripts/demo/generate_ar.py
|
||||||
|
```
|
||||||
|
|
||||||
|
观看 [bilibili](https://www.bilibili.com/video/BV1z5RPYHEkd) 上的视频演示。
|
||||||
|
|
||||||
|
### 文档
|
||||||
|
|
||||||
|
| 文档 | 说明 |
|
||||||
|
|------|------|
|
||||||
|
| [参数说明](./params.md) | 训练与推理参数配置 |
|
||||||
|
| [设计文档](./design.md) | 系统架构与模块设计 |
|
||||||
|
| [数据流程](./dataflow.md) | 数据处理管道详解 |
|
||||||
|
| [模型介绍](./introduction.md) | 模型架构与技术细节 |
|
||||||
|
|
||||||
|
### 贡献
|
||||||
|
|
||||||
|
我们欢迎贡献!请参阅[贡献指南](../../CONTRIBUTING.md)了解详情。
|
||||||
|
|
||||||
|
1. Fork 本仓库。
|
||||||
|
2. 创建功能分支。
|
||||||
|
3. 提交更改。
|
||||||
|
4. 发起 Pull Request。
|
||||||
|
|
||||||
|
重大更改请先开 issue 讨论。
|
||||||
|
|
||||||
|
### 社区
|
||||||
|
|
||||||
|
- **GitHub Issues**: [问题追踪](https://github.com/ViperEkura/AstrAI/issues)
|
||||||
|
- **Discussions**: [GitHub 讨论区](https://github.com/ViperEkura/AstrAI/discussions)
|
||||||
|
- **HuggingFace**: [模型中心](https://huggingface.co/ViperEk)
|
||||||
|
|
||||||
|
### 许可证
|
||||||
|
|
||||||
|
本项目采用 [GPL-3.0 许可证](../../LICENSE)。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
<div align="center">
|
||||||
|
<em>专为高性能与易用性设计的轻量级 Transformer 框架。</em>
|
||||||
|
</div>
|
||||||
@@ -1,220 +0,0 @@
|
|||||||
## 1. 为什么我要做这个项目?
|
|
||||||
|
|
||||||
现在市面上有很多大模型,比如GPT、LLaMA这些,动不动就是几十亿甚至上千亿参数。但说实话,这些模型对硬件要求太高了,普通开发者根本玩不起。我就想:**能不能做一个既好用又能在普通电脑上跑起来的模型呢?** 这其实也是目前大部分人的期望, 能有一个可以本地部署的ai小型项目,实现完全私有化并且有一定的智能能力。
|
|
||||||
|
|
||||||
于是就有了这个KHAOSZ项目,1B参数,中英双语,支持对话、文本生成、RAG检索,而且训练代码都是开源的!
|
|
||||||
|
|
||||||
## 2. 系统架构
|
|
||||||
|
|
||||||
系统分为以下板块
|
|
||||||
|
|
||||||
```mermaid
|
|
||||||
graph LR
|
|
||||||
%% 样式定义
|
|
||||||
classDef config fill:#e1f5fe,stroke:#01579b;
|
|
||||||
classDef trainer fill:#f3e5f5,stroke:#4a148c;
|
|
||||||
classDef data fill:#e8f5e8,stroke:#1b5e20;
|
|
||||||
classDef model fill:#fff3e0,stroke:#e65100;
|
|
||||||
classDef inference fill:#fce4ec,stroke:#880e4f;
|
|
||||||
classDef parallel fill:#e0f2f1,stroke:#004d40;
|
|
||||||
|
|
||||||
%% 配置模块
|
|
||||||
subgraph Config["Config(配置模块)"]
|
|
||||||
C1[model_config.py]
|
|
||||||
C2[train_config.py]
|
|
||||||
C3[scheduler_config.py]
|
|
||||||
end
|
|
||||||
class Config config;
|
|
||||||
|
|
||||||
%% 训练器模块
|
|
||||||
subgraph Trainer["Trainer(训练器模块)"]
|
|
||||||
T1[trainer.py]
|
|
||||||
T2[train_content.py]
|
|
||||||
T3[schedule.py]
|
|
||||||
T4[strategy.py]
|
|
||||||
T5[train_callback.py]
|
|
||||||
end
|
|
||||||
class Trainer trainer;
|
|
||||||
|
|
||||||
%% 数据模块
|
|
||||||
subgraph Data["Data(数据模块)"]
|
|
||||||
D1[dataset.py]
|
|
||||||
D2[sampler.py]
|
|
||||||
D3[mmap.py]
|
|
||||||
D4[tokenizer.py]
|
|
||||||
D5[checkpoint.py]
|
|
||||||
end
|
|
||||||
class Data data;
|
|
||||||
|
|
||||||
%% 模型模块
|
|
||||||
subgraph Model["Model(模型模块)"]
|
|
||||||
M1[transformer.py]
|
|
||||||
M2[module.py]
|
|
||||||
end
|
|
||||||
class Model model;
|
|
||||||
|
|
||||||
%% 推理模块
|
|
||||||
subgraph Inference["Inference(推理模块)"]
|
|
||||||
I1[generator.py]
|
|
||||||
I2[core.py]
|
|
||||||
end
|
|
||||||
class Inference inference;
|
|
||||||
|
|
||||||
%% 并行模块
|
|
||||||
subgraph Parallel["Parallel(并行模块)"]
|
|
||||||
P1[setup.py]
|
|
||||||
P2[module.py]
|
|
||||||
end
|
|
||||||
class Parallel parallel;
|
|
||||||
|
|
||||||
%% 配置依赖
|
|
||||||
C2 -.-> T1
|
|
||||||
C1 -.-> M1
|
|
||||||
C3 -.-> T3
|
|
||||||
|
|
||||||
%% 训练器内部依赖
|
|
||||||
T1 --> T5
|
|
||||||
T1 --> T2
|
|
||||||
T2 --> T3
|
|
||||||
T2 --> T4
|
|
||||||
|
|
||||||
%% 数据流
|
|
||||||
D1 --> D2
|
|
||||||
D1 --> D3
|
|
||||||
D1 --> D4
|
|
||||||
D1 --> D5
|
|
||||||
|
|
||||||
%% 模型依赖
|
|
||||||
M1 --> M2
|
|
||||||
|
|
||||||
%% 推理依赖
|
|
||||||
I1 --> I2
|
|
||||||
|
|
||||||
%% 跨模块依赖
|
|
||||||
T2 -.-> M1
|
|
||||||
I1 -.-> M1
|
|
||||||
T2 -.-> D1
|
|
||||||
T1 -.-> P1
|
|
||||||
```
|
|
||||||
|
|
||||||
|
|
||||||
### 1. 配置管理(/config/)
|
|
||||||
- **模型配置**:定义模型结构参数(如层数、头数、维度等),通过 `ModelConfig` 统一管理。
|
|
||||||
- **训练配置**:设置训练参数(如批次大小、训练阶段 PT/SFT/DPO、优化器等),由 `TrainConfig` 加载。
|
|
||||||
- **调度配置**:控制学习率策略(如余弦退火)和训练进度。
|
|
||||||
|
|
||||||
### 2. 硬件与并行(/parallel/)
|
|
||||||
- **分布式初始化**:通过 `setup_parallel` 函数,根据配置初始化多卡/多机训练环境。
|
|
||||||
|
|
||||||
### 3. 数据处理(/data/)
|
|
||||||
- **高效加载**:使用内存映射(mmap)技术加载超大语料,避免内存溢出,实现零拷贝读取。
|
|
||||||
|
|
||||||
### 4. 模型与训练(/model/, /trainer/)
|
|
||||||
- **统一模型架构**:基于 Transformer,支持灵活配置不同规模(如7B、13B)。
|
|
||||||
- **策略化训练器**:`Trainer` 根据训练阶段(PT/SFT/DPO)自动切换训练策略,复用同一训练循环。
|
|
||||||
- **训练上下文管理**:统一管理模型、优化器、调度器和指标,支持多阶段无缝衔接。
|
|
||||||
|
|
||||||
### 5. 推理服务(/inference/, /utils/)
|
|
||||||
- **统一生成接口**:提供同步、批量、流式生成方法,适配所有训练阶段。
|
|
||||||
- **KV缓存优化**:在自回归生成中缓存 Key/Value,昇腾XPU下利用高速片上内存加速。
|
|
||||||
- **RAG支持**:结合检索器和嵌入模型,从外部知识库注入相关信息,提升回答质量。
|
|
||||||
- **智能文本分割**:
|
|
||||||
- **结构优先分割**:按标题、段落等切分;
|
|
||||||
- **语义分割**:基于句子嵌入相似度,确保片段语义完整,提升微调效果。
|
|
||||||
|
|
||||||
|
|
||||||
## 3. 训练流程
|
|
||||||
|
|
||||||
常见大语言模型(Large Language Model, LLM)的训练流程通常包含三个阶段:**预训练(Pre-training, PT)**、**监督微调(Supervised Fine-Tuning, SFT)** 以及 **基于人类反馈的强化学习(Reinforcement Learning from Human Feedback, RLHF)**。本系统设计支持全流程无缝衔接,通过模块化策略实现不同训练阶段的高效切换与状态管理,确保模型能力从通用语言理解逐步对齐至符合人类偏好的对话与指令执行。
|
|
||||||
|
|
||||||
### **2.1 预训练阶段**
|
|
||||||
|
|
||||||
预训练阶段旨在构建模型的基础语言能力与通用知识表示。该阶段在大规模、无标注的语料库(通常涵盖数百GB至数TB的文本数据)上进行自监督学习。模型架构基于标准的Transformer Decoder,通过掩码语言建模(如因果语言建模)目标进行训练,使模型能够学习词汇、语法、语义及蕴含于文本中的世界知识。
|
|
||||||
|
|
||||||
**核心公式:因果语言建模(Causal Language Modeling)**
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{PT}} = - \sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta)
|
|
||||||
$$
|
|
||||||
|
|
||||||
**符号说明:**
|
|
||||||
|
|
||||||
- $T$:序列长度
|
|
||||||
- $x_t$:序列中第 $ t $ 个词元(token)
|
|
||||||
- $x_{<t}$:位置 $ t $ 之前的所有词元
|
|
||||||
- $\theta$:模型参数
|
|
||||||
- $P(x_t \mid x_{<t}; \theta)$:模型在给定上文条件下预测下一个词元的概率
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
本阶段的核心在于利用分布式的并行计算资源,实现模型参数的稳定优化。训练器模块中的`PTStrategy`策略,专门负责管理预训练特有的数据采样、长序列分段与梯度累积逻辑。同时,硬件适配模块会根据运行环境(如华为昇腾NPU集群或标准GPU集群)自动选择最优的并行通信后端(如HCCL或NCCL),并进行计算图优化,以最大化硬件利用率和训练吞吐量。
|
|
||||||
|
|
||||||
另外系统通过数据模块中的高效内存映射加载器(`MmapFileHandler`),实现海量数据的零拷贝读取,以克服传统IO瓶颈。
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
### **2.2 监督微调阶段**
|
|
||||||
|
|
||||||
预训练模型虽具备强大的语言生成能力,但尚未对齐至遵循人类指令、进行安全有益对话的行为模式。监督微调阶段旨在弥合这一差距。该阶段使用由人工精心编写的、高质量的“指令-响应”配对数据集。
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
**核心公式:序列到序列条件语言建模**
|
|
||||||
|
|
||||||
设完整序列 $S = [s_1, s_2, \ldots, s_{P+L}]$,其中:
|
|
||||||
|
|
||||||
- 前 $P$ 个token是prompt 以及对应控制token: $X = [s_1, \ldots, s_P]$
|
|
||||||
- 后 $L$ 个token是response以及对应控制token: $Y = [s_{P+1}, \ldots, s_{P+L}]$
|
|
||||||
|
|
||||||
损失函数为:
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{SFT}} = - \sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
|
|
||||||
$$
|
|
||||||
|
|
||||||
训练器模块将动态切换到`SFTStrategy`策略。此策略的核心是引入序列级的监督学习目标,例如预测给定指令下完整、正确的响应序列。训练上下文管理器(`TrainContext`)负责平滑地从PT阶段检查点加载模型状态,并初始化新的优化器和学习率调度器。本阶段不仅优化模型参数,更重要的是引导模型学习“对话”这一特定任务范式,使其输出风格、内容与格式均符合人类期望。
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
### **2.3 基于人类反馈的强化学习阶段**
|
|
||||||
|
|
||||||
为生成更具帮助性、无害性且符合人类偏好的高质量输出,系统进一步集成强化学习阶段。传统的RLHF流程包括**奖励模型训练**与**策略模型微调**两个核心步骤。系统支持以直接偏好优化(Direct Preference Optimization,DPO)算法为代表的策略微调,并针对稳定性与收敛性进行了多项工程优化。
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
#### **2.3.1 传统 RLHF(奖励模型训练)**
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{RM}} = -\mathbb{E}_{(x, y_w, y_l) \sim D} \left[ \log \sigma\left( r_\phi(x, y_w) - r_\phi(x, y_l) \right) \right]
|
|
||||||
$$
|
|
||||||
|
|
||||||
**符号说明:**
|
|
||||||
|
|
||||||
- $r_\phi(x, y)$:参数为 $phi$ 的奖励模型给出的标量分数
|
|
||||||
- $y_w, y_l $:同一提示 $ x $ 下的优选和劣选回答
|
|
||||||
- $\sigma $:sigmoid 函数
|
|
||||||
- $\mathcal{D} $:人类偏好数据集
|
|
||||||
|
|
||||||
|
|
||||||
#### **2.3.2 DPO 直接偏好优化**(推荐)
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{DPO}} = -E_{(x, y_w, y_l) \sim D} \left[ \log \sigma\left( \beta \log \frac{\pi_\theta(y_w \mid x)}{\pi_{\text{ref}}(y_w \mid x)} - \beta \log \frac{\pi_\theta(y_l \mid x)}{\pi_{\text{ref}}(y_l \mid x)} \right) \right]
|
|
||||||
$$
|
|
||||||
|
|
||||||
**符号说明:**
|
|
||||||
|
|
||||||
- $\pi_\theta(y \mid x) $:当前策略模型生成回答的概率
|
|
||||||
- $\pi_{\text{ref}}(y \mid x) $:参考模型生成回答的概率
|
|
||||||
- $\beta $:温度参数(通常设为 0.1-0.5)
|
|
||||||
- 注意:隐式学习奖励函数 $r(x, y) = \beta \log \frac{\pi_\theta(y \mid x)}{\pi_{\text{ref}}(y \mid x)} $
|
|
||||||
|
|
||||||
|
|
||||||
在本阶段,训练器模块启用`RLHFStrategy`策略(或类似的`DPOStrategy`直接偏好优化策略)。该策略管理一个复杂的训练循环,其中包含策略模型(待优化的LLM)、参考模型(通常为SFT后的模型快照)和奖励模型。系统流程如下:
|
|
||||||
|
|
||||||
1. **偏好数据收集与奖励建模**:首先,通过收集人类标注员对同一提示词下多个模型生成结果的排序偏好数据,训练一个独立的奖励模型(Reward Model, RM)。该模型学习为生成文本输出一个标量奖励分数,以量化其符合人类偏好的程度。
|
|
||||||
2. **策略优化**:随后,使用奖励模型作为优化信号,通过强化学习算法对SFT模型(作为策略)进行微调。策略优化的目标是最大化从奖励模型获得的期望累计奖励,同时通过KL散度惩罚项约束策略模型与参考模型的输出分布不过度偏离,以防止模式崩溃并保持生成多样性。训练上下文管理器在此阶段同时维护策略模型、参考模型和奖励模型(或价值函数模型)的状态,并协调复杂的多阶段梯度计算。
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
通过上述三阶段的递进式训练,模型完成了从通用语言基座到专业化、高对齐度对话智能体的进化。系统通过统一的`Trainer`接口和策略模式设计,使得各阶段训练在代码层面高度复用,在流程层面清晰解耦,为大规模语言模型的研发与迭代提供了高效、灵活且可扩展的工程基础。
|
|
||||||
@@ -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`): Checkpoint management with safetensors
|
||||||
|
|
||||||
|
## Data Flow Diagram
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
flowchart LR
|
||||||
|
subgraph A[Data Preparation]
|
||||||
|
direction TB
|
||||||
|
A1[Raw Text] --> A2[AutoTokenizer]
|
||||||
|
A2 --> A3[Tokenized .h5 files]
|
||||||
|
A3 --> A4[BaseDataset]
|
||||||
|
A4 --> A5[ResumableDistributedSampler]
|
||||||
|
A5 --> A6[DataLoader]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph B[Training]
|
||||||
|
direction TB
|
||||||
|
B1[DataLoader] --> B2[BaseStrategy]
|
||||||
|
B2 --> B3[Transformer Forward]
|
||||||
|
B3 --> B4[Loss + Backward]
|
||||||
|
B4 --> B5[Gradient Accumulation]
|
||||||
|
B5 -->|every accum_steps| B6[Optimizer Step]
|
||||||
|
B6 --> B7[LR Scheduler]
|
||||||
|
B7 -->|next batch| B2
|
||||||
|
B6 --> B8[CheckpointCallback]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph C[Inference]
|
||||||
|
direction TB
|
||||||
|
C1[Checkpoint] --> C2[AutoModel]
|
||||||
|
C1 --> C3[AutoTokenizer]
|
||||||
|
C2 --> C4[InferenceEngine]
|
||||||
|
C3 --> C4
|
||||||
|
C4 --> C5[InferenceScheduler]
|
||||||
|
C5 --> C6[Transformer Forward]
|
||||||
|
C6 --> C7[sample]
|
||||||
|
C7 --> C8{End?}
|
||||||
|
C8 -->|No| C6
|
||||||
|
C8 -->|Yes| C9[Generated Text]
|
||||||
|
end
|
||||||
|
|
||||||
|
A --> B
|
||||||
|
B --> C
|
||||||
|
```
|
||||||
|
|
||||||
|
## Detailed Module Descriptions
|
||||||
|
|
||||||
|
### 1. Data Serialization (`astrai/dataset/storage.py` & `astrai/serialization.py`)
|
||||||
|
|
||||||
|
- **`save_h5`**: Saves tensors by groups as HDF5 files (`.h5`), each key maps to a list of tensors
|
||||||
|
- **`load_h5`**: Loads `.h5` files, returns `Dict[str, List[Tensor]]`, supports shared memory
|
||||||
|
- **`Checkpoint`**: Encapsulates model state dict + epoch + iteration; uses safetensors
|
||||||
|
|
||||||
|
### 2. Dataset Module
|
||||||
|
|
||||||
|
#### 2.1 Dataset (`dataset.py`)
|
||||||
|
- **`BaseDataset`**: Abstract base class for windowed sequence sampling
|
||||||
|
- **`BaseSegmentFetcher` / `MultiSegmentFetcher`**: Fetch tensor segments by index range
|
||||||
|
- **`DatasetFactory`**: Creates dataset instances by `train_type` (`seq`, `sft`, `dpo`, `grpo`)
|
||||||
|
- Data keys: `"sequence"` (SEQ), `"loss_mask"` (SFT), `"chosen_mask"/"rejected_mask"` (DPO), `"masks"` (GRPO)
|
||||||
|
|
||||||
|
#### 2.2 Sampler (`sampler.py`)
|
||||||
|
- **`ResumableDistributedSampler`**: Tracks `epoch` and `iter` for breakpoint resume; supports shuffle and drop_last
|
||||||
|
|
||||||
|
### 3. Model Module
|
||||||
|
|
||||||
|
#### 3.1 Transformer / AutoModel
|
||||||
|
- **`AutoModel`**: Base class with `from_pretrained()` / `save_pretrained()`
|
||||||
|
- **`Transformer`**: Decoder-only architecture, registered via `@AutoModel.register('transformer')`
|
||||||
|
- Embedding → N×DecoderBlock → RMSNorm → Linear lm_head
|
||||||
|
- RoPE position encoding, optional weight tying
|
||||||
|
|
||||||
|
#### 3.2 Submodules (`module.py`)
|
||||||
|
- **`DecoderBlock`**: GQA attention + residual + MLP + RMSNorm
|
||||||
|
- **`GQA`**: Grouped Query Attention (also `MLA` for multi-latent attention)
|
||||||
|
- **`MLP`**: `SiLU(gate(x)) * up(x)` → down projection
|
||||||
|
- **`RotaryEmbedding`**: RoPE complex cache (freqs_cis)
|
||||||
|
- **`RMSNorm`**: Layer normalization
|
||||||
|
|
||||||
|
### 4. Training Module
|
||||||
|
|
||||||
|
#### 4.1 Training Context (`train_context.py`)
|
||||||
|
- **`TrainContext`**: Dataclass holding model, optimizer, dataloader, strategy, scheduler, checkpoint state
|
||||||
|
- **`TrainContextBuilder`**: Builder pattern — takes checkpoint for resume, builds all components
|
||||||
|
|
||||||
|
#### 4.2 Trainer (`trainer.py`)
|
||||||
|
|
||||||
|
The training loop is nested: **epoch** → **batch** (with step phase interspersed):
|
||||||
|
|
||||||
|
```
|
||||||
|
on_train_begin
|
||||||
|
on_epoch_begin
|
||||||
|
for each accumulation window of batches: ← step phase
|
||||||
|
on_step_begin
|
||||||
|
for each batch in window: ← batch phase
|
||||||
|
on_batch_begin → strategy(batch) → loss → backward → on_batch_end
|
||||||
|
iteration += 1
|
||||||
|
on_step_end
|
||||||
|
optimizer.step() → zero_grad
|
||||||
|
|
||||||
|
on_epoch_end
|
||||||
|
on_train_end
|
||||||
|
```
|
||||||
|
|
||||||
|
Key points:
|
||||||
|
- `on_step_*` fires every `accumulation_steps` batches, wrapping optimizer step AFTER the hook
|
||||||
|
- `on_batch_*` fires every batch, wrapping loss computation
|
||||||
|
- `GradientClippingCallback` fires on `on_step_end`
|
||||||
|
- LR scheduler steps inline (no `SchedulerCallback` class)
|
||||||
|
|
||||||
|
#### 4.3 Strategy (`strategy.py`)
|
||||||
|
- **`SEQStrategy`**: Next-token prediction, cross-entropy with label smoothing
|
||||||
|
- **`SFTStrategy`**: Supervised fine-tuning with loss masking
|
||||||
|
- **`DPOStrategy`**: Direct Preference Optimization with reference model
|
||||||
|
- **`GRPOStrategy`**: Group Relative Policy Optimization with clipped ratio
|
||||||
|
|
||||||
|
#### 4.4 Scheduler (`schedule.py`)
|
||||||
|
- **`CosineScheduler`**: Cosine decay + linear warmup
|
||||||
|
- **`SGDRScheduler`**: Cosine annealing with warm restarts
|
||||||
|
- Created by `SchedulerFactory` and bound to optimizer
|
||||||
|
|
||||||
|
#### 4.5 Callbacks
|
||||||
|
- **`CheckpointCallback`**: Saves safetensors at `ckpt_interval` iterations
|
||||||
|
- **`ProgressBarCallback`**: tqdm progress display
|
||||||
|
- **`MetricLoggerCallback`**: Writes JSONL metrics to `{ckpt_dir}/logs/`
|
||||||
|
- **`GradientClippingCallback`**: `clip_grad_norm_` on `on_step_end`
|
||||||
|
|
||||||
|
### 5. Inference Module
|
||||||
|
|
||||||
|
#### 5.1 Inference Engine (`engine.py`)
|
||||||
|
- **`InferenceEngine`**: Facade over scheduler; provides `generate()`, `generate_with_request()`, `generate_async()`
|
||||||
|
- Accepts `prompt: str | List[str]`, returns generator (stream) or string (non-stream)
|
||||||
|
|
||||||
|
#### 5.2 Scheduler 4-Phase Loop (`scheduler.py`)
|
||||||
|
|
||||||
|
Background thread runs continuously:
|
||||||
|
|
||||||
|
```
|
||||||
|
1. Cleanup → Remove finished tasks, free KV cache pages
|
||||||
|
2. Refill → Pop from waiting_queue, alloc pages, add to active
|
||||||
|
3. Prefill → Group active tasks by prompt_len, run full forward pass
|
||||||
|
4. Decode → Pick largest same-position group, run single-token forward
|
||||||
|
```
|
||||||
|
|
||||||
|
- **`Task`**: Tracks prompt_ids, output_ids, status (PENDING/RUNNING/FINISHED/ABORTED)
|
||||||
|
- **`KVCache`**: Facade over `Allocator` + `PrefixCache` + `PagePool` + `Storage` for paged KV cache
|
||||||
|
- **`KvcacheView`**: Batch view bundling cache + page table for attention layers
|
||||||
|
- **`sample()`**: Temperature → top-k → top-p → multinomial
|
||||||
|
|
||||||
|
#### 5.3 Server (`server.py`)
|
||||||
|
- FastAPI with OpenAI `/v1/chat/completions` and Anthropic `/v1/messages` endpoints
|
||||||
|
- Streaming via SSE, health check at `/health`, stats at `/stats`
|
||||||
|
|
||||||
|
### 6. Tokenizer Module
|
||||||
|
|
||||||
|
- **`AutoTokenizer`**: Wraps HuggingFace tokenizers (BBPE); `encode`/`decode`/`apply_chat_template`
|
||||||
|
- **`ChatTemplate`**: Jinja2-based template rendering for multi-turn chat
|
||||||
|
|
||||||
|
### 7. Factory & Parallel
|
||||||
|
|
||||||
|
- **`Registry` / `BaseFactory`**: Decorator-based component registration
|
||||||
|
- **`spawn_parallel_fn`**: Multi-process DDP launcher with NCCL backend
|
||||||
|
- **`ParallelModel` / `ColumnParallelLinear` / `RowParallelLinear`**: Tensor model parallelism
|
||||||
|
|
||||||
|
## Training Data Flow — Detailed Steps
|
||||||
|
|
||||||
|
1. **Data Preparation**
|
||||||
|
- Raw text → token IDs via `AutoTokenizer.encode()`
|
||||||
|
- Save as `.h5` files (groups of tensor lists per data key)
|
||||||
|
|
||||||
|
2. **Dataset Loading**
|
||||||
|
- `BaseDataset.load()` calls `load_h5()`, builds `MultiSegmentFetcher`
|
||||||
|
- Sliding window of `window_size` with `stride` determines sample boundaries
|
||||||
|
|
||||||
|
3. **Sampling & Batching**
|
||||||
|
- `ResumableDistributedSampler` produces shuffled index sequences
|
||||||
|
- `DataLoader` fetches `[batch_size, window_size]` tensors via `__getitem__`
|
||||||
|
|
||||||
|
4. **Strategy Forward**
|
||||||
|
- Strategy receives batch, calls `Transformer.forward()` for logits
|
||||||
|
- Computes task-specific loss (cross-entropy, DPO, GRPO)
|
||||||
|
|
||||||
|
5. **Backward & Accumulation**
|
||||||
|
- `loss = raw_loss / accumulation_steps`
|
||||||
|
- `loss.backward()` accumulates gradients
|
||||||
|
- Every `accumulation_steps` batches: `optimizer.step()` → `zero_grad()`
|
||||||
|
- Every batch: `scheduler.step()` updates learning rate
|
||||||
|
|
||||||
|
6. **Checkpoint**
|
||||||
|
- `CheckpointCallback` saves `model.state_dict()` + metadata to safetensors at `ckpt_interval` iterations
|
||||||
|
- Does NOT save optimizer/scheduler state (resume resets those)
|
||||||
|
|
||||||
|
## Inference Data Flow — Detailed Steps
|
||||||
|
|
||||||
|
1. **Model Loading**
|
||||||
|
- `AutoModel.from_pretrained(path)` loads weights from safetensors
|
||||||
|
- `torch.inference_mode()` wraps generation
|
||||||
|
|
||||||
|
2. **Prompt Construction**
|
||||||
|
- Messages → `apply_chat_template(messages, tokenize=False)` → prompt string
|
||||||
|
- `tokenizer.encode(prompt)` → token IDs (truncated to `max_prompt_len`)
|
||||||
|
|
||||||
|
3. **Continuous Batching Loop**
|
||||||
|
- **Cleanup**: Finished tasks → `stream_callback(STOP)`, free KV pages
|
||||||
|
- **Refill**: Pop from waiting queue, `PagePool.task_alloc()` for prompt pages
|
||||||
|
- **Prefill**: Group by prompt length, run full forward with `start_pos=0`
|
||||||
|
- **Decode**: Pick position group with most tasks, single-token forward:
|
||||||
|
- Model forward → `logits` → `sample()` → next token ID
|
||||||
|
- Append to `output_ids`, update `output_tokens`
|
||||||
|
- `PagePool.task_alloc()` allocates pages as needed
|
||||||
|
- `stream_callback(token)` for streaming clients
|
||||||
|
|
||||||
|
4. **Output**
|
||||||
|
- `tokenizer.decode(output_ids)` → text
|
||||||
|
- Return to caller (streaming: token-by-token; non-streaming: complete string)
|
||||||
|
|
||||||
|
## Checkpoint & Serialization
|
||||||
|
|
||||||
|
- **Training Checkpoint**: safetensors weights + epoch/iteration metadata. Optimizer/scheduler state is NOT persisted.
|
||||||
|
- **Inference Loading**: `AutoModel.from_pretrained()` loads from the same safetensors format.
|
||||||
|
- **Dataset Serialization**: HDF5 with shared memory support for large-scale pre-training data.
|
||||||
|
|
||||||
|
> Document Update Time: 2026-05-14
|
||||||
@@ -0,0 +1,779 @@
|
|||||||
|
## 1. Why I Created This Project
|
||||||
|
|
||||||
|
There are many large language models on the market today, such as GPT, LLaMA, and others, with tens of billions or even hundreds of billions of parameters. But honestly, these models have extremely high hardware requirements, making them inaccessible for ordinary developers. I thought: **Can we create a model that is both useful and can run on ordinary computers?** This is also what most people currently hope for - a locally deployable AI project that achieves complete privatization while maintaining some level of intelligence.
|
||||||
|
|
||||||
|
Thus, the AstrAI project was born - 1B parameters, Chinese-English bilingual, supporting dialogue, text generation, and the training code is open source!
|
||||||
|
|
||||||
|
## 2. System Architecture
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
classDiagram
|
||||||
|
namespace config {
|
||||||
|
class ModelConfig {
|
||||||
|
+int vocab_size
|
||||||
|
+int dim
|
||||||
|
+int n_layers
|
||||||
|
+float norm_eps
|
||||||
|
+int dim_ffn
|
||||||
|
+bool tie_weight
|
||||||
|
+int max_len
|
||||||
|
+float rope_theta
|
||||||
|
+int n_heads
|
||||||
|
+int n_kv_heads
|
||||||
|
+bool use_qk_norm
|
||||||
|
+bool use_gated_attention
|
||||||
|
+load(config_path) ModelConfig
|
||||||
|
+save(config_path)
|
||||||
|
}
|
||||||
|
|
||||||
|
class TrainConfig {
|
||||||
|
+nn.Module model
|
||||||
|
+str strategy
|
||||||
|
+Dataset dataset
|
||||||
|
+Callable optimizer_fn
|
||||||
|
+Callable scheduler_fn
|
||||||
|
+int n_epoch
|
||||||
|
+int batch_size
|
||||||
|
+int accumulation_steps
|
||||||
|
+float max_grad_norm
|
||||||
|
+int start_epoch
|
||||||
|
+int start_batch
|
||||||
|
+str ckpt_dir
|
||||||
|
+int ckpt_interval
|
||||||
|
+int random_seed
|
||||||
|
+int num_workers
|
||||||
|
+int prefetch_factor
|
||||||
|
+bool pin_memory
|
||||||
|
+int nprocs
|
||||||
|
+str backend
|
||||||
|
+str master_addr
|
||||||
|
+str master_port
|
||||||
|
+Callable parallel_wrapper
|
||||||
|
+Callable state_dict_fn
|
||||||
|
+str device_type
|
||||||
|
+dict extra_kwargs
|
||||||
|
+validate()
|
||||||
|
}
|
||||||
|
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace dataset {
|
||||||
|
class BaseDataset {
|
||||||
|
+int window_size
|
||||||
|
+int stride
|
||||||
|
+BaseStorage storage
|
||||||
|
+load(load_path, storage_type, tokenizer)
|
||||||
|
+__getitem__(index)
|
||||||
|
+__len__()
|
||||||
|
}
|
||||||
|
|
||||||
|
class SEQDataset {
|
||||||
|
+__getitem__(index) Dict
|
||||||
|
}
|
||||||
|
|
||||||
|
class SFTDataset {
|
||||||
|
+__getitem__(index) Dict
|
||||||
|
}
|
||||||
|
|
||||||
|
class DPODataset {
|
||||||
|
+__getitem__(index) Dict
|
||||||
|
}
|
||||||
|
|
||||||
|
class GRPODataset {
|
||||||
|
+__getitem__(index) Dict
|
||||||
|
}
|
||||||
|
|
||||||
|
class BaseSegmentFetcher {
|
||||||
|
+List[Tensor] segments
|
||||||
|
+List[int] cum_lengths
|
||||||
|
+int total_length
|
||||||
|
+fetch_data(begin_idx, end_idx) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class BaseStorage {
|
||||||
|
+MultiSegmentFetcher _fetcher
|
||||||
|
+keys (property)
|
||||||
|
+load(load_path, tokenizer)
|
||||||
|
+fetch(begin, end, keys)
|
||||||
|
+__len__()
|
||||||
|
}
|
||||||
|
|
||||||
|
class H5Storage {
|
||||||
|
+load(load_path, tokenizer)
|
||||||
|
+fetch(begin, end, keys) Dict
|
||||||
|
+keys() List
|
||||||
|
}
|
||||||
|
|
||||||
|
class JSONStorage {
|
||||||
|
+load(load_path, tokenizer)
|
||||||
|
+fetch(begin, end, keys) Dict
|
||||||
|
+keys() List
|
||||||
|
}
|
||||||
|
|
||||||
|
class MultiSegmentFetcher {
|
||||||
|
+Dict multi_fetchers
|
||||||
|
+List multi_keys
|
||||||
|
+key_fetch(begin_idx, end_idx, keys) Dict
|
||||||
|
+fetch_data(begin_idx, end_idx) Dict
|
||||||
|
}
|
||||||
|
|
||||||
|
class ResumableDistributedSampler {
|
||||||
|
+int epoch
|
||||||
|
+int iter
|
||||||
|
}
|
||||||
|
|
||||||
|
class DatasetFactory {
|
||||||
|
+Registry _registry
|
||||||
|
+register(name) decorator
|
||||||
|
+create(train_type, window_size, stride) BaseDataset
|
||||||
|
+load(train_type, load_path, window_size, stride) BaseDataset
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace serialization {
|
||||||
|
class Checkpoint {
|
||||||
|
+dict state_dict
|
||||||
|
+int epoch
|
||||||
|
+int iteration
|
||||||
|
+save(save_dir)
|
||||||
|
+load(save_dir) Checkpoint
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace model {
|
||||||
|
class AutoModel {
|
||||||
|
+ModelConfig config
|
||||||
|
+Registry _registry
|
||||||
|
+register(model_type) decorator
|
||||||
|
+get_component_class(model_type) Type
|
||||||
|
+from_pretrained(path, disable_random_init) nn.Module
|
||||||
|
+save_pretrained(save_directory)
|
||||||
|
+to(*args, **kwargs) Self
|
||||||
|
}
|
||||||
|
|
||||||
|
class Transformer {
|
||||||
|
+ModelConfig config
|
||||||
|
+RotaryEmbedding rotary_embedding
|
||||||
|
+Embedding embed_tokens
|
||||||
|
+ModuleList layers
|
||||||
|
+RMSNorm norm
|
||||||
|
+Linear lm_head
|
||||||
|
+forward(input_ids, input_mask, paged_cache, position_ids) Tensor
|
||||||
|
+load_state_dict(state_dict)
|
||||||
|
+state_dict()
|
||||||
|
}
|
||||||
|
|
||||||
|
class DecoderBlock {
|
||||||
|
+GQA attention
|
||||||
|
+RMSNorm input_norm
|
||||||
|
+MLP mlp
|
||||||
|
+RMSNorm post_attention_norm
|
||||||
|
+forward(x, rotary_emb, attention_mask, paged_cache) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class GQA {
|
||||||
|
+int n_heads
|
||||||
|
+int n_kv_heads
|
||||||
|
+int head_dim
|
||||||
|
+Linear q_proj, k_proj, v_proj, o_proj
|
||||||
|
+RMSNorm q_norm, k_norm
|
||||||
|
+forward(x, rotary_emb, attn_mask, paged_cache) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class MLA {
|
||||||
|
+int n_heads
|
||||||
|
+int n_kv_heads
|
||||||
|
+int head_dim
|
||||||
|
+int kv_lora_rank
|
||||||
|
+int qk_nope_head_dim
|
||||||
|
+int qk_rope_head_dim
|
||||||
|
+Linear q_proj, kv_a_proj, kv_b_proj
|
||||||
|
+Linear o_proj
|
||||||
|
+RMSNorm kv_norm
|
||||||
|
+forward(x, rotary_emb, attn_mask, paged_cache) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class MLP {
|
||||||
|
+Linear up, gate, down
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class RMSNorm {
|
||||||
|
+Parameter weight
|
||||||
|
+float norm_eps
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class Linear {
|
||||||
|
+Parameter weight
|
||||||
|
+Parameter bias
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class RotaryEmbedding {
|
||||||
|
+int dim
|
||||||
|
+int max_len
|
||||||
|
+float base
|
||||||
|
+forward(x, position_ids=None) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class Embedding {
|
||||||
|
+Parameter weight
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace tokenize {
|
||||||
|
class AutoTokenizer {
|
||||||
|
+vocab_size int
|
||||||
|
+encode(tokens, out_ids, add_special_tokens) List[int]
|
||||||
|
+decode(tokens, skip_special_tokens) str
|
||||||
|
+__getattr__(name) Any (bos_id, eos_id, pad_id, stop_ids)
|
||||||
|
+apply_chat_template(messages, tokenize) Union[str, List[int]]
|
||||||
|
+set_chat_template(template)
|
||||||
|
+load(path)
|
||||||
|
+from_pretrained(path) AutoTokenizer
|
||||||
|
+save_pretrained(save_path)
|
||||||
|
}
|
||||||
|
|
||||||
|
class ChatTemplate {
|
||||||
|
+String template_str
|
||||||
|
+render(messages, system_prompt, **extra_variables) str
|
||||||
|
+from_string(template) ChatTemplate
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace factory {
|
||||||
|
class Registry {
|
||||||
|
+Dict _entries
|
||||||
|
+register(name, component_cls, category, priority)
|
||||||
|
+get(name) Type
|
||||||
|
+list_names() List[str]
|
||||||
|
}
|
||||||
|
|
||||||
|
class BaseFactory {
|
||||||
|
+Registry _registry
|
||||||
|
+register(name, category, priority) decorator
|
||||||
|
+create(name, *args, **kwargs) T
|
||||||
|
+list_registered() list
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace trainer {
|
||||||
|
class Trainer {
|
||||||
|
+TrainConfig train_config
|
||||||
|
+List[TrainCallback] callbacks
|
||||||
|
+train(checkpoint)
|
||||||
|
+_build_context(checkpoint) TrainContext
|
||||||
|
+_get_default_callbacks() List[TrainCallback]
|
||||||
|
}
|
||||||
|
|
||||||
|
class TrainContext {
|
||||||
|
+nn.Module model
|
||||||
|
+BaseStrategy strategy
|
||||||
|
+DataLoader dataloader
|
||||||
|
+Optimizer optimizer
|
||||||
|
+LRScheduler scheduler
|
||||||
|
+Checkpoint checkpoint
|
||||||
|
+int epoch
|
||||||
|
+int iteration
|
||||||
|
+float loss
|
||||||
|
+int world_size
|
||||||
|
+int rank
|
||||||
|
}
|
||||||
|
|
||||||
|
class TrainContextBuilder {
|
||||||
|
+TrainConfig config
|
||||||
|
+with_checkpoint(checkpoint) TrainContextBuilder
|
||||||
|
+build() TrainContext
|
||||||
|
}
|
||||||
|
|
||||||
|
class BaseStrategy {
|
||||||
|
+nn.Module model
|
||||||
|
+str device
|
||||||
|
+compute_loss(batch) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class StrategyFactory {
|
||||||
|
+Registry _registry
|
||||||
|
+register(name) decorator
|
||||||
|
+create(model, train_type, device, **kwargs) BaseStrategy
|
||||||
|
}
|
||||||
|
|
||||||
|
class SEQStrategy {
|
||||||
|
+float label_smoothing
|
||||||
|
+compute_loss(batch) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class SFTStrategy {
|
||||||
|
+float label_smoothing
|
||||||
|
+compute_loss(batch) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class DPOStrategy {
|
||||||
|
+nn.Module ref_model
|
||||||
|
+float beta
|
||||||
|
+str reduction
|
||||||
|
+compute_loss(batch) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class GRPOStrategy {
|
||||||
|
+nn.Module ref_model
|
||||||
|
+float clip_eps
|
||||||
|
+float kl_coef
|
||||||
|
+int group_size
|
||||||
|
+str reduction
|
||||||
|
+int sync_interval
|
||||||
|
+compute_loss(batch) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class BaseScheduler {
|
||||||
|
+get_lr() List[float]
|
||||||
|
+step()
|
||||||
|
}
|
||||||
|
|
||||||
|
class SchedulerFactory {
|
||||||
|
+Registry _registry
|
||||||
|
+register(name) decorator
|
||||||
|
+create(optimizer, schedule_type, **kwargs) BaseScheduler
|
||||||
|
}
|
||||||
|
|
||||||
|
class CosineScheduler {
|
||||||
|
+int warmup_steps
|
||||||
|
+int lr_decay_steps
|
||||||
|
+float min_rate
|
||||||
|
}
|
||||||
|
|
||||||
|
class SGDRScheduler {
|
||||||
|
+int warmup_steps
|
||||||
|
+int cycle_length
|
||||||
|
+float min_rate
|
||||||
|
+int t_mult
|
||||||
|
}
|
||||||
|
|
||||||
|
class TrainCallback {
|
||||||
|
+on_train_begin(context)
|
||||||
|
+on_train_end(context)
|
||||||
|
+on_epoch_begin(context)
|
||||||
|
+on_epoch_end(context)
|
||||||
|
+on_step_begin(context)
|
||||||
|
+on_step_end(context)
|
||||||
|
+on_batch_begin(context)
|
||||||
|
+on_batch_end(context)
|
||||||
|
+on_error(context)
|
||||||
|
}
|
||||||
|
|
||||||
|
class GradientClippingCallback {
|
||||||
|
+float max_grad_norm
|
||||||
|
+on_step_begin(context)
|
||||||
|
}
|
||||||
|
|
||||||
|
class CheckpointCallback {
|
||||||
|
+str save_dir
|
||||||
|
+int interval
|
||||||
|
+_save_checkpoint(context)
|
||||||
|
+on_batch_end(context)
|
||||||
|
+on_train_end(context)
|
||||||
|
+on_error(context)
|
||||||
|
}
|
||||||
|
|
||||||
|
class ProgressBarCallback {
|
||||||
|
+int num_epoch
|
||||||
|
+on_epoch_begin(context)
|
||||||
|
+on_batch_end(context)
|
||||||
|
+on_epoch_end(context)
|
||||||
|
}
|
||||||
|
|
||||||
|
class MetricLoggerCallback {
|
||||||
|
+str log_dir
|
||||||
|
+int save_interval
|
||||||
|
+on_batch_end(context)
|
||||||
|
+on_train_end(context)
|
||||||
|
}
|
||||||
|
|
||||||
|
class CallbackFactory {
|
||||||
|
+Registry _registry
|
||||||
|
+register(name) decorator
|
||||||
|
+create(name, **kwargs) TrainCallback
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace inference {
|
||||||
|
class InferenceEngine {
|
||||||
|
+nn.Module model
|
||||||
|
+AutoTokenizer tokenizer
|
||||||
|
+InferenceScheduler scheduler
|
||||||
|
+generate(prompt, stream, max_tokens, temperature, top_p, top_k) Union[Generator, str, List[str]]
|
||||||
|
+generate_with_request(request) Union[Generator, str, List[str]]
|
||||||
|
+generate_async(prompt, max_tokens, temperature, top_p, top_k) AsyncGenerator
|
||||||
|
+get_stats() Dict
|
||||||
|
+shutdown()
|
||||||
|
}
|
||||||
|
|
||||||
|
class InferenceScheduler {
|
||||||
|
+nn.Module model
|
||||||
|
+AutoTokenizer tokenizer
|
||||||
|
+KVCache _page_cache
|
||||||
|
+int max_batch_size
|
||||||
|
+int max_seq_len
|
||||||
|
+int max_prompt_len
|
||||||
|
+int page_size
|
||||||
|
+TaskManager _task_mgr
|
||||||
|
+add_task(prompt, max_tokens, temperature, top_p, top_k, stream_callback) str
|
||||||
|
+remove_task(task_id)
|
||||||
|
+start()
|
||||||
|
+stop()
|
||||||
|
+get_stats() Dict
|
||||||
|
}
|
||||||
|
|
||||||
|
class Allocator {
|
||||||
|
+int _free_mask
|
||||||
|
+int refs_count
|
||||||
|
+LRU _lru
|
||||||
|
+alloc() int
|
||||||
|
+free(idx, keep_cached)
|
||||||
|
+inc_ref(idx)
|
||||||
|
+touch(idx)
|
||||||
|
+ref_count(idx) int
|
||||||
|
}
|
||||||
|
|
||||||
|
class PrefixCache {
|
||||||
|
+int _page_size
|
||||||
|
+evict(page_idx)
|
||||||
|
+has_page(idx) bool
|
||||||
|
+lookup(token_ids) List[int]
|
||||||
|
+record(page_idx, token_ids, logical_page_idx)
|
||||||
|
}
|
||||||
|
|
||||||
|
class PagePool {
|
||||||
|
-Allocator _alloc
|
||||||
|
-PrefixCache _prefix
|
||||||
|
+alloc() int
|
||||||
|
+free(idx)
|
||||||
|
+inc_ref(idx)
|
||||||
|
+lookup(token_ids) List[int]
|
||||||
|
+record(page_idx, token_ids, logical_page_idx)
|
||||||
|
}
|
||||||
|
|
||||||
|
class Storage {
|
||||||
|
+int n_layers
|
||||||
|
+int page_size
|
||||||
|
+int head_dim
|
||||||
|
+int n_kv_heads
|
||||||
|
+Tensor k_cache
|
||||||
|
+Tensor v_cache
|
||||||
|
+write(layer_id, page_table, start_pos, k, v)
|
||||||
|
+gather(layer_id, page_table, total_len) Tuple[Tensor, Tensor]
|
||||||
|
}
|
||||||
|
|
||||||
|
class KVCache {
|
||||||
|
-PagePool _pool
|
||||||
|
-Storage _storage
|
||||||
|
-TaskTable _table
|
||||||
|
+int page_size
|
||||||
|
+task_alloc(task_id, prompt_ids) bool
|
||||||
|
+task_free(task_id)
|
||||||
|
+task_extend(task_id, pos) bool
|
||||||
|
+task_cached(task_id) int
|
||||||
|
+task_record_hashes(task_id, prompt_ids, start_logical_page)
|
||||||
|
+make_table_tensor(task_ids, device) Tensor
|
||||||
|
+bind(page_table, total_len) KvcacheView
|
||||||
|
}
|
||||||
|
|
||||||
|
class KvcacheView {
|
||||||
|
-Storage _storage
|
||||||
|
+Tensor _page_table
|
||||||
|
+int _total_len
|
||||||
|
+write(layer_id, k, v)
|
||||||
|
+gather(layer_id) Tuple[Tensor, Tensor]
|
||||||
|
}
|
||||||
|
|
||||||
|
class TaskTable {
|
||||||
|
+set(task_id, page_table, cached)
|
||||||
|
+get(task_id) List[int]
|
||||||
|
+get_cached(task_id) int
|
||||||
|
+get_ref(task_id) List[int]
|
||||||
|
+pop(task_id) Tuple[List[int], int]
|
||||||
|
+table_tensor(task_ids, device) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class Task {
|
||||||
|
+str task_id
|
||||||
|
+List prompt_ids
|
||||||
|
+int max_tokens
|
||||||
|
+float temperature
|
||||||
|
+float top_p
|
||||||
|
+int top_k
|
||||||
|
+TaskStatus status
|
||||||
|
+List output_ids
|
||||||
|
+int input_tokens
|
||||||
|
+int output_tokens
|
||||||
|
+float arrival_time
|
||||||
|
+float finish_time
|
||||||
|
+Callable stream_callback
|
||||||
|
+int next_pos
|
||||||
|
+is_finished(stop_ids) bool
|
||||||
|
}
|
||||||
|
|
||||||
|
class TaskStatus {
|
||||||
|
<<enumeration>>
|
||||||
|
PENDING
|
||||||
|
RUNNING
|
||||||
|
FINISHED
|
||||||
|
ABORTED
|
||||||
|
}
|
||||||
|
|
||||||
|
class GenerationRequest {
|
||||||
|
+List[Dict] messages
|
||||||
|
+int top_k
|
||||||
|
+float top_p
|
||||||
|
+float temperature
|
||||||
|
+Optional[int] max_tokens
|
||||||
|
+bool stream
|
||||||
|
}
|
||||||
|
|
||||||
|
class BaseSamplingStrategy {
|
||||||
|
<<abstract>>
|
||||||
|
+apply(logits, filter_value) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class TemperatureStrategy {
|
||||||
|
+float temperature
|
||||||
|
+apply(logits, filter_value) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class TopKStrategy {
|
||||||
|
+int top_k
|
||||||
|
+apply(logits, filter_value) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class TopPStrategy {
|
||||||
|
+float top_p
|
||||||
|
+apply(logits, filter_value) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class SamplingPipeline {
|
||||||
|
+List strategies
|
||||||
|
+apply(logits, filter_value) Tensor
|
||||||
|
+sample(logits, filter_value) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class GenerateResult {
|
||||||
|
+List[Tuple[int, str]] tokens
|
||||||
|
+List[str] results
|
||||||
|
+List[bool] _done
|
||||||
|
+append(token, idx)
|
||||||
|
+get_results() List[str]
|
||||||
|
+pop_all() List[str]
|
||||||
|
+wait(timeout) bool
|
||||||
|
+wait_completion()
|
||||||
|
}
|
||||||
|
|
||||||
|
class ChatMessage {
|
||||||
|
+str role
|
||||||
|
+str content
|
||||||
|
}
|
||||||
|
|
||||||
|
class ChatCompletionRequest {
|
||||||
|
+List[ChatMessage] messages
|
||||||
|
+float temperature
|
||||||
|
+float top_p
|
||||||
|
+int top_k
|
||||||
|
+int max_tokens
|
||||||
|
+bool stream
|
||||||
|
+Optional[str] stop
|
||||||
|
+Optional[int] n
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace parallel {
|
||||||
|
class Functions {
|
||||||
|
+spawn_parallel_fn(fn, nprocs)
|
||||||
|
+setup_parallel(rank, world_size, backend, master_addr, master_port, device_type)
|
||||||
|
+get_current_device() str
|
||||||
|
+get_world_size() int
|
||||||
|
+get_rank() int
|
||||||
|
}
|
||||||
|
|
||||||
|
class ParallelModel {
|
||||||
|
+dist.ProcessGroup process_group
|
||||||
|
+int rank
|
||||||
|
+int world_size
|
||||||
|
}
|
||||||
|
|
||||||
|
class ColumnParallelLinear {
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class RowParallelLinear {
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
%% Relationships
|
||||||
|
TrainConfig --> BaseDataset : uses
|
||||||
|
TrainConfig ..> BaseStrategy : selects
|
||||||
|
StrategyFactory ..> BaseStrategy : creates
|
||||||
|
BaseStrategy <|-- SEQStrategy
|
||||||
|
BaseStrategy <|-- SFTStrategy
|
||||||
|
BaseStrategy <|-- DPOStrategy
|
||||||
|
BaseStrategy <|-- GRPOStrategy
|
||||||
|
DPOStrategy --> Transformer : uses
|
||||||
|
GRPOStrategy --> Transformer : uses
|
||||||
|
Trainer --> TrainConfig : uses
|
||||||
|
Trainer --> TrainContextBuilder : uses
|
||||||
|
Trainer --> TrainCallback : manages
|
||||||
|
TrainContextBuilder --> TrainContext : creates
|
||||||
|
TrainContextBuilder --> StrategyFactory : uses
|
||||||
|
Checkpoint ..> Checkpoint : serializes
|
||||||
|
TrainContext --> Checkpoint : manages
|
||||||
|
TrainContext --> BaseStrategy : uses
|
||||||
|
TrainContext --> BaseScheduler : uses
|
||||||
|
SchedulerFactory ..> BaseScheduler : creates
|
||||||
|
BaseScheduler <|-- CosineScheduler
|
||||||
|
BaseScheduler <|-- SGDRScheduler
|
||||||
|
CallbackFactory ..> TrainCallback : creates
|
||||||
|
TrainCallback <|-- GradientClippingCallback
|
||||||
|
TrainCallback <|-- CheckpointCallback
|
||||||
|
TrainCallback <|-- ProgressBarCallback
|
||||||
|
TrainCallback <|-- MetricLoggerCallback
|
||||||
|
PagePool --> Allocator : composes
|
||||||
|
PagePool --> PrefixCache : composes
|
||||||
|
KVCache --> PagePool : composes
|
||||||
|
KVCache --> Storage : composes
|
||||||
|
KVCache --> TaskTable : composes
|
||||||
|
KvcacheView --> Storage : wraps
|
||||||
|
InferenceEngine --> InferenceScheduler : uses
|
||||||
|
InferenceEngine --> GenerationRequest : uses
|
||||||
|
InferenceEngine --> GenerateResult : creates
|
||||||
|
InferenceScheduler --> Task : manages
|
||||||
|
InferenceScheduler --> TaskStatus : uses
|
||||||
|
InferenceScheduler --> KVCache : uses
|
||||||
|
InferenceScheduler --> Transformer : uses
|
||||||
|
Task --> TaskStatus : uses
|
||||||
|
InferenceEngine --> Transformer : uses
|
||||||
|
BaseSamplingStrategy <|-- TemperatureStrategy
|
||||||
|
BaseSamplingStrategy <|-- TopKStrategy
|
||||||
|
BaseSamplingStrategy <|-- TopPStrategy
|
||||||
|
SamplingPipeline --> BaseSamplingStrategy : composes
|
||||||
|
BaseDataset <|-- SEQDataset
|
||||||
|
BaseDataset <|-- SFTDataset
|
||||||
|
BaseDataset <|-- DPODataset
|
||||||
|
BaseDataset <|-- GRPODataset
|
||||||
|
DatasetFactory ..> BaseDataset : creates
|
||||||
|
BaseStorage <|-- H5Storage
|
||||||
|
BaseStorage <|-- JSONStorage
|
||||||
|
BaseDataset --> BaseStorage : uses
|
||||||
|
MultiSegmentFetcher --> BaseSegmentFetcher : uses
|
||||||
|
AutoModel <|-- Transformer
|
||||||
|
AutoModel --> ModelConfig : contains
|
||||||
|
Transformer --> DecoderBlock : uses
|
||||||
|
Transformer --> RotaryEmbedding : uses
|
||||||
|
Transformer --> Embedding : uses
|
||||||
|
DecoderBlock --> GQA : uses
|
||||||
|
DecoderBlock --> MLP : uses
|
||||||
|
DecoderBlock --> RMSNorm : uses
|
||||||
|
TrainContextBuilder --> ResumableDistributedSampler : creates
|
||||||
|
ResumableDistributedSampler --> BaseDataset : samples
|
||||||
|
ParallelModel <|-- RowParallelLinear
|
||||||
|
ParallelModel <|-- ColumnParallelLinear
|
||||||
|
AutoTokenizer --> ChatTemplate : uses
|
||||||
|
BaseFactory <|-- AutoModel
|
||||||
|
BaseFactory <|-- DatasetFactory
|
||||||
|
BaseFactory <|-- StrategyFactory
|
||||||
|
BaseFactory <|-- SchedulerFactory
|
||||||
|
BaseFactory <|-- CallbackFactory
|
||||||
|
```
|
||||||
|
|
||||||
|
### Module Overview
|
||||||
|
|
||||||
|
| Module | Components | Description |
|
||||||
|
|--------|------------|-------------|
|
||||||
|
| **astrai.config** | ModelConfig, TrainConfig | Configuration management |
|
||||||
|
| **astrai.dataset** | BaseDataset, SEQDataset, SFTDataset, DPODataset, GRPODataset, BaseStorage, H5Storage, JSONStorage, BaseSegmentFetcher, MultiSegmentFetcher, ResumableDistributedSampler, DatasetFactory, save_h5, load_h5 | Dataset loading and management |
|
||||||
|
| **astrai.serialization** | Checkpoint | Model serialization and checkpoint management |
|
||||||
|
| **astrai.model** | AutoModel, Transformer, DecoderBlock, GQA, MLA, MLP, RMSNorm, Linear, RotaryEmbedding, Embedding | Neural network model |
|
||||||
|
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
||||||
|
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy, StrategyFactory, BaseScheduler, SchedulerFactory, TrainCallback, CallbackFactory | Training workflow management |
|
||||||
|
| **astrai.inference** | InferenceEngine, InferenceScheduler, KVCache, KvcacheView, Allocator, PrefixCache, PagePool, Storage, TaskTable, Task, TaskStatus, GenerationRequest, BaseSamplingStrategy, TemperatureStrategy, TopKStrategy, TopPStrategy, SamplingPipeline, ChatMessage, ChatCompletionRequest | Inference service with continuous batching and paged KV cache |
|
||||||
|
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank, get_world_size, get_current_device, ParallelModel, ColumnParallelLinear, RowParallelLinear | Distributed parallel |
|
||||||
|
| **astrai.factory** | Registry, BaseFactory | Generic component registration |
|
||||||
|
|
||||||
|
### Design Patterns
|
||||||
|
|
||||||
|
| Pattern | Classes | Purpose |
|
||||||
|
|---------|---------|---------|
|
||||||
|
| **Strategy** | `BaseStrategy`, `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy`, `StrategyFactory` | Flexible training strategy switching, supports SEQ/SFT/DPO/GRPO |
|
||||||
|
| **Builder** | `TrainContextBuilder` | Chain-building training context, step-by-step initialization of components |
|
||||||
|
| **Factory** | `StrategyFactory`, `SchedulerFactory`, `DatasetFactory`, `CallbackFactory`, `BaseFactory` | Decorator registration mechanism, dynamically create training strategies, schedulers, datasets, and callbacks |
|
||||||
|
| **Observer** | `TrainCallback`, `CallbackFactory` | Callback mechanism for training process monitoring (checkpoint, early stopping, metrics) |
|
||||||
|
| **Context** | `TrainContext` | Training process state container with model, optimizer, scheduler and checkpoint |
|
||||||
|
| **Registry** | `BaseFactory`, `Registry` | Generic component registration with category and priority support |
|
||||||
|
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with O(1) alloc/free via bitmask + LRU eviction |
|
||||||
|
| **Strategy (Sampling)** | `BaseSamplingStrategy`, `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations with temperature, top-k, top-p |
|
||||||
|
| **Producer-Consumer** | `InferenceScheduler`, `Task`, `waiting_queue`, `active_tasks` | Continuous batching with dynamic task queue management |
|
||||||
|
| **Event-Driven** | `threading.Event`, `_task_event` | Non-blocking wait mechanism for task scheduling using Python's `threading` module |
|
||||||
|
| **AutoModel Registry** | `AutoModel`, `Transformer` | Model type registration and dynamic loading via decorator pattern |
|
||||||
|
| **Generator Pattern** | `GenerateResult`, `GenerationRequest` | Event-based result notification for streaming/non-streaming generation |
|
||||||
|
|
||||||
|
### Core Relationships
|
||||||
|
|
||||||
|
1. **Configuration → Training**: `TrainConfig` holds model, dataset, optimizer_fn, scheduler_fn and other training configuration references
|
||||||
|
2. **Training Flow**: `Trainer` → `TrainContextBuilder` → `TrainContext`, uses `BaseStrategy` to compute loss
|
||||||
|
3. **Strategy Selection**: `StrategyFactory` creates corresponding strategy instance based on `train_type`
|
||||||
|
4. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `Transformer`, uses `KVCache` (backed by `Allocator` + `PrefixCache` + `PagePool` + `Storage`) for paged KV cache management and `SamplingPipeline` for efficient continuous batching with streaming/non-streaming
|
||||||
|
5. **Distributed Support**: `spawn_parallel_fn` and `setup_parallel` provide multi-process training capability for `Trainer`
|
||||||
|
6. **Dataset Loading**: `DatasetFactory` creates datasets (SEQDataset, SFTDataset, DPODataset, GRPODataset), supports HDF5 loading via `BaseSegmentFetcher` and `MultiSegmentFetcher`
|
||||||
|
7. **Checkpoint Management**: `Checkpoint` handles model state serialization/deserialization with safetensors
|
||||||
|
8. **Scheduler Support**: `SchedulerFactory` creates learning rate schedulers (CosineScheduler, SGDRScheduler)
|
||||||
|
9. **AutoModel Loading**: `AutoModel.from_pretrained()` dynamically loads model based on `config.json` model_type, uses `Registry` pattern for model type registration
|
||||||
|
|
||||||
|
## 3. Training Process
|
||||||
|
|
||||||
|
The common training process for large language models (LLM) typically includes three stages: **Pre-training (SEQ)**, **Supervised Fine-Tuning (SFT)**, and **Reinforcement Learning from Human Feedback (DPO/GRPO)**. This system is designed to support seamless end-to-end flow, achieving efficient switching and state management of different training stages through modular strategies.
|
||||||
|
|
||||||
|
### Core Formulas
|
||||||
|
|
||||||
|
**Pre-training (SEQ):**
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{\text{PT}} = - \sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta)
|
||||||
|
$$
|
||||||
|
|
||||||
|
**SFT:**
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{\text{SFT}} = - \sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
|
||||||
|
$$
|
||||||
|
|
||||||
|
**DPO:**
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{\text{DPO}} = -\mathbb{E}_{(x, y_w, y_l) \sim D} \left[ \log \sigma\left( \beta \log \frac{\pi_\theta(y_w \mid x)}{\pi_{\text{ref}}(y_w \mid x)} - \beta \log \frac{\pi_\theta(y_l \mid x)}{\pi_{\text{ref}}(y_l \mid x)} \right) \right]
|
||||||
|
$$
|
||||||
|
|
||||||
|
**GRPO:**
|
||||||
|
|
||||||
|
GRPO (Group Relative Policy Optimization) computes advantages from multiple responses to the same prompt, then optimizes using a PPO-style clipped objective:
|
||||||
|
|
||||||
|
$$
|
||||||
|
\text{Advantage}_i = \frac{r_i - \mu}{\sigma + \epsilon}
|
||||||
|
$$
|
||||||
|
|
||||||
|
Where $r_i$ is the reward for the $i$-th response, $\mu$ and $\sigma$ are the mean and standard deviation of group rewards.
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{\text{GRPO}} = -\mathbb{E} \left[ \min\left( \frac{\pi_\theta(a|s)}{\pi_{\text{ref}}(a|s)} \cdot A, \text{clip}\left(\frac{\pi_\theta(a|s)}{\pi_{\text{ref}}(a|s)}, 1-\epsilon, 1+\epsilon\right) \cdot A \right) \right] + \lambda \cdot D_{KL}
|
||||||
|
$$
|
||||||
|
|
||||||
|
The KL divergence term uses mean squared error approximation:
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{KL} = \lambda \cdot \mathbb{E} \left[ (\log \pi_\theta - \log \pi_{\text{ref}})^2 \right]
|
||||||
|
$$
|
||||||
|
|
||||||
|
The final loss is the sum of both: $L = L_{\text{policy}} + L_{KL}$
|
||||||
|
|
||||||
|
Through the above three-stage progressive training, the model completes its evolution from a general language foundation to a specialized, highly-aligned dialogue intelligence.
|
||||||
|
|
||||||
|
> Document Update Time: 2026-05-14
|
||||||
+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 multiple layers of Transformer blocks, with 1.0 billion parameters. Transformer is an autoregressive model that calculates the relationship between all previous tokens to obtain the probability distribution of the next token.
|
||||||
|
|
||||||
### 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 ...]
|
||||||
|
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_tokens=None,
|
||||||
|
stream=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Generate with streaming
|
||||||
|
for token in engine.generate_with_request(request):
|
||||||
|
print(token, end="", flush=True)
|
||||||
|
```
|
||||||
|
|
||||||
|
The continuous batching feature allows dynamic batch composition where new requests can join at any time and completed requests are released immediately.
|
||||||
|
|
||||||
|
## HTTP API Usage
|
||||||
|
|
||||||
|
The inference server provides HTTP endpoints for remote inference. Start the server first:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python -m scripts.tools.server --port 8000
|
||||||
|
```
|
||||||
|
|
||||||
|
### OpenAI-Compatible Endpoint
|
||||||
|
|
||||||
|
The server provides an OpenAI-compatible chat completion endpoint at `/v1/chat/completions`:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"messages": [
|
||||||
|
{"role": "system", "content": "You are a helpful assistant."},
|
||||||
|
{"role": "user", "content": "Hello, how are you?"}
|
||||||
|
],
|
||||||
|
"temperature": 0.8,
|
||||||
|
"max_tokens": 2048,
|
||||||
|
"stream": false
|
||||||
|
}'
|
||||||
|
```
|
||||||
|
|
||||||
|
**Request Parameters:**
|
||||||
|
| Parameter | Type | Default | Description |
|
||||||
|
|-----------|------|---------|-------------|
|
||||||
|
| `messages` | List[dict] | Required | Chat messages with role and content |
|
||||||
|
| `temperature` | float | 1.0 | Sampling temperature (0.0-2.0) |
|
||||||
|
| `top_p` | float | 1.0 | Nucleus sampling threshold |
|
||||||
|
| `top_k` | int | 50 | Top-k sampling parameter |
|
||||||
|
| `max_tokens` | int | 1024 | Maximum tokens to generate |
|
||||||
|
| `stream` | bool | false | Enable streaming response |
|
||||||
|
|
||||||
|
**Response (non-streaming):**
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"id": "chatcmpl-1234567890",
|
||||||
|
"object": "chat.completion",
|
||||||
|
"created": 1234567890,
|
||||||
|
"model": "astrai",
|
||||||
|
"choices": [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"message": {"role": "assistant", "content": "Hello! I'm doing well..."},
|
||||||
|
"finish_reason": "stop"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"usage": {
|
||||||
|
"prompt_tokens": 20,
|
||||||
|
"completion_tokens": 15,
|
||||||
|
"total_tokens": 35
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Streaming Response
|
||||||
|
|
||||||
|
Enable streaming for real-time token-by-token output:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"messages": [{"role": "user", "content": "Write a story"}],
|
||||||
|
"stream": true,
|
||||||
|
"max_tokens": 500
|
||||||
|
}'
|
||||||
|
```
|
||||||
|
|
||||||
|
The server uses Server-Sent Events (SSE) with content type `text/event-stream`.
|
||||||
|
|
||||||
|
### Anthropic-Compatible Endpoint
|
||||||
|
|
||||||
|
The server also provides an Anthropic-compatible endpoint at `/v1/messages`:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"model": "astrai",
|
||||||
|
"system": "You are a helpful assistant.",
|
||||||
|
"messages": [{"role": "user", "content": "Hello, how are you?"}],
|
||||||
|
"max_tokens": 2048
|
||||||
|
}'
|
||||||
|
```
|
||||||
|
|
||||||
|
Response:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"id": "msg_abc123...",
|
||||||
|
"type": "message",
|
||||||
|
"role": "assistant",
|
||||||
|
"model": "astrai",
|
||||||
|
"content": [{"type": "text", "text": "Hello! I am doing well..."}],
|
||||||
|
"stop_reason": "end_turn",
|
||||||
|
"stop_sequence": null,
|
||||||
|
"usage": {"input_tokens": 20, "output_tokens": 15}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Streaming:
|
||||||
|
```bash
|
||||||
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"model": "astrai",
|
||||||
|
"system": "You are a helpful assistant.",
|
||||||
|
"messages": [{"role": "user", "content": "Write a short poem"}],
|
||||||
|
"max_tokens": 500,
|
||||||
|
"stream": true
|
||||||
|
}'
|
||||||
|
```
|
||||||
|
|
||||||
|
Supports `stop_sequences` for early termination:
|
||||||
|
```bash
|
||||||
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"model": "astrai",
|
||||||
|
"messages": [{"role": "user", "content": "Write a story"}],
|
||||||
|
"max_tokens": 500,
|
||||||
|
"stop_sequences": ["The end", "THE END"]
|
||||||
|
}'
|
||||||
|
```
|
||||||
|
|
||||||
|
### Health Check
|
||||||
|
|
||||||
|
Monitor server and model status:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl http://localhost:8000/health
|
||||||
|
# {"status": "ok", "model_loaded": true}
|
||||||
|
|
||||||
|
curl http://localhost:8000/stats
|
||||||
|
# {"total_tasks": 10, "total_tokens": 5000, "active_tasks": 1, "waiting_queue": 0}
|
||||||
|
```
|
||||||
|
|
||||||
|
> Document Update Time: 2026-05-14
|
||||||
@@ -1,27 +0,0 @@
|
|||||||
## kv_cache 实现
|
|
||||||
|
|
||||||
根据注意力的计算公式
|
|
||||||
|
|
||||||
$$
|
|
||||||
\begin{align*}
|
|
||||||
o_i &= \sum_j s_{ij} v_{j} \newline
|
|
||||||
s_{ij} &= \text{softmax}\left( \frac{q_{i} k_{j}}{\sqrt{d_k}} \right)
|
|
||||||
\end{align*}
|
|
||||||
$$
|
|
||||||
|
|
||||||
由于模型是自回归模型, 我们只用求序列最后一个部分,也就是说 $ i $ 的下标是确定的, 是序列最后一个元素, 我们求的是 $o_{n} $
|
|
||||||
|
|
||||||
$$
|
|
||||||
\begin{align*}
|
|
||||||
o_n &= \sum_j s_{j}v_{j} \newline
|
|
||||||
s_j &= \text{softmax}\left(\frac{q_n k_{j}}{\sqrt{d_k}} \right)
|
|
||||||
\end{align*}
|
|
||||||
$$
|
|
||||||
|
|
||||||
如果我们把式子展开
|
|
||||||
|
|
||||||
$$
|
|
||||||
o_n = \sum_j \text{softmax}\left(\frac{q_n k_{j}}{\sqrt{d_k}}\right)v_{j}
|
|
||||||
$$
|
|
||||||
|
|
||||||
以上表达式只有k和v存在长度下标, 而 $q$ 没有, 所以计算过程中 $q$ 的输入是确定的上次输入的最后一个token, 而 $k, v$ 是需要对不同长度的部分进行缓存的,同时缓存的时候应该注意位置编码的计算应该在kvcache的计算之前进行,否则会存在位置编码的计算错误
|
|
||||||
@@ -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_tokens` | Maximum generation length | None (unlimited) |
|
||||||
|
| `stream` | Whether to stream output | False |
|
||||||
|
|
||||||
|
### Usage Example
|
||||||
|
|
||||||
|
```python
|
||||||
|
import torch
|
||||||
|
from astrai.model import AutoModel
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
from astrai.inference import InferenceEngine, GenerationRequest
|
||||||
|
|
||||||
|
# Load model using AutoModel
|
||||||
|
model = AutoModel.from_pretrained("your_model_dir")
|
||||||
|
|
||||||
|
# Load tokenizer
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained("your_model_dir")
|
||||||
|
|
||||||
|
# Create engine with separate model and tokenizer
|
||||||
|
engine = InferenceEngine(
|
||||||
|
model=model,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Build request with messages format
|
||||||
|
request = GenerationRequest(
|
||||||
|
messages=[
|
||||||
|
{"role": "system", "content": "You are a helpful assistant."},
|
||||||
|
{"role": "user", "content": "Hello"},
|
||||||
|
],
|
||||||
|
temperature=0.8,
|
||||||
|
top_p=0.95,
|
||||||
|
top_k=50,
|
||||||
|
max_tokens=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Generate (streaming)
|
||||||
|
for token in engine.generate_with_request(request):
|
||||||
|
print(token, end="", flush=True)
|
||||||
|
|
||||||
|
# Or use simple generate interface
|
||||||
|
result = engine.generate(
|
||||||
|
prompt="Hello",
|
||||||
|
stream=False,
|
||||||
|
max_tokens=1024,
|
||||||
|
temperature=0.8,
|
||||||
|
top_p=0.95,
|
||||||
|
top_k=50,
|
||||||
|
)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Generation Modes
|
||||||
|
|
||||||
|
| Mode | Description |
|
||||||
|
|------|-------------|
|
||||||
|
| `stream=True` | Streaming output, yields token by token |
|
||||||
|
| `stream=False` | Non-streaming output, returns complete result |
|
||||||
|
|
||||||
|
> Document Update Time: 2026-05-14
|
||||||
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.5"
|
||||||
|
__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",
|
||||||
|
]
|
||||||
@@ -1,5 +1,4 @@
|
|||||||
import json
|
import json
|
||||||
|
|
||||||
from dataclasses import asdict, dataclass
|
from dataclasses import asdict, dataclass
|
||||||
from typing import Optional, Self
|
from typing import Optional, Self
|
||||||
|
|
||||||
@@ -7,6 +6,7 @@ from typing import Optional, Self
|
|||||||
@dataclass
|
@dataclass
|
||||||
class ModelConfig:
|
class ModelConfig:
|
||||||
# basic config
|
# basic config
|
||||||
|
model_type: Optional[str] = None
|
||||||
vocab_size: Optional[int] = None
|
vocab_size: Optional[int] = None
|
||||||
dim: Optional[int] = None
|
dim: Optional[int] = None
|
||||||
|
|
||||||
@@ -25,10 +25,9 @@ class ModelConfig:
|
|||||||
use_qk_norm: Optional[bool] = None
|
use_qk_norm: Optional[bool] = None
|
||||||
use_gated_attention: Optional[bool] = None
|
use_gated_attention: Optional[bool] = None
|
||||||
|
|
||||||
|
|
||||||
def load(self, config_path: str) -> Self:
|
def load(self, config_path: str) -> Self:
|
||||||
config = {}
|
config = {}
|
||||||
with open(config_path, 'r') as f:
|
with open(config_path, "r") as f:
|
||||||
config.update(json.load(f))
|
config.update(json.load(f))
|
||||||
|
|
||||||
for key, value in config.items():
|
for key, value in config.items():
|
||||||
@@ -39,5 +38,5 @@ class ModelConfig:
|
|||||||
|
|
||||||
def save(self, config_path: str):
|
def save(self, config_path: str):
|
||||||
config_dict = {k: v for k, v in asdict(self).items() if v is not None}
|
config_dict = {k: v for k, v in asdict(self).items() if v is not None}
|
||||||
with open(config_path, 'w') as f:
|
with open(config_path, "w") as f:
|
||||||
json.dump(config_dict, f, indent=4)
|
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,37 @@
|
|||||||
|
from astrai.dataset.dataset import (
|
||||||
|
BaseDataset,
|
||||||
|
DatasetFactory,
|
||||||
|
)
|
||||||
|
from astrai.dataset.sampler import ResumableDistributedSampler
|
||||||
|
from astrai.dataset.storage import (
|
||||||
|
BaseSegmentFetcher,
|
||||||
|
BaseStorage,
|
||||||
|
H5Storage,
|
||||||
|
JSONStorage,
|
||||||
|
MultiSegmentFetcher,
|
||||||
|
available_storage_types,
|
||||||
|
create_storage,
|
||||||
|
detect_format,
|
||||||
|
load_h5,
|
||||||
|
load_json,
|
||||||
|
save_h5,
|
||||||
|
save_json,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"BaseDataset",
|
||||||
|
"DatasetFactory",
|
||||||
|
"BaseSegmentFetcher",
|
||||||
|
"MultiSegmentFetcher",
|
||||||
|
"BaseStorage",
|
||||||
|
"H5Storage",
|
||||||
|
"JSONStorage",
|
||||||
|
"create_storage",
|
||||||
|
"detect_format",
|
||||||
|
"available_storage_types",
|
||||||
|
"save_h5",
|
||||||
|
"load_h5",
|
||||||
|
"save_json",
|
||||||
|
"load_json",
|
||||||
|
"ResumableDistributedSampler",
|
||||||
|
]
|
||||||
@@ -0,0 +1,278 @@
|
|||||||
|
"""Dataset implementations with factory pattern for training."""
|
||||||
|
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from typing import Dict, List, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
|
from astrai.dataset.storage import (
|
||||||
|
BaseStorage,
|
||||||
|
create_storage,
|
||||||
|
detect_format,
|
||||||
|
)
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
|
||||||
|
|
||||||
|
class BaseDataset(Dataset, ABC):
|
||||||
|
"""Abstract base class for all dataset types.
|
||||||
|
|
||||||
|
Implements common functionality for window-based data fetching.
|
||||||
|
Uses a storage abstraction for format-agnostic data loading.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, window_size: int, stride: int):
|
||||||
|
super().__init__()
|
||||||
|
self.window_size = window_size
|
||||||
|
self.stride = stride
|
||||||
|
self.storage: Optional[BaseStorage] = None
|
||||||
|
|
||||||
|
def load(self, load_path: str, storage_type: Optional[str] = None, tokenizer=None):
|
||||||
|
"""Load dataset from the given path.
|
||||||
|
|
||||||
|
Auto-detects the storage format if not specified.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
load_path: Path to the data directory or file
|
||||||
|
storage_type: Force a specific storage type ("h5", "json"),
|
||||||
|
or None for auto-detection
|
||||||
|
tokenizer: Callable str -> List[int], used to tokenize raw text
|
||||||
|
in JSON files. Ignored for HDF5.
|
||||||
|
"""
|
||||||
|
if storage_type is None:
|
||||||
|
storage_type = detect_format(load_path)
|
||||||
|
self.storage = create_storage(storage_type)
|
||||||
|
self.storage.load(load_path, tokenizer=tokenizer)
|
||||||
|
|
||||||
|
def load_json(self, load_path: str, tokenizer=None):
|
||||||
|
"""Load dataset from JSON files explicitly.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
load_path: Path to the JSON data file or directory
|
||||||
|
tokenizer: Optional tokenizer callable for raw text JSON.
|
||||||
|
"""
|
||||||
|
self.load(load_path, storage_type="json", tokenizer=tokenizer)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def count(self) -> int:
|
||||||
|
"""Return the total number of raw elements (tokens) in the dataset."""
|
||||||
|
if self.storage is None:
|
||||||
|
return 0
|
||||||
|
return len(self.storage)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def keys(self) -> List[str]:
|
||||||
|
"""Return the available data keys."""
|
||||||
|
if self.storage is None:
|
||||||
|
return []
|
||||||
|
return self.storage.keys
|
||||||
|
|
||||||
|
def get_index(self, index: int) -> tuple:
|
||||||
|
"""Calculate begin and end indices for a sample.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
index: Sample index
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (begin_idx, end_idx)
|
||||||
|
"""
|
||||||
|
if self.storage is None:
|
||||||
|
raise RuntimeError("Dataset not loaded, call load() first")
|
||||||
|
total = len(self.storage)
|
||||||
|
if total <= self.window_size:
|
||||||
|
raise ValueError(
|
||||||
|
f"Data too short: {total} tokens <= window_size {self.window_size}"
|
||||||
|
)
|
||||||
|
|
||||||
|
begin_idx = min(index * self.stride, total - 1 - self.window_size)
|
||||||
|
end_idx = min(begin_idx + self.window_size, total - 1)
|
||||||
|
|
||||||
|
return begin_idx, end_idx
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
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:
|
||||||
|
if self.storage is None:
|
||||||
|
return 0
|
||||||
|
total = len(self.storage)
|
||||||
|
if total <= self.window_size:
|
||||||
|
return 0
|
||||||
|
return (total - 1 - self.window_size) // self.stride + 1
|
||||||
|
|
||||||
|
|
||||||
|
class DatasetFactory(BaseFactory["BaseDataset"]):
|
||||||
|
"""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,
|
||||||
|
storage_type: Optional[str] = None,
|
||||||
|
tokenizer=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)
|
||||||
|
storage_type: Storage type ("h5", "json") or None for auto-detection
|
||||||
|
tokenizer: Callable str -> List[int] for raw text JSON tokenization
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Loaded dataset instance
|
||||||
|
"""
|
||||||
|
if stride is None:
|
||||||
|
stride = window_size
|
||||||
|
|
||||||
|
dataset = cls.create(train_type, window_size, stride)
|
||||||
|
dataset.load(load_path, storage_type=storage_type, tokenizer=tokenizer)
|
||||||
|
|
||||||
|
return dataset
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def available_types(cls) -> list:
|
||||||
|
"""Return list of registered dataset type names."""
|
||||||
|
return cls.list_registered()
|
||||||
|
|
||||||
|
|
||||||
|
@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.storage.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.storage.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.storage.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.storage.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,
|
||||||
|
}
|
||||||
@@ -1,20 +1,20 @@
|
|||||||
|
from typing import Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
|
|
||||||
from torch.utils.data import Dataset, Sampler
|
from torch.utils.data import Dataset, Sampler
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
|
|
||||||
class ResumableDistributedSampler(Sampler[int]):
|
class ResumableDistributedSampler(Sampler[int]):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
data_source: Dataset,
|
data_source: Dataset,
|
||||||
start_epoch: int=0,
|
start_epoch: int = 0,
|
||||||
start_iter: int=0,
|
start_iter: int = 0,
|
||||||
seed: int=42,
|
seed: int = 42,
|
||||||
drop_last: bool=False,
|
drop_last: bool = False,
|
||||||
shuffle: bool=True,
|
shuffle: bool = True,
|
||||||
process_group: Optional[dist.ProcessGroup]=None,
|
process_group: Optional[dist.ProcessGroup] = None,
|
||||||
):
|
):
|
||||||
self.epoch = start_epoch
|
self.epoch = start_epoch
|
||||||
self.iter = start_iter
|
self.iter = start_iter
|
||||||
@@ -58,10 +58,10 @@ class ResumableDistributedSampler(Sampler[int]):
|
|||||||
padding_size = self.total_size - len(indices)
|
padding_size = self.total_size - len(indices)
|
||||||
indices += indices[:padding_size]
|
indices += indices[:padding_size]
|
||||||
|
|
||||||
local_indices = indices[self.rank:self.total_size:self.num_replicas]
|
local_indices = indices[self.rank : self.total_size : self.num_replicas]
|
||||||
|
|
||||||
self.iter = self.iter % self.num_samples_per_replica
|
self.iter = self.iter % self.num_samples_per_replica
|
||||||
self._indices = local_indices[self.iter:]
|
self._indices = local_indices[self.iter :]
|
||||||
|
|
||||||
def __iter__(self):
|
def __iter__(self):
|
||||||
if self._indices is None:
|
if self._indices is None:
|
||||||
@@ -0,0 +1,312 @@
|
|||||||
|
"""Storage backends for different data formats.
|
||||||
|
|
||||||
|
Each storage handles format-specific loading (HDF5, JSON, etc.) and provides
|
||||||
|
a uniform interface for data access and length observation via fetchers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import bisect
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Callable, Dict, List, Optional, Union
|
||||||
|
|
||||||
|
import h5py
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
def save_json(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
|
||||||
|
os.makedirs(file_path, exist_ok=True)
|
||||||
|
full_file_path = os.path.join(file_path, f"{file_name}.json")
|
||||||
|
json_data = {}
|
||||||
|
for key, tensors in tensor_group.items():
|
||||||
|
json_data[key] = [tensor.tolist() for tensor in tensors]
|
||||||
|
with open(full_file_path, "w", encoding="utf-8") as f:
|
||||||
|
json.dump(json_data, f, ensure_ascii=False)
|
||||||
|
|
||||||
|
|
||||||
|
def load_json(
|
||||||
|
file_path: str,
|
||||||
|
share_memory: bool = True,
|
||||||
|
tokenizer: Optional[Callable[[str], List[int]]] = None,
|
||||||
|
) -> Dict[str, List[Tensor]]:
|
||||||
|
"""Load tensor data from JSON files.
|
||||||
|
|
||||||
|
Supports two modes:
|
||||||
|
- Pre-tokenized: JSON values are List[List[int]] (token IDs), loaded as-is.
|
||||||
|
- Raw text: JSON values are List[str], tokenized via ``tokenizer`` callable
|
||||||
|
at load time. A ``tokenizer`` receives a str and returns List[int].
|
||||||
|
|
||||||
|
Non-data JSON files (e.g. config.json) with scalar/object values are
|
||||||
|
silently skipped.
|
||||||
|
"""
|
||||||
|
tensor_group: Dict[str, List[Tensor]] = {}
|
||||||
|
root_path = Path(file_path)
|
||||||
|
json_files = list(root_path.rglob("*.json")) + list(root_path.rglob("*.jsonl"))
|
||||||
|
for json_file in json_files:
|
||||||
|
with open(json_file, "r", encoding="utf-8") as f:
|
||||||
|
data = json.load(f)
|
||||||
|
if not isinstance(data, dict):
|
||||||
|
continue
|
||||||
|
for key, sequences in data.items():
|
||||||
|
if not isinstance(sequences, list):
|
||||||
|
continue
|
||||||
|
tensors = []
|
||||||
|
for seq in sequences:
|
||||||
|
if tokenizer is not None and isinstance(seq, str):
|
||||||
|
seq = tokenizer(seq)
|
||||||
|
tensor = torch.tensor(seq, dtype=torch.long)
|
||||||
|
if share_memory:
|
||||||
|
tensor = tensor.share_memory_()
|
||||||
|
tensors.append(tensor)
|
||||||
|
if tensor_group.get(key) is None:
|
||||||
|
tensor_group[key] = []
|
||||||
|
tensor_group[key].extend(tensors)
|
||||||
|
return tensor_group
|
||||||
|
|
||||||
|
|
||||||
|
def detect_format(load_path: str) -> str:
|
||||||
|
"""Auto-detect storage format from files in the directory.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
load_path: Directory or file path
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Format string ("h5" or "json")
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
FileNotFoundError: If no supported data files are found
|
||||||
|
"""
|
||||||
|
root = Path(load_path)
|
||||||
|
if root.is_file():
|
||||||
|
suffix = root.suffix.lower()
|
||||||
|
if suffix in (".h5", ".hdf5"):
|
||||||
|
return "h5"
|
||||||
|
if suffix in (".json", ".jsonl"):
|
||||||
|
return "json"
|
||||||
|
raise ValueError(f"Unsupported file format: {suffix}")
|
||||||
|
|
||||||
|
h5_files = list(root.rglob("*.h5")) + list(root.rglob("*.hdf5"))
|
||||||
|
if h5_files:
|
||||||
|
return "h5"
|
||||||
|
json_files = list(root.rglob("*.json")) + list(root.rglob("*.jsonl"))
|
||||||
|
if json_files:
|
||||||
|
return "json"
|
||||||
|
raise FileNotFoundError(f"No supported data files found at {load_path}")
|
||||||
|
|
||||||
|
|
||||||
|
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)."""
|
||||||
|
if not (
|
||||||
|
0 <= begin_idx < self.total_length and 0 <= end_idx <= self.total_length
|
||||||
|
):
|
||||||
|
raise ValueError("begin_idx or end_idx out of bounds")
|
||||||
|
if begin_idx >= end_idx:
|
||||||
|
return torch.tensor([], dtype=torch.long)
|
||||||
|
|
||||||
|
seg_start_idx = bisect.bisect_right(self.cum_lengths, begin_idx)
|
||||||
|
seg_end_idx = bisect.bisect_left(self.cum_lengths, end_idx)
|
||||||
|
|
||||||
|
result_segments = []
|
||||||
|
|
||||||
|
for i in range(seg_start_idx, seg_end_idx + 1):
|
||||||
|
prev_cum = self.cum_lengths[i - 1] if i > 0 else 0
|
||||||
|
start = max(begin_idx - prev_cum, 0)
|
||||||
|
end = min(end_idx - prev_cum, len(self.segments[i]))
|
||||||
|
result_segments.append(self.segments[i][start:end])
|
||||||
|
|
||||||
|
return torch.cat(result_segments, dim=0)
|
||||||
|
|
||||||
|
|
||||||
|
class MultiSegmentFetcher:
|
||||||
|
"""Manages multiple segment fetchers for different data keys."""
|
||||||
|
|
||||||
|
def __init__(self, multi_segments: Dict):
|
||||||
|
self.multi_keys = list(multi_segments.keys())
|
||||||
|
self.multi_fetchers = {
|
||||||
|
key: BaseSegmentFetcher(segments)
|
||||||
|
for key, segments in multi_segments.items()
|
||||||
|
}
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
"""Returns the minimum length across all fetchers."""
|
||||||
|
if not self.multi_fetchers:
|
||||||
|
return 0
|
||||||
|
len_list = [len(seg) for seg in self.multi_fetchers.values()]
|
||||||
|
return min(len_list)
|
||||||
|
|
||||||
|
def key_fetch(
|
||||||
|
self, begin_idx: int, end_idx: int, keys: Union[str, List[str]]
|
||||||
|
) -> Dict:
|
||||||
|
"""Fetch data for specific keys."""
|
||||||
|
fetch_dict = {}
|
||||||
|
keys = [keys] if isinstance(keys, str) else keys
|
||||||
|
|
||||||
|
for key in keys:
|
||||||
|
fetcher = self.multi_fetchers[key]
|
||||||
|
fetch_tensor = fetcher.fetch_data(begin_idx, end_idx)
|
||||||
|
fetch_dict[key] = fetch_tensor
|
||||||
|
|
||||||
|
return fetch_dict if len(keys) > 1 else fetch_dict[keys[0]]
|
||||||
|
|
||||||
|
def fetch_data(self, begin_idx: int, end_idx: int) -> Dict:
|
||||||
|
"""Fetch all keys."""
|
||||||
|
return self.key_fetch(begin_idx, end_idx, self.multi_keys)
|
||||||
|
|
||||||
|
|
||||||
|
class BaseStorage(ABC):
|
||||||
|
"""Abstract storage backend for loading and dispatching data.
|
||||||
|
|
||||||
|
Storage encapsulates format-specific loading and provides a uniform
|
||||||
|
interface for data access and length observation. Subclasses handle
|
||||||
|
different data formats (HDF5, JSON, etc.) while exposing the same
|
||||||
|
fetch interface.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._fetcher: Optional[MultiSegmentFetcher] = None
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def load(self, load_path: str, tokenizer=None) -> None:
|
||||||
|
"""Load data from the given path into internal fetcher."""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
"""Total number of raw elements (tokens) in storage."""
|
||||||
|
if self._fetcher is None:
|
||||||
|
return 0
|
||||||
|
return len(self._fetcher)
|
||||||
|
|
||||||
|
def fetch(self, begin_idx: int, end_idx: int, keys: Union[str, List[str]]):
|
||||||
|
"""Fetch data for the given keys and index range.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
begin_idx: Starting index (inclusive)
|
||||||
|
end_idx: Ending index (exclusive)
|
||||||
|
keys: Single key or list of keys to fetch
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tensor if single key, Dict[str, Tensor] if multiple keys
|
||||||
|
"""
|
||||||
|
if self._fetcher is None:
|
||||||
|
raise RuntimeError("Storage not loaded")
|
||||||
|
return self._fetcher.key_fetch(begin_idx, end_idx, keys)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def keys(self) -> List[str]:
|
||||||
|
"""Return the data keys available in this storage."""
|
||||||
|
if self._fetcher is None:
|
||||||
|
return []
|
||||||
|
return self._fetcher.multi_keys
|
||||||
|
|
||||||
|
|
||||||
|
class H5Storage(BaseStorage):
|
||||||
|
"""HDF5-based storage backend (pre-tokenized data)."""
|
||||||
|
|
||||||
|
def load(self, load_path: str, tokenizer=None) -> None:
|
||||||
|
segments = load_h5(load_path)
|
||||||
|
self._fetcher = MultiSegmentFetcher(segments)
|
||||||
|
|
||||||
|
|
||||||
|
class JSONStorage(BaseStorage):
|
||||||
|
"""JSON-based storage backend.
|
||||||
|
|
||||||
|
Supports two modes:
|
||||||
|
- Pre-tokenized: JSON values are List[List[int]], loaded as-is.
|
||||||
|
- Raw text: JSON values are List[str], tokenized via ``tokenizer``
|
||||||
|
callable (str -> List[int]) at load time.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def load(self, load_path: str, tokenizer=None) -> None:
|
||||||
|
segments = load_json(load_path, tokenizer=tokenizer)
|
||||||
|
self._fetcher = MultiSegmentFetcher(segments)
|
||||||
|
|
||||||
|
|
||||||
|
_STORAGE_REGISTRY: Dict[str, type] = {
|
||||||
|
"h5": H5Storage,
|
||||||
|
"json": JSONStorage,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def create_storage(storage_type: str) -> BaseStorage:
|
||||||
|
"""Create a storage instance by type name.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
storage_type: Storage type name ("h5", "json")
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Storage instance
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If the storage type is unknown
|
||||||
|
"""
|
||||||
|
storage_cls = _STORAGE_REGISTRY.get(storage_type)
|
||||||
|
if storage_cls is None:
|
||||||
|
raise ValueError(
|
||||||
|
f"Unknown storage type: '{storage_type}'. "
|
||||||
|
f"Available: {sorted(_STORAGE_REGISTRY.keys())}"
|
||||||
|
)
|
||||||
|
return storage_cls()
|
||||||
|
|
||||||
|
|
||||||
|
def available_storage_types() -> List[str]:
|
||||||
|
"""Return list of registered storage type names."""
|
||||||
|
return sorted(_STORAGE_REGISTRY.keys())
|
||||||
@@ -0,0 +1,210 @@
|
|||||||
|
"""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 get_component_class(cls, name: str) -> Type[T]:
|
||||||
|
"""Get the registered component class by name without instantiating it.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
name: Registered name of the component
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The component class itself
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If the component name is not registered
|
||||||
|
"""
|
||||||
|
if not cls._registry.contains(name):
|
||||||
|
raise ValueError(
|
||||||
|
f"Unknown component: '{name}'. "
|
||||||
|
f"Supported types: {sorted(cls._registry.list_names())}"
|
||||||
|
)
|
||||||
|
return cls._registry.get(name)
|
||||||
|
|
||||||
|
@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,92 @@
|
|||||||
|
"""Inference module for continuous batching.
|
||||||
|
|
||||||
|
Layers:
|
||||||
|
- core/: Core inference loop (cache, executor, scheduler, task)
|
||||||
|
- api/: HTTP protocol handlers (OpenAI, Anthropic)
|
||||||
|
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
|
||||||
|
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy)
|
||||||
|
"""
|
||||||
|
|
||||||
|
from astrai.inference.api import (
|
||||||
|
AnthropicHandler,
|
||||||
|
AnthropicMessage,
|
||||||
|
ChatCompletionRequest,
|
||||||
|
ChatMessage,
|
||||||
|
MessagesRequest,
|
||||||
|
OpenAIHandler,
|
||||||
|
ProtocolHandler,
|
||||||
|
StopChecker,
|
||||||
|
StreamContext,
|
||||||
|
app,
|
||||||
|
run_server,
|
||||||
|
)
|
||||||
|
from astrai.inference.core import (
|
||||||
|
STOP,
|
||||||
|
Allocator,
|
||||||
|
Executor,
|
||||||
|
InferenceScheduler,
|
||||||
|
KVCache,
|
||||||
|
KvcacheView,
|
||||||
|
PagePool,
|
||||||
|
PrefixCache,
|
||||||
|
Storage,
|
||||||
|
Task,
|
||||||
|
TaskManager,
|
||||||
|
TaskStatus,
|
||||||
|
TaskTable,
|
||||||
|
page_hash,
|
||||||
|
)
|
||||||
|
from astrai.inference.engine import (
|
||||||
|
GenerationRequest,
|
||||||
|
InferenceEngine,
|
||||||
|
)
|
||||||
|
from astrai.inference.sample import (
|
||||||
|
BaseSamplingStrategy,
|
||||||
|
SamplingPipeline,
|
||||||
|
TemperatureStrategy,
|
||||||
|
TopKStrategy,
|
||||||
|
TopPStrategy,
|
||||||
|
sample,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
# Engine / Requests
|
||||||
|
"InferenceEngine",
|
||||||
|
"GenerationRequest",
|
||||||
|
# Core scheduler
|
||||||
|
"InferenceScheduler",
|
||||||
|
"Executor",
|
||||||
|
"STOP",
|
||||||
|
"Task",
|
||||||
|
"TaskManager",
|
||||||
|
"TaskStatus",
|
||||||
|
# Core cache
|
||||||
|
"Allocator",
|
||||||
|
"KVCache",
|
||||||
|
"KvcacheView",
|
||||||
|
"PagePool",
|
||||||
|
"PrefixCache",
|
||||||
|
"Storage",
|
||||||
|
"TaskTable",
|
||||||
|
"page_hash",
|
||||||
|
# Sampling (Strategy pattern)
|
||||||
|
"sample",
|
||||||
|
"BaseSamplingStrategy",
|
||||||
|
"TemperatureStrategy",
|
||||||
|
"TopKStrategy",
|
||||||
|
"TopPStrategy",
|
||||||
|
"SamplingPipeline",
|
||||||
|
# Protocol
|
||||||
|
"ProtocolHandler",
|
||||||
|
"StopChecker",
|
||||||
|
"StreamContext",
|
||||||
|
"AnthropicHandler",
|
||||||
|
"OpenAIHandler",
|
||||||
|
# Server
|
||||||
|
"ChatMessage",
|
||||||
|
"ChatCompletionRequest",
|
||||||
|
"AnthropicMessage",
|
||||||
|
"MessagesRequest",
|
||||||
|
"app",
|
||||||
|
"run_server",
|
||||||
|
]
|
||||||
@@ -0,0 +1,31 @@
|
|||||||
|
"""Inference API: protocol handlers and FastAPI server."""
|
||||||
|
|
||||||
|
from astrai.inference.api.protocol import (
|
||||||
|
AnthropicHandler,
|
||||||
|
OpenAIHandler,
|
||||||
|
ProtocolHandler,
|
||||||
|
StopChecker,
|
||||||
|
StreamContext,
|
||||||
|
)
|
||||||
|
from astrai.inference.api.server import (
|
||||||
|
AnthropicMessage,
|
||||||
|
ChatCompletionRequest,
|
||||||
|
ChatMessage,
|
||||||
|
MessagesRequest,
|
||||||
|
app,
|
||||||
|
run_server,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"AnthropicHandler",
|
||||||
|
"OpenAIHandler",
|
||||||
|
"ProtocolHandler",
|
||||||
|
"StopChecker",
|
||||||
|
"StreamContext",
|
||||||
|
"AnthropicMessage",
|
||||||
|
"ChatCompletionRequest",
|
||||||
|
"ChatMessage",
|
||||||
|
"MessagesRequest",
|
||||||
|
"app",
|
||||||
|
"run_server",
|
||||||
|
]
|
||||||
@@ -0,0 +1,434 @@
|
|||||||
|
"""Protocol handlers for OpenAI and Anthropic chat completion APIs.
|
||||||
|
|
||||||
|
Template Method + Builder patterns eliminate the 45% code duplication between
|
||||||
|
stream/non-stream branches and across protocol adapters.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
|
from fastapi.responses import StreamingResponse
|
||||||
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
from astrai.inference.engine import InferenceEngine
|
||||||
|
|
||||||
|
|
||||||
|
def _sse_event(data: Dict[str, Any], event: Optional[str] = None) -> str:
|
||||||
|
lines: List[str] = []
|
||||||
|
if event:
|
||||||
|
lines.append(f"event: {event}")
|
||||||
|
lines.append(f"data: {json.dumps(data, ensure_ascii=False)}")
|
||||||
|
lines.append("")
|
||||||
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
|
||||||
|
def _sse_done() -> str:
|
||||||
|
return "data: [DONE]\n\n"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class StreamContext:
|
||||||
|
"""Shared state across the streaming generation lifecycle."""
|
||||||
|
|
||||||
|
resp_id: str
|
||||||
|
created: int
|
||||||
|
model: str
|
||||||
|
prompt_tokens: int
|
||||||
|
completion_tokens: int = 0
|
||||||
|
accumulated: str = ""
|
||||||
|
stop_matched: Optional[str] = None
|
||||||
|
last_yield_trimmed: str = ""
|
||||||
|
|
||||||
|
|
||||||
|
class StopChecker:
|
||||||
|
"""Scans accumulated text for stop sequence matches."""
|
||||||
|
|
||||||
|
def __init__(self, sequences: List[str]):
|
||||||
|
self._sequences = [s for s in sequences if s]
|
||||||
|
|
||||||
|
def check(self, text: str) -> Optional[str]:
|
||||||
|
for seq in self._sequences:
|
||||||
|
if seq in text:
|
||||||
|
return seq
|
||||||
|
return None
|
||||||
|
|
||||||
|
def trim(self, text: str, matched: str) -> str:
|
||||||
|
idx = text.rfind(matched)
|
||||||
|
return text[:idx] if idx != -1 else text
|
||||||
|
|
||||||
|
@property
|
||||||
|
def has_sequences(self) -> bool:
|
||||||
|
return len(self._sequences) > 0
|
||||||
|
|
||||||
|
|
||||||
|
class ProtocolHandler(ABC):
|
||||||
|
"""Template-method base for API protocol handlers.
|
||||||
|
|
||||||
|
Subclasses implement format hooks; the base class orchestrates the
|
||||||
|
generate-async loop and SSE/JSON response construction.
|
||||||
|
|
||||||
|
Lifecycle::
|
||||||
|
|
||||||
|
handle()
|
||||||
|
├─ build_prompt() # protocol-specific prompt assembly
|
||||||
|
├─ create_response_id() # unique response identifier
|
||||||
|
├─ [stream]
|
||||||
|
│ ├─ format_stream_start()
|
||||||
|
│ ├─ format_stream_token() × N
|
||||||
|
│ │ └─ on_token() hook for stop-sequence interception
|
||||||
|
│ └─ format_stream_end()
|
||||||
|
└─ [non-stream]
|
||||||
|
├─ (accumulate tokens)
|
||||||
|
└─ format_non_stream_response()
|
||||||
|
"""
|
||||||
|
|
||||||
|
request_model: type[BaseModel]
|
||||||
|
|
||||||
|
def __init__(self, request: BaseModel, engine: InferenceEngine):
|
||||||
|
self.request = request
|
||||||
|
self.engine = engine
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def build_prompt(self) -> str:
|
||||||
|
"""Build the full prompt string from the request messages."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def create_response_id(self) -> str:
|
||||||
|
"""Generate a unique response ID following the protocol convention."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def format_stream_start(self, ctx: StreamContext) -> List[str]:
|
||||||
|
"""Yield SSE events that open the stream (role marker, metadata)."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def format_stream_token(self, ctx: StreamContext, token: str) -> str:
|
||||||
|
"""Yield an SSE event for a single generated token."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def format_stream_end(self, ctx: StreamContext) -> List[str]:
|
||||||
|
"""Yield SSE events that close the stream (finish reason, usage stats)."""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def format_non_stream_response(
|
||||||
|
self, ctx: StreamContext, content: str
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""Build the JSON response body for non-streaming mode."""
|
||||||
|
|
||||||
|
def get_stop_sequences(self) -> List[str]:
|
||||||
|
return []
|
||||||
|
|
||||||
|
def create_stop_checker(self) -> StopChecker:
|
||||||
|
return StopChecker(self.get_stop_sequences())
|
||||||
|
|
||||||
|
def on_token(
|
||||||
|
self, ctx: StreamContext, token: str, stop_checker: StopChecker
|
||||||
|
) -> Optional[str]:
|
||||||
|
"""Hook after each token is appended to accumulated.
|
||||||
|
|
||||||
|
Return a matched stop-sequence string to break the loop,
|
||||||
|
or None to continue.
|
||||||
|
|
||||||
|
"""
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def handle(self) -> Union[StreamingResponse, Dict[str, Any]]:
|
||||||
|
ctx = StreamContext(
|
||||||
|
resp_id=self.create_response_id(),
|
||||||
|
created=int(time.time()),
|
||||||
|
model=self.request.model,
|
||||||
|
prompt_tokens=self._count_prompt_tokens(),
|
||||||
|
)
|
||||||
|
|
||||||
|
agen = self.engine.generate_async(
|
||||||
|
prompt=self.build_prompt(),
|
||||||
|
max_tokens=self.request.max_tokens,
|
||||||
|
temperature=self.request.temperature,
|
||||||
|
top_p=self.request.top_p,
|
||||||
|
top_k=self.request.top_k,
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.request.stream:
|
||||||
|
return self._handle_stream(agen, ctx)
|
||||||
|
else:
|
||||||
|
return await self._handle_non_stream(agen, ctx)
|
||||||
|
|
||||||
|
def _count_prompt_tokens(self) -> int:
|
||||||
|
return len(self.engine.tokenizer.encode(self.build_prompt()))
|
||||||
|
|
||||||
|
def _handle_stream(self, agen, ctx: StreamContext) -> StreamingResponse:
|
||||||
|
stop_checker = self.create_stop_checker()
|
||||||
|
|
||||||
|
async def event_stream():
|
||||||
|
for event in self.format_stream_start(ctx):
|
||||||
|
yield event
|
||||||
|
|
||||||
|
async for token in agen:
|
||||||
|
ctx.completion_tokens += 1
|
||||||
|
ctx.accumulated += token
|
||||||
|
|
||||||
|
matched = self.on_token(ctx, token, stop_checker)
|
||||||
|
if matched:
|
||||||
|
break
|
||||||
|
|
||||||
|
yield self.format_stream_token(ctx, token)
|
||||||
|
|
||||||
|
for event in self.format_stream_end(ctx):
|
||||||
|
yield event
|
||||||
|
yield _sse_done()
|
||||||
|
|
||||||
|
return StreamingResponse(
|
||||||
|
event_stream(),
|
||||||
|
media_type="text/event-stream",
|
||||||
|
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _handle_non_stream(self, agen, ctx: StreamContext) -> Dict[str, Any]:
|
||||||
|
stop_checker = self.create_stop_checker()
|
||||||
|
chunks: List[str] = []
|
||||||
|
|
||||||
|
async for token in agen:
|
||||||
|
ctx.completion_tokens += 1
|
||||||
|
ctx.accumulated += token
|
||||||
|
chunks.append(token)
|
||||||
|
|
||||||
|
matched = self.on_token(ctx, token, stop_checker)
|
||||||
|
if matched:
|
||||||
|
break
|
||||||
|
|
||||||
|
content = "".join(chunks)
|
||||||
|
return self.format_non_stream_response(ctx, content)
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_text_content(content: Union[str, List[Dict[str, Any]]]) -> str:
|
||||||
|
"""Extract plain text from an Anthropic content block (string or list)."""
|
||||||
|
if isinstance(content, str):
|
||||||
|
return content
|
||||||
|
if isinstance(content, list):
|
||||||
|
for block in content:
|
||||||
|
if isinstance(block, dict) and block.get("type") == "text":
|
||||||
|
return block.get("text", "")
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
class OpenAIHandler(ProtocolHandler):
|
||||||
|
"""OpenAI-compatible /v1/chat/completions handler."""
|
||||||
|
|
||||||
|
def build_prompt(self) -> str:
|
||||||
|
messages = [
|
||||||
|
{"role": m.role, "content": m.content} for m in self.request.messages
|
||||||
|
]
|
||||||
|
return self.engine.tokenizer.apply_chat_template(messages, tokenize=False)
|
||||||
|
|
||||||
|
def create_response_id(self) -> str:
|
||||||
|
return f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
||||||
|
|
||||||
|
def format_stream_start(self, ctx: StreamContext) -> List[str]:
|
||||||
|
return [
|
||||||
|
_sse_event(
|
||||||
|
{
|
||||||
|
"id": ctx.resp_id,
|
||||||
|
"object": "chat.completion.chunk",
|
||||||
|
"created": ctx.created,
|
||||||
|
"model": ctx.model,
|
||||||
|
"choices": [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"delta": {"role": "assistant"},
|
||||||
|
"finish_reason": None,
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
def format_stream_token(self, ctx: StreamContext, token: str) -> str:
|
||||||
|
return _sse_event(
|
||||||
|
{
|
||||||
|
"id": ctx.resp_id,
|
||||||
|
"object": "chat.completion.chunk",
|
||||||
|
"created": ctx.created,
|
||||||
|
"model": ctx.model,
|
||||||
|
"choices": [
|
||||||
|
{"index": 0, "delta": {"content": token}, "finish_reason": None}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
def format_stream_end(self, ctx: StreamContext) -> List[str]:
|
||||||
|
return [
|
||||||
|
_sse_event(
|
||||||
|
{
|
||||||
|
"id": ctx.resp_id,
|
||||||
|
"object": "chat.completion.chunk",
|
||||||
|
"created": ctx.created,
|
||||||
|
"model": ctx.model,
|
||||||
|
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
|
||||||
|
}
|
||||||
|
),
|
||||||
|
_sse_event(
|
||||||
|
{
|
||||||
|
"prompt_tokens": ctx.prompt_tokens,
|
||||||
|
"completion_tokens": ctx.completion_tokens,
|
||||||
|
"total_tokens": ctx.prompt_tokens + ctx.completion_tokens,
|
||||||
|
}
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
def format_non_stream_response(
|
||||||
|
self, ctx: StreamContext, content: str
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"id": ctx.resp_id,
|
||||||
|
"object": "chat.completion",
|
||||||
|
"created": ctx.created,
|
||||||
|
"model": ctx.model,
|
||||||
|
"choices": [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"message": {"role": "assistant", "content": content},
|
||||||
|
"finish_reason": "stop",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"usage": {
|
||||||
|
"prompt_tokens": ctx.prompt_tokens,
|
||||||
|
"completion_tokens": ctx.completion_tokens,
|
||||||
|
"total_tokens": ctx.prompt_tokens + ctx.completion_tokens,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class AnthropicHandler(ProtocolHandler):
|
||||||
|
"""Anthropic-compatible /v1/messages handler."""
|
||||||
|
|
||||||
|
def __init__(self, *args, **kwargs):
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
self._yielded = ""
|
||||||
|
|
||||||
|
def build_prompt(self) -> str:
|
||||||
|
messages: List[Dict[str, str]] = []
|
||||||
|
system = getattr(self.request, "system", None)
|
||||||
|
if system:
|
||||||
|
messages.append({"role": "system", "content": system})
|
||||||
|
for m in self.request.messages:
|
||||||
|
content = _extract_text_content(m.content)
|
||||||
|
if content:
|
||||||
|
messages.append({"role": m.role, "content": content})
|
||||||
|
return self.engine.tokenizer.apply_chat_template(messages, tokenize=False)
|
||||||
|
|
||||||
|
def create_response_id(self) -> str:
|
||||||
|
return f"msg_{uuid.uuid4().hex[:24]}"
|
||||||
|
|
||||||
|
def get_stop_sequences(self) -> List[str]:
|
||||||
|
return getattr(self.request, "stop_sequences", None) or []
|
||||||
|
|
||||||
|
def on_token(
|
||||||
|
self, ctx: StreamContext, token: str, stop_checker: StopChecker
|
||||||
|
) -> Optional[str]:
|
||||||
|
matched = stop_checker.check(ctx.accumulated)
|
||||||
|
if not matched:
|
||||||
|
return None
|
||||||
|
|
||||||
|
ctx.stop_matched = matched
|
||||||
|
trimmed = ctx.accumulated[: ctx.accumulated.rfind(matched)]
|
||||||
|
unyielded = trimmed[len(self._yielded) :]
|
||||||
|
if unyielded:
|
||||||
|
ctx.last_yield_trimmed = unyielded
|
||||||
|
return matched
|
||||||
|
|
||||||
|
def format_stream_start(self, ctx: StreamContext) -> List[str]:
|
||||||
|
return [
|
||||||
|
_sse_event(
|
||||||
|
{
|
||||||
|
"type": "message_start",
|
||||||
|
"message": {
|
||||||
|
"id": ctx.resp_id,
|
||||||
|
"type": "message",
|
||||||
|
"role": "assistant",
|
||||||
|
"model": ctx.model,
|
||||||
|
"content": [],
|
||||||
|
"usage": {"input_tokens": ctx.prompt_tokens},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
event="message_start",
|
||||||
|
),
|
||||||
|
_sse_event(
|
||||||
|
{
|
||||||
|
"type": "content_block_start",
|
||||||
|
"index": 0,
|
||||||
|
"content_block": {"type": "text", "text": ""},
|
||||||
|
},
|
||||||
|
event="content_block_start",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
def format_stream_token(self, ctx: StreamContext, token: str) -> str:
|
||||||
|
self._yielded += token
|
||||||
|
return _sse_event(
|
||||||
|
{
|
||||||
|
"type": "content_block_delta",
|
||||||
|
"index": 0,
|
||||||
|
"delta": {"type": "text_delta", "text": token},
|
||||||
|
},
|
||||||
|
event="content_block_delta",
|
||||||
|
)
|
||||||
|
|
||||||
|
def format_stream_end(self, ctx: StreamContext) -> List[str]:
|
||||||
|
matched = ctx.stop_matched
|
||||||
|
events: List[str] = []
|
||||||
|
last_yielded = ctx.last_yield_trimmed
|
||||||
|
if last_yielded:
|
||||||
|
events.append(
|
||||||
|
_sse_event(
|
||||||
|
{
|
||||||
|
"type": "content_block_delta",
|
||||||
|
"index": 0,
|
||||||
|
"delta": {"type": "text_delta", "text": last_yielded},
|
||||||
|
},
|
||||||
|
event="content_block_delta",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
events.append(
|
||||||
|
_sse_event(
|
||||||
|
{"type": "content_block_stop", "index": 0},
|
||||||
|
event="content_block_stop",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
events.append(
|
||||||
|
_sse_event(
|
||||||
|
{
|
||||||
|
"type": "message_delta",
|
||||||
|
"delta": {
|
||||||
|
"stop_reason": "stop_sequence" if matched else "end_turn",
|
||||||
|
"stop_sequence": matched,
|
||||||
|
},
|
||||||
|
"usage": {"output_tokens": ctx.completion_tokens},
|
||||||
|
},
|
||||||
|
event="message_delta",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
events.append(_sse_event({"type": "message_stop"}, event="message_stop"))
|
||||||
|
return events
|
||||||
|
|
||||||
|
def format_non_stream_response(
|
||||||
|
self, ctx: StreamContext, content: str
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
matched = ctx.stop_matched
|
||||||
|
if matched:
|
||||||
|
content = content[: content.rfind(matched)]
|
||||||
|
return {
|
||||||
|
"id": ctx.resp_id,
|
||||||
|
"type": "message",
|
||||||
|
"role": "assistant",
|
||||||
|
"model": ctx.model,
|
||||||
|
"content": [{"type": "text", "text": content}],
|
||||||
|
"stop_reason": "stop_sequence" if matched else "end_turn",
|
||||||
|
"stop_sequence": matched,
|
||||||
|
"usage": {
|
||||||
|
"input_tokens": ctx.prompt_tokens,
|
||||||
|
"output_tokens": ctx.completion_tokens,
|
||||||
|
},
|
||||||
|
}
|
||||||
@@ -0,0 +1,166 @@
|
|||||||
|
"""
|
||||||
|
OpenAI / Anthropic-compatible chat completion server backed by continuous-batching inference.
|
||||||
|
|
||||||
|
Protocol-specific formatting is delegated to ``astrai.inference.protocol``.
|
||||||
|
This module owns the FastAPI app, request/response schemas, and dependency wiring.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from contextlib import asynccontextmanager
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import uvicorn
|
||||||
|
from fastapi import FastAPI, HTTPException, Request
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
from astrai.inference.api.protocol import AnthropicHandler, OpenAIHandler
|
||||||
|
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 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 _create_engine(
|
||||||
|
param_path: Optional[Path] = None,
|
||||||
|
device: str = "cuda",
|
||||||
|
dtype: torch.dtype = torch.bfloat16,
|
||||||
|
max_batch_size: int = 16,
|
||||||
|
) -> InferenceEngine:
|
||||||
|
if param_path is None:
|
||||||
|
param_path = _project_root / "params"
|
||||||
|
if not param_path.exists():
|
||||||
|
raise FileNotFoundError(f"Parameter directory not found: {param_path}")
|
||||||
|
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||||
|
model = AutoModel.from_pretrained(param_path)
|
||||||
|
model.to(device=device, dtype=dtype)
|
||||||
|
logger.info(f"Model loaded on {device} with dtype {dtype}")
|
||||||
|
|
||||||
|
engine = InferenceEngine(
|
||||||
|
model=model,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
max_batch_size=max_batch_size,
|
||||||
|
)
|
||||||
|
logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}")
|
||||||
|
return engine
|
||||||
|
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
|
async def lifespan(app: FastAPI):
|
||||||
|
config = app.state.server_config
|
||||||
|
if not config.get("_test", False):
|
||||||
|
try:
|
||||||
|
app.state.engine = _create_engine(**config)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Failed to load model: {e}")
|
||||||
|
raise
|
||||||
|
yield
|
||||||
|
if app.state.engine:
|
||||||
|
app.state.engine.shutdown()
|
||||||
|
logger.info("Inference engine shutdown complete")
|
||||||
|
|
||||||
|
|
||||||
|
app = FastAPI(title="AstrAI Inference Server", version="0.2.0", lifespan=lifespan)
|
||||||
|
|
||||||
|
|
||||||
|
def _get_engine(request: Request) -> InferenceEngine:
|
||||||
|
engine = request.app.state.engine
|
||||||
|
if engine is None:
|
||||||
|
raise HTTPException(status_code=503, detail="Engine not initialized")
|
||||||
|
return engine
|
||||||
|
|
||||||
|
|
||||||
|
@app.get("/health")
|
||||||
|
async def health(request: Request):
|
||||||
|
return {
|
||||||
|
"status": "ok",
|
||||||
|
"model_loaded": request.app.state.engine is not None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@app.get("/stats")
|
||||||
|
async def get_stats(request: Request):
|
||||||
|
return _get_engine(request).get_stats()
|
||||||
|
|
||||||
|
|
||||||
|
@app.post("/v1/chat/completions")
|
||||||
|
async def chat_completion(request: ChatCompletionRequest, req: Request):
|
||||||
|
engine = _get_engine(req)
|
||||||
|
handler = OpenAIHandler(request, engine)
|
||||||
|
return await handler.handle()
|
||||||
|
|
||||||
|
|
||||||
|
@app.post("/v1/messages")
|
||||||
|
async def create_message(request: MessagesRequest, req: Request):
|
||||||
|
engine = _get_engine(req)
|
||||||
|
handler = AnthropicHandler(request, engine)
|
||||||
|
return await handler.handle()
|
||||||
|
|
||||||
|
|
||||||
|
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,
|
||||||
|
):
|
||||||
|
app.state.server_config = {
|
||||||
|
"device": device,
|
||||||
|
"dtype": dtype,
|
||||||
|
"param_path": param_path,
|
||||||
|
"max_batch_size": max_batch_size,
|
||||||
|
}
|
||||||
|
uvicorn.run(
|
||||||
|
app,
|
||||||
|
host=host,
|
||||||
|
port=port,
|
||||||
|
)
|
||||||
@@ -0,0 +1,32 @@
|
|||||||
|
"""Inference core: cache, executor, scheduler, task management."""
|
||||||
|
|
||||||
|
from astrai.inference.core.cache import (
|
||||||
|
Allocator,
|
||||||
|
KVCache,
|
||||||
|
KvcacheView,
|
||||||
|
PagePool,
|
||||||
|
PrefixCache,
|
||||||
|
Storage,
|
||||||
|
TaskTable,
|
||||||
|
page_hash,
|
||||||
|
)
|
||||||
|
from astrai.inference.core.executor import Executor
|
||||||
|
from astrai.inference.core.scheduler import InferenceScheduler
|
||||||
|
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"Allocator",
|
||||||
|
"KVCache",
|
||||||
|
"KvcacheView",
|
||||||
|
"PagePool",
|
||||||
|
"PrefixCache",
|
||||||
|
"Storage",
|
||||||
|
"TaskTable",
|
||||||
|
"page_hash",
|
||||||
|
"Executor",
|
||||||
|
"InferenceScheduler",
|
||||||
|
"STOP",
|
||||||
|
"Task",
|
||||||
|
"TaskManager",
|
||||||
|
"TaskStatus",
|
||||||
|
]
|
||||||
@@ -0,0 +1,372 @@
|
|||||||
|
import threading
|
||||||
|
from collections import OrderedDict
|
||||||
|
from typing import Callable, Dict, List, Optional, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
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 Allocator:
|
||||||
|
"""Bitmask-based page allocator with ref-counting and LRU eviction."""
|
||||||
|
|
||||||
|
def __init__(self, n_pages: int):
|
||||||
|
self._free_mask = (1 << n_pages) - 1
|
||||||
|
self._refs: List[int] = [0] * n_pages
|
||||||
|
self._lru: OrderedDict[int, None] = OrderedDict()
|
||||||
|
self.on_evict: Optional[Callable[[int], None]] = None
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
def alloc(self) -> int:
|
||||||
|
with self._lock:
|
||||||
|
if self._free_mask:
|
||||||
|
lsb = self._free_mask & -self._free_mask
|
||||||
|
idx = lsb.bit_length() - 1
|
||||||
|
self._free_mask ^= lsb
|
||||||
|
self._refs[idx] = 1
|
||||||
|
return idx
|
||||||
|
if self._lru:
|
||||||
|
idx, _ = self._lru.popitem(last=False)
|
||||||
|
if self.on_evict:
|
||||||
|
self.on_evict(idx)
|
||||||
|
self._refs[idx] = 1
|
||||||
|
self._free_mask &= ~(1 << idx)
|
||||||
|
return idx
|
||||||
|
return -1
|
||||||
|
|
||||||
|
def free(self, idx: int, keep_cached: bool = False) -> None:
|
||||||
|
with self._lock:
|
||||||
|
self._refs[idx] -= 1
|
||||||
|
if self._refs[idx] == 0:
|
||||||
|
if keep_cached:
|
||||||
|
self._lru[idx] = None
|
||||||
|
else:
|
||||||
|
self._free_mask |= 1 << idx
|
||||||
|
|
||||||
|
def inc_ref(self, idx: int) -> None:
|
||||||
|
with self._lock:
|
||||||
|
self._refs[idx] += 1
|
||||||
|
self._lru.pop(idx, None)
|
||||||
|
|
||||||
|
def ref_count(self, idx: int) -> int:
|
||||||
|
with self._lock:
|
||||||
|
return self._refs[idx]
|
||||||
|
|
||||||
|
def touch(self, idx: int) -> None:
|
||||||
|
with self._lock:
|
||||||
|
self._lru.move_to_end(idx)
|
||||||
|
|
||||||
|
|
||||||
|
class PrefixCache:
|
||||||
|
"""Hash-based prefix matching: maps page hashes to physical page indices."""
|
||||||
|
|
||||||
|
def __init__(self, page_size: int):
|
||||||
|
self._page_size = page_size
|
||||||
|
self._page_to_hash: Dict[int, int] = {}
|
||||||
|
self._hash_to_page: Dict[int, int] = {}
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
def evict(self, idx: int) -> None:
|
||||||
|
with self._lock:
|
||||||
|
h = self._page_to_hash.pop(idx, None)
|
||||||
|
if h is not None:
|
||||||
|
self._hash_to_page.pop(h, None)
|
||||||
|
|
||||||
|
def has_page(self, idx: int) -> bool:
|
||||||
|
with self._lock:
|
||||||
|
return idx in self._page_to_hash
|
||||||
|
|
||||||
|
def lookup(self, token_ids: List[int]) -> List[int]:
|
||||||
|
with self._lock:
|
||||||
|
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 record(
|
||||||
|
self, page_idx: int, token_ids: List[int], logical_page_idx: int
|
||||||
|
) -> None:
|
||||||
|
with self._lock:
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
class PagePool:
|
||||||
|
"""Orchestrates allocator (page management) and PrefixCache (content addressing)."""
|
||||||
|
|
||||||
|
def __init__(self, allocator: Allocator, prefix: PrefixCache):
|
||||||
|
self._alloc = allocator
|
||||||
|
self._prefix = prefix
|
||||||
|
self._alloc.on_evict = prefix.evict
|
||||||
|
|
||||||
|
@property
|
||||||
|
def allocator(self) -> Allocator:
|
||||||
|
return self._alloc
|
||||||
|
|
||||||
|
@property
|
||||||
|
def prefix(self) -> PrefixCache:
|
||||||
|
return self._prefix
|
||||||
|
|
||||||
|
def alloc(self) -> int:
|
||||||
|
return self._alloc.alloc()
|
||||||
|
|
||||||
|
def free(self, idx: int) -> None:
|
||||||
|
keep = self._prefix.has_page(idx)
|
||||||
|
self._alloc.free(idx, keep_cached=keep)
|
||||||
|
if not keep:
|
||||||
|
self._prefix.evict(idx)
|
||||||
|
|
||||||
|
def inc_ref(self, idx: int) -> None:
|
||||||
|
self._alloc.inc_ref(idx)
|
||||||
|
|
||||||
|
def lookup(self, token_ids: List[int]) -> List[int]:
|
||||||
|
hits = self._prefix.lookup(token_ids)
|
||||||
|
for p in hits:
|
||||||
|
self._alloc.touch(p)
|
||||||
|
return hits
|
||||||
|
|
||||||
|
def record(
|
||||||
|
self, page_idx: int, token_ids: List[int], logical_page_idx: int
|
||||||
|
) -> None:
|
||||||
|
self._prefix.record(page_idx, token_ids, logical_page_idx)
|
||||||
|
|
||||||
|
|
||||||
|
class TaskTable:
|
||||||
|
"""Maps task_ids to page tables and cached token counts."""
|
||||||
|
|
||||||
|
def __init__(self, page_size: int):
|
||||||
|
self._page_size = page_size
|
||||||
|
self._pages: Dict[str, List[int]] = {}
|
||||||
|
self._cached: Dict[str, int] = {}
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
def set(self, task_id: str, page_table: List[int], cached: int) -> None:
|
||||||
|
with self._lock:
|
||||||
|
self._pages[task_id] = page_table
|
||||||
|
self._cached[task_id] = cached
|
||||||
|
|
||||||
|
def get(self, task_id: str) -> List[int]:
|
||||||
|
with self._lock:
|
||||||
|
return self._pages.get(task_id, [])
|
||||||
|
|
||||||
|
def get_cached(self, task_id: str) -> int:
|
||||||
|
with self._lock:
|
||||||
|
return self._cached.get(task_id, 0)
|
||||||
|
|
||||||
|
def pop(self, task_id: str) -> Tuple[List[int], int]:
|
||||||
|
with self._lock:
|
||||||
|
pages = self._pages.pop(task_id, [])
|
||||||
|
cached = self._cached.pop(task_id, 0)
|
||||||
|
return pages, cached
|
||||||
|
|
||||||
|
def get_ref(self, task_id: str) -> List[int]:
|
||||||
|
with self._lock:
|
||||||
|
return self._pages.setdefault(task_id, [])
|
||||||
|
|
||||||
|
def table_tensor(self, task_ids: List[str], device: torch.device) -> Tensor:
|
||||||
|
with self._lock:
|
||||||
|
states = [self._pages.get(tid, []) for tid in task_ids]
|
||||||
|
max_pages = max((len(s) for s in states), default=0)
|
||||||
|
rows = [s + [-1] * (max_pages - len(s)) for s in states]
|
||||||
|
return torch.tensor(rows, dtype=torch.long, device=device)
|
||||||
|
|
||||||
|
|
||||||
|
class Storage:
|
||||||
|
"""KV-cache tensor storage with paged write/gather."""
|
||||||
|
|
||||||
|
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.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,
|
||||||
|
)
|
||||||
|
|
||||||
|
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
|
||||||
|
valid = phys_pages >= 0
|
||||||
|
if not valid.all():
|
||||||
|
if valid.any():
|
||||||
|
valid_pages = phys_pages[valid]
|
||||||
|
self.k_cache[layer_id, valid_pages, offset : offset + chunk] = k[
|
||||||
|
valid, written : written + chunk
|
||||||
|
]
|
||||||
|
self.v_cache[layer_id, valid_pages, offset : offset + chunk] = v[
|
||||||
|
valid, written : written + chunk
|
||||||
|
]
|
||||||
|
written += chunk
|
||||||
|
continue
|
||||||
|
self.k_cache[layer_id, phys_pages, offset : offset + chunk] = k[
|
||||||
|
:, written : written + chunk
|
||||||
|
]
|
||||||
|
self.v_cache[layer_id, phys_pages, offset : offset + chunk] = v[
|
||||||
|
:, written : written + chunk
|
||||||
|
]
|
||||||
|
written += chunk
|
||||||
|
|
||||||
|
def gather(
|
||||||
|
self, layer_id: int, page_table: Tensor, total_len: int
|
||||||
|
) -> Tuple[Tensor, Tensor]:
|
||||||
|
safe = page_table.clamp(min=0)
|
||||||
|
k = self.k_cache[layer_id, safe]
|
||||||
|
v = self.v_cache[layer_id, safe]
|
||||||
|
k = k.flatten(1, 2)
|
||||||
|
v = v.flatten(1, 2)
|
||||||
|
if (page_table < 0).any():
|
||||||
|
invalid = (
|
||||||
|
(page_table < 0)
|
||||||
|
.unsqueeze(-1)
|
||||||
|
.expand(-1, -1, self.page_size)
|
||||||
|
.flatten(1, 2)
|
||||||
|
)
|
||||||
|
invalid = invalid[:, :, None, None].expand_as(k)
|
||||||
|
k = k.masked_fill(invalid, 0.0)
|
||||||
|
v = v.masked_fill(invalid, 0.0)
|
||||||
|
k = k[:, :total_len]
|
||||||
|
v = v[:, :total_len]
|
||||||
|
return k, v
|
||||||
|
|
||||||
|
|
||||||
|
class KvcacheView:
|
||||||
|
"""Bundles Storage + page_table + total_len for attention layers."""
|
||||||
|
|
||||||
|
def __init__(self, storage: Storage, page_table: Tensor, total_len: int = 0):
|
||||||
|
self._storage = storage
|
||||||
|
self._page_table = page_table
|
||||||
|
self._total_len = total_len
|
||||||
|
|
||||||
|
def write(self, layer_id: int, k: Tensor, v: Tensor) -> None:
|
||||||
|
start_pos = self._total_len - k.size(1)
|
||||||
|
self._storage.write(layer_id, self._page_table, start_pos, k, v)
|
||||||
|
|
||||||
|
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
||||||
|
return self._storage.gather(layer_id, self._page_table, self._total_len)
|
||||||
|
|
||||||
|
|
||||||
|
class KVCache:
|
||||||
|
"""Facade: page management + KV-cache I/O for continuous batching."""
|
||||||
|
|
||||||
|
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._pool = PagePool(Allocator(n_pages), PrefixCache(page_size))
|
||||||
|
self._table = TaskTable(page_size)
|
||||||
|
self._storage = Storage(
|
||||||
|
n_layers, n_pages, page_size, n_kv_heads, head_dim, device, dtype
|
||||||
|
)
|
||||||
|
|
||||||
|
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
||||||
|
hits = self._pool.lookup(prompt_ids)
|
||||||
|
cached = len(hits) * self.page_size
|
||||||
|
for p in hits:
|
||||||
|
self._pool.inc_ref(p)
|
||||||
|
|
||||||
|
remaining = len(prompt_ids) - cached
|
||||||
|
n_new = (
|
||||||
|
(remaining + self.page_size - 1) // self.page_size if remaining > 0 else 0
|
||||||
|
)
|
||||||
|
new_pages: List[int] = []
|
||||||
|
if n_new > 0:
|
||||||
|
for _ in range(n_new):
|
||||||
|
p = self._pool.alloc()
|
||||||
|
if p < 0:
|
||||||
|
for hp in hits:
|
||||||
|
self._pool.free(hp)
|
||||||
|
for np in new_pages:
|
||||||
|
self._pool.free(np)
|
||||||
|
return False
|
||||||
|
new_pages.append(p)
|
||||||
|
|
||||||
|
self._table.set(task_id, hits + new_pages, cached)
|
||||||
|
return True
|
||||||
|
|
||||||
|
def task_free(self, task_id: str) -> None:
|
||||||
|
page_table, _ = self._table.pop(task_id)
|
||||||
|
for idx in page_table:
|
||||||
|
self._pool.free(idx)
|
||||||
|
|
||||||
|
def task_extend(self, task_id: str, pos: int) -> bool:
|
||||||
|
page_table = self._table.get(task_id)
|
||||||
|
needed = (pos + 1 + self.page_size - 1) // self.page_size
|
||||||
|
while len(page_table) < needed:
|
||||||
|
p = self._pool.alloc()
|
||||||
|
if p < 0:
|
||||||
|
return False
|
||||||
|
page_table.append(p)
|
||||||
|
return True
|
||||||
|
|
||||||
|
def task_cached(self, task_id: str) -> int:
|
||||||
|
return self._table.get_cached(task_id)
|
||||||
|
|
||||||
|
def task_record_hashes(
|
||||||
|
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
||||||
|
) -> None:
|
||||||
|
page_table = self._table.get(task_id)
|
||||||
|
full_pages = len(prompt_ids) // self.page_size
|
||||||
|
for i in range(start_logical_page, full_pages):
|
||||||
|
self._pool.record(page_table[i], prompt_ids, i)
|
||||||
|
|
||||||
|
def make_table_tensor(self, task_ids: List[str], device: torch.device) -> Tensor:
|
||||||
|
return self._table.table_tensor(task_ids, device)
|
||||||
|
|
||||||
|
def bind(self, page_table: Tensor, total_len: int = 0) -> KvcacheView:
|
||||||
|
return KvcacheView(self._storage, page_table, total_len)
|
||||||
@@ -0,0 +1,96 @@
|
|||||||
|
import logging
|
||||||
|
from typing import List, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.inference.core.cache import KVCache
|
||||||
|
from astrai.inference.core.task import Task
|
||||||
|
from astrai.inference.sample import sample
|
||||||
|
from astrai.model.automodel import AutoModel
|
||||||
|
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class Executor:
|
||||||
|
"""Model forward passes for prefill and decode phases."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: AutoModel,
|
||||||
|
tokenizer: AutoTokenizer,
|
||||||
|
page_cache: KVCache,
|
||||||
|
device: Optional[str] = None,
|
||||||
|
dtype: Optional[torch.dtype] = None,
|
||||||
|
):
|
||||||
|
self.model = model
|
||||||
|
self.tokenizer = tokenizer
|
||||||
|
self.page_cache = page_cache
|
||||||
|
self.device = device or next(model.parameters()).device
|
||||||
|
self.dtype = dtype or next(model.parameters()).dtype
|
||||||
|
|
||||||
|
def execute_prefill(
|
||||||
|
self, tasks: List[Task], prompt_len: int, start_pos: int = 0
|
||||||
|
) -> None:
|
||||||
|
if start_pos >= prompt_len:
|
||||||
|
return
|
||||||
|
|
||||||
|
tasks = sorted(tasks, key=lambda t: t.task_id)
|
||||||
|
batch_sz = len(tasks)
|
||||||
|
|
||||||
|
input_ids = torch.tensor(
|
||||||
|
[t.prompt_ids[start_pos:prompt_len] for t in tasks],
|
||||||
|
dtype=torch.long,
|
||||||
|
device=self.device,
|
||||||
|
)
|
||||||
|
|
||||||
|
task_ids = [t.task_id for t in tasks]
|
||||||
|
page_tables = self.page_cache.make_table_tensor(task_ids, self.device)
|
||||||
|
|
||||||
|
with torch.inference_mode():
|
||||||
|
self.model(
|
||||||
|
input_ids,
|
||||||
|
position_ids=torch.arange(
|
||||||
|
start_pos, prompt_len, dtype=torch.long, device=self.device
|
||||||
|
)
|
||||||
|
.unsqueeze(0)
|
||||||
|
.expand(batch_sz, -1),
|
||||||
|
paged_cache=self.page_cache.bind(page_tables, total_len=prompt_len),
|
||||||
|
)
|
||||||
|
|
||||||
|
def execute_decode(self, tasks: List[Task]) -> List[int]:
|
||||||
|
if not tasks:
|
||||||
|
return []
|
||||||
|
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
position_ids = torch.tensor(
|
||||||
|
[t.next_pos for t in tasks], dtype=torch.long, device=self.device
|
||||||
|
)
|
||||||
|
total_len = position_ids.max().item() + 1
|
||||||
|
|
||||||
|
task_ids = [t.task_id for t in tasks]
|
||||||
|
page_tables = self.page_cache.make_table_tensor(task_ids, self.device)
|
||||||
|
|
||||||
|
temperatures = torch.tensor([t.temperature for t in tasks], device=self.device)
|
||||||
|
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),
|
||||||
|
paged_cache=self.page_cache.bind(page_tables, total_len=total_len),
|
||||||
|
position_ids=position_ids.unsqueeze(1),
|
||||||
|
)
|
||||||
|
logits = outputs["logits"][:, -1, :]
|
||||||
|
|
||||||
|
return sample(
|
||||||
|
logits,
|
||||||
|
temperature=temperatures,
|
||||||
|
top_k=top_ks,
|
||||||
|
top_p=top_ps,
|
||||||
|
).tolist()
|
||||||
@@ -0,0 +1,187 @@
|
|||||||
|
import logging
|
||||||
|
import threading
|
||||||
|
from typing import Any, Dict, List, Optional, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.inference.core.cache import KVCache
|
||||||
|
from astrai.inference.core.executor import Executor
|
||||||
|
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
||||||
|
from astrai.model.automodel import AutoModel
|
||||||
|
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class InferenceScheduler:
|
||||||
|
"""Four-phase continuous batching loop: cleanup -> refill -> prefill -> decode."""
|
||||||
|
|
||||||
|
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.max_seq_len = max_seq_len or config.max_len
|
||||||
|
self.device = device or next(model.parameters()).device
|
||||||
|
self.dtype = dtype or next(model.parameters()).dtype
|
||||||
|
|
||||||
|
n_pages = (
|
||||||
|
max_batch_size * (self.max_seq_len + page_size) + page_size - 1
|
||||||
|
) // page_size
|
||||||
|
|
||||||
|
self._page_cache = KVCache(
|
||||||
|
config.n_layers,
|
||||||
|
n_pages,
|
||||||
|
page_size,
|
||||||
|
config.n_kv_heads,
|
||||||
|
config.dim // config.n_heads,
|
||||||
|
self.device,
|
||||||
|
self.dtype,
|
||||||
|
)
|
||||||
|
|
||||||
|
self._task_mgr = TaskManager(
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
max_batch_size=max_batch_size,
|
||||||
|
max_seq_len=self.max_seq_len,
|
||||||
|
max_prompt_len=max_prompt_len,
|
||||||
|
)
|
||||||
|
|
||||||
|
self._executor = Executor(
|
||||||
|
model=model,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
page_cache=self._page_cache,
|
||||||
|
device=self.device,
|
||||||
|
dtype=self.dtype,
|
||||||
|
)
|
||||||
|
|
||||||
|
self._running = False
|
||||||
|
|
||||||
|
def add_task(self, prompt: str, **kwargs) -> str:
|
||||||
|
return self._task_mgr.add_task(prompt, **kwargs)
|
||||||
|
|
||||||
|
def remove_task(self, task_id: str) -> None:
|
||||||
|
for task in self._task_mgr.remove_task(task_id):
|
||||||
|
self._page_cache.task_free(task.task_id)
|
||||||
|
|
||||||
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
|
return self._task_mgr.get_stats()
|
||||||
|
|
||||||
|
def _run_generation_loop(self) -> None:
|
||||||
|
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||||
|
try:
|
||||||
|
while self._running:
|
||||||
|
finished = self._task_mgr.remove_finished_tasks(stop_ids)
|
||||||
|
for task in finished:
|
||||||
|
self._page_cache.task_free(task.task_id)
|
||||||
|
|
||||||
|
active = self._task_mgr.get_active_tasks()
|
||||||
|
available = self._task_mgr.max_batch_size - len(active)
|
||||||
|
if available > 0:
|
||||||
|
candidates = self._task_mgr.pull_candidates(available)
|
||||||
|
failed = []
|
||||||
|
for task in candidates:
|
||||||
|
if self._page_cache.task_alloc(task.task_id, task.prompt_ids):
|
||||||
|
self._task_mgr.activate(task)
|
||||||
|
else:
|
||||||
|
failed.append(task)
|
||||||
|
if failed:
|
||||||
|
self._task_mgr.return_to_waiting(failed)
|
||||||
|
|
||||||
|
if not self._task_mgr.has_work():
|
||||||
|
self._task_mgr.wait_for_tasks(timeout=1.0)
|
||||||
|
continue
|
||||||
|
|
||||||
|
to_prefill = [
|
||||||
|
t for t in self._task_mgr.get_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),
|
||||||
|
self._page_cache.task_cached(t.task_id),
|
||||||
|
)
|
||||||
|
groups.setdefault(key, []).append(t)
|
||||||
|
|
||||||
|
for (prompt_len, start_pos), group in groups.items():
|
||||||
|
self._executor.execute_prefill(group, prompt_len, start_pos)
|
||||||
|
start_logical_page = start_pos // self._page_cache.page_size
|
||||||
|
for t in group:
|
||||||
|
self._page_cache.task_record_hashes(
|
||||||
|
t.task_id,
|
||||||
|
t.prompt_ids,
|
||||||
|
start_logical_page=start_logical_page,
|
||||||
|
)
|
||||||
|
|
||||||
|
pos_groups: Dict[int, List[Task]] = {}
|
||||||
|
for t in self._task_mgr.get_active_tasks():
|
||||||
|
pos_groups.setdefault(t.next_pos, []).append(t)
|
||||||
|
|
||||||
|
if pos_groups:
|
||||||
|
best_key = max(pos_groups, key=lambda k: len(pos_groups[k]))
|
||||||
|
group = sorted(pos_groups[best_key], key=lambda t: t.task_id)
|
||||||
|
|
||||||
|
valid: List[Task] = []
|
||||||
|
for t in group:
|
||||||
|
if self._page_cache.task_extend(t.task_id, t.next_pos):
|
||||||
|
valid.append(t)
|
||||||
|
else:
|
||||||
|
t.status = TaskStatus.ABORTED
|
||||||
|
if t.stream_callback:
|
||||||
|
t.stream_callback(STOP)
|
||||||
|
|
||||||
|
if valid:
|
||||||
|
next_tokens = self._executor.execute_decode(valid)
|
||||||
|
|
||||||
|
for t, ntok in zip(valid, next_tokens):
|
||||||
|
t.output_ids.append(ntok)
|
||||||
|
t.output_tokens += 1
|
||||||
|
pos = t.input_tokens + t.output_tokens
|
||||||
|
self._page_cache.task_extend(t.task_id, pos)
|
||||||
|
if t.stream_callback:
|
||||||
|
t.stream_callback(
|
||||||
|
self._task_mgr.tokenizer.decode([ntok])
|
||||||
|
)
|
||||||
|
|
||||||
|
for t in valid:
|
||||||
|
if t.is_finished(stop_ids):
|
||||||
|
if t.stream_callback:
|
||||||
|
t.stream_callback(STOP)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
||||||
|
for task in self._task_mgr.get_active_tasks():
|
||||||
|
if task.stream_callback:
|
||||||
|
task.stream_callback(STOP)
|
||||||
|
self._page_cache.task_free(task.task_id)
|
||||||
|
self._task_mgr.clear_queues()
|
||||||
|
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_mgr.wake()
|
||||||
|
if hasattr(self, "_loop_thread"):
|
||||||
|
self._loop_thread.join(timeout=2.0)
|
||||||
|
for task in self._task_mgr.get_active_tasks():
|
||||||
|
self._page_cache.task_free(task.task_id)
|
||||||
|
self._task_mgr.clear_queues()
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.empty_cache()
|
||||||
@@ -0,0 +1,202 @@
|
|||||||
|
import logging
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
from collections import deque
|
||||||
|
from enum import Enum
|
||||||
|
from typing import Any, Callable, Deque, Dict, List, Optional
|
||||||
|
|
||||||
|
from astrai.tokenize.tokenizer import AutoTokenizer
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
STOP = object()
|
||||||
|
|
||||||
|
|
||||||
|
class TaskStatus(Enum):
|
||||||
|
"""Task lifecycle states."""
|
||||||
|
|
||||||
|
PENDING = "pending"
|
||||||
|
RUNNING = "running"
|
||||||
|
FINISHED = "finished"
|
||||||
|
ABORTED = "aborted"
|
||||||
|
|
||||||
|
|
||||||
|
class Task:
|
||||||
|
"""Single generation request: prompt, sampling params, output state."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
prompt_ids: List[int],
|
||||||
|
max_tokens: Optional[int] = None,
|
||||||
|
temperature: float = 1.0,
|
||||||
|
top_p: float = 1.0,
|
||||||
|
top_k: int = 50,
|
||||||
|
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.arrival_time = time.time()
|
||||||
|
self.finish_time: Optional[float] = None
|
||||||
|
self.stream_callback = stream_callback
|
||||||
|
|
||||||
|
@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.max_tokens is not None and self.output_tokens >= self.max_tokens:
|
||||||
|
return True
|
||||||
|
if self.output_ids and self.output_ids[-1] in stop_ids:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
class TaskManager:
|
||||||
|
"""Thread-safe task queues and lifecycle transitions (no page ops)."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
tokenizer: AutoTokenizer,
|
||||||
|
max_batch_size: int = 16,
|
||||||
|
max_seq_len: int = 8192,
|
||||||
|
max_prompt_len: int = 512,
|
||||||
|
):
|
||||||
|
self.tokenizer = tokenizer
|
||||||
|
self.max_batch_size = max_batch_size
|
||||||
|
self.max_seq_len = max_seq_len
|
||||||
|
self.max_prompt_len = max_prompt_len
|
||||||
|
|
||||||
|
self.waiting_queue: Deque[Task] = deque()
|
||||||
|
self.active_tasks: List[Task] = []
|
||||||
|
|
||||||
|
self._task_event = threading.Event()
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
self._total_tasks = 0
|
||||||
|
self._total_tokens = 0
|
||||||
|
|
||||||
|
def add_task(
|
||||||
|
self,
|
||||||
|
prompt: str,
|
||||||
|
max_tokens: Optional[int] = None,
|
||||||
|
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 :]
|
||||||
|
|
||||||
|
if len(prompt_ids) >= self.max_seq_len:
|
||||||
|
if stream_callback:
|
||||||
|
stream_callback(STOP)
|
||||||
|
return task_id
|
||||||
|
|
||||||
|
if max_tokens is None:
|
||||||
|
max_tokens = self.max_seq_len - len(prompt_ids)
|
||||||
|
else:
|
||||||
|
max_tokens = min(max_tokens, self.max_seq_len - len(prompt_ids))
|
||||||
|
|
||||||
|
task = Task(
|
||||||
|
task_id=task_id,
|
||||||
|
prompt_ids=prompt_ids,
|
||||||
|
max_tokens=max_tokens,
|
||||||
|
temperature=temperature,
|
||||||
|
top_p=top_p,
|
||||||
|
top_k=top_k,
|
||||||
|
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) -> List[Task]:
|
||||||
|
with self._lock:
|
||||||
|
removed_active = [t for t in self.active_tasks if t.task_id == task_id]
|
||||||
|
self.waiting_queue = deque(
|
||||||
|
t for t in self.waiting_queue if t.task_id != task_id
|
||||||
|
)
|
||||||
|
self.active_tasks = [t for t in self.active_tasks if t.task_id != task_id]
|
||||||
|
return removed_active
|
||||||
|
|
||||||
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"total_tasks": self._total_tasks,
|
||||||
|
"total_tokens": self._total_tokens,
|
||||||
|
"active_tasks": len(self.active_tasks),
|
||||||
|
"waiting_queue": len(self.waiting_queue),
|
||||||
|
}
|
||||||
|
|
||||||
|
def remove_finished_tasks(self, stop_ids: List[int]) -> List[Task]:
|
||||||
|
with self._lock:
|
||||||
|
finished = []
|
||||||
|
for task in self.active_tasks:
|
||||||
|
if task.status == TaskStatus.ABORTED:
|
||||||
|
task.finish_time = time.time()
|
||||||
|
finished.append(task)
|
||||||
|
elif task.is_finished(stop_ids):
|
||||||
|
task.status = TaskStatus.FINISHED
|
||||||
|
task.finish_time = time.time()
|
||||||
|
finished.append(task)
|
||||||
|
self._total_tokens += task.output_tokens
|
||||||
|
|
||||||
|
self.active_tasks = [
|
||||||
|
t
|
||||||
|
for t in self.active_tasks
|
||||||
|
if t.status not in (TaskStatus.FINISHED, TaskStatus.ABORTED)
|
||||||
|
]
|
||||||
|
return finished
|
||||||
|
|
||||||
|
def pull_candidates(self, n: int) -> List[Task]:
|
||||||
|
to_add: List[Task] = []
|
||||||
|
with self._lock:
|
||||||
|
take = min(n, len(self.waiting_queue))
|
||||||
|
for _ in range(take):
|
||||||
|
to_add.append(self.waiting_queue.popleft())
|
||||||
|
return to_add
|
||||||
|
|
||||||
|
def activate(self, task: Task) -> None:
|
||||||
|
task.status = TaskStatus.RUNNING
|
||||||
|
with self._lock:
|
||||||
|
self.active_tasks.append(task)
|
||||||
|
|
||||||
|
def return_to_waiting(self, tasks: List[Task]) -> None:
|
||||||
|
with self._lock:
|
||||||
|
for task in reversed(tasks):
|
||||||
|
self.waiting_queue.appendleft(task)
|
||||||
|
|
||||||
|
def has_work(self) -> bool:
|
||||||
|
return bool(self.active_tasks or self.waiting_queue)
|
||||||
|
|
||||||
|
def wait_for_tasks(self, timeout: float = 1.0) -> None:
|
||||||
|
self._task_event.clear()
|
||||||
|
self._task_event.wait(timeout=timeout)
|
||||||
|
|
||||||
|
def get_active_tasks(self) -> List[Task]:
|
||||||
|
with self._lock:
|
||||||
|
return list(self.active_tasks)
|
||||||
|
|
||||||
|
def clear_queues(self) -> None:
|
||||||
|
with self._lock:
|
||||||
|
self.waiting_queue.clear()
|
||||||
|
self.active_tasks.clear()
|
||||||
|
|
||||||
|
def wake(self) -> None:
|
||||||
|
self._task_event.set()
|
||||||
@@ -0,0 +1,296 @@
|
|||||||
|
"""Unified inference engine for continuous batching."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import gc
|
||||||
|
import threading
|
||||||
|
from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
from astrai.inference.core.scheduler import InferenceScheduler
|
||||||
|
from astrai.inference.core.task import STOP
|
||||||
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_sampling_params(
|
||||||
|
top_k: int, top_p: float, temperature: float, max_tokens: Optional[int] = None
|
||||||
|
):
|
||||||
|
if not (isinstance(top_k, int) and top_k >= 0):
|
||||||
|
raise ValueError("top_k must be a non-negative integer")
|
||||||
|
if not (0.0 <= top_p <= 1.0):
|
||||||
|
raise ValueError("top_p must be a float between 0.0 and 1.0")
|
||||||
|
if not (isinstance(temperature, (int, float)) and temperature >= 0):
|
||||||
|
raise ValueError("temperature must be a non-negative number")
|
||||||
|
|
||||||
|
|
||||||
|
class GenerateResult:
|
||||||
|
"""Thread-safe token accumulator for streaming and non-streaming modes."""
|
||||||
|
|
||||||
|
def __init__(self, count: int = 1):
|
||||||
|
self._cond = threading.Condition()
|
||||||
|
self._event = threading.Event()
|
||||||
|
self.tokens: List[Tuple[int, str]] = []
|
||||||
|
self.results: List[str] = [""] * count
|
||||||
|
self._done: List[bool] = [False] * count
|
||||||
|
self._completed = 0
|
||||||
|
self._total = count
|
||||||
|
|
||||||
|
def append(self, token: str, idx: int = 0):
|
||||||
|
with self._cond:
|
||||||
|
self.tokens.append((idx, token))
|
||||||
|
if token is not STOP:
|
||||||
|
self.results[idx] += token
|
||||||
|
else:
|
||||||
|
if not self._done[idx]:
|
||||||
|
self._done[idx] = True
|
||||||
|
self._completed += 1
|
||||||
|
self._cond.notify_all()
|
||||||
|
self._event.set()
|
||||||
|
|
||||||
|
def pop_all(self) -> List[Tuple[int, str]]:
|
||||||
|
with self._cond:
|
||||||
|
out = self.tokens.copy()
|
||||||
|
self.tokens.clear()
|
||||||
|
if not out:
|
||||||
|
self._event.clear()
|
||||||
|
return out
|
||||||
|
|
||||||
|
def wait(self, timeout: Optional[float] = None) -> bool:
|
||||||
|
return self._event.wait(timeout=timeout)
|
||||||
|
|
||||||
|
def wait_completion(self, timeout: float = 300.0) -> None:
|
||||||
|
with self._cond:
|
||||||
|
if not self._cond.wait_for(
|
||||||
|
lambda: self._completed >= self._total, timeout=timeout
|
||||||
|
):
|
||||||
|
raise TimeoutError(
|
||||||
|
f"Generation timeout after {timeout}s "
|
||||||
|
f"({self._completed}/{self._total} completed)"
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_results(self) -> List[str]:
|
||||||
|
with self._cond:
|
||||||
|
return self.results.copy()
|
||||||
|
|
||||||
|
|
||||||
|
class GenerationRequest:
|
||||||
|
"""Request parameters for text generation."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
messages: List[Dict[str, str]],
|
||||||
|
top_k: int = 50,
|
||||||
|
top_p: float = 1.0,
|
||||||
|
temperature: float = 1.0,
|
||||||
|
max_tokens: Optional[int] = None,
|
||||||
|
stream: bool = False,
|
||||||
|
):
|
||||||
|
_validate_sampling_params(top_k, top_p, temperature, max_tokens)
|
||||||
|
|
||||||
|
self.messages = messages
|
||||||
|
self.top_k = top_k
|
||||||
|
self.top_p = top_p
|
||||||
|
self.temperature = temperature
|
||||||
|
self.max_tokens = max_tokens
|
||||||
|
self.stream = stream
|
||||||
|
|
||||||
|
|
||||||
|
class InferenceEngine:
|
||||||
|
"""Unified inference engine backed by continuous-batching scheduler."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model: nn.Module,
|
||||||
|
tokenizer: AutoTokenizer,
|
||||||
|
max_batch_size: int = 1,
|
||||||
|
max_seq_len: Optional[int] = None,
|
||||||
|
max_prompt_len: int = 2048,
|
||||||
|
page_size: int = 128,
|
||||||
|
):
|
||||||
|
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: Optional[int] = None,
|
||||||
|
temperature: float = 1.0,
|
||||||
|
top_p: float = 1.0,
|
||||||
|
top_k: int = 50,
|
||||||
|
) -> Union[Generator, str, List[str]]:
|
||||||
|
_validate_sampling_params(top_k, top_p, temperature, max_tokens)
|
||||||
|
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: Optional[int] = None,
|
||||||
|
temperature: float = 1.0,
|
||||||
|
top_p: float = 1.0,
|
||||||
|
top_k: int = 50,
|
||||||
|
) -> AsyncGenerator[str, None]:
|
||||||
|
_validate_sampling_params(top_k, top_p, temperature, max_tokens)
|
||||||
|
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]:
|
||||||
|
try:
|
||||||
|
return next(gen)
|
||||||
|
except StopIteration:
|
||||||
|
return None
|
||||||
|
|
||||||
|
def generate_with_request(
|
||||||
|
self, request: GenerationRequest
|
||||||
|
) -> Union[Generator[str, None, None], str, List[str]]:
|
||||||
|
prompt = self.tokenizer.apply_chat_template(request.messages, tokenize=False)
|
||||||
|
return self.generate(
|
||||||
|
prompt=prompt,
|
||||||
|
stream=request.stream,
|
||||||
|
max_tokens=request.max_tokens,
|
||||||
|
temperature=request.temperature,
|
||||||
|
top_p=request.top_p,
|
||||||
|
top_k=request.top_k,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _submit_tasks(
|
||||||
|
self,
|
||||||
|
prompts: List[str],
|
||||||
|
max_tokens: Optional[int],
|
||||||
|
temperature: float,
|
||||||
|
top_p: float,
|
||||||
|
top_k: int,
|
||||||
|
) -> Tuple[GenerateResult, List[str]]:
|
||||||
|
n = len(prompts)
|
||||||
|
result = GenerateResult(count=n)
|
||||||
|
task_ids = []
|
||||||
|
for i, p in enumerate(prompts):
|
||||||
|
cb = self._make_callback(result, i)
|
||||||
|
task_id = self.scheduler.add_task(
|
||||||
|
prompt=p,
|
||||||
|
max_tokens=max_tokens,
|
||||||
|
temperature=temperature,
|
||||||
|
top_p=top_p,
|
||||||
|
top_k=top_k,
|
||||||
|
stream_callback=cb,
|
||||||
|
)
|
||||||
|
task_ids.append(task_id)
|
||||||
|
return result, task_ids
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _make_callback(result: GenerateResult, idx: int):
|
||||||
|
def cb(token):
|
||||||
|
result.append(token, idx)
|
||||||
|
|
||||||
|
return cb
|
||||||
|
|
||||||
|
def _generate_streaming(
|
||||||
|
self,
|
||||||
|
prompts: List[str],
|
||||||
|
is_batch: bool,
|
||||||
|
max_tokens: Optional[int],
|
||||||
|
temperature: float,
|
||||||
|
top_p: float,
|
||||||
|
top_k: int,
|
||||||
|
) -> Generator:
|
||||||
|
result, task_ids = self._submit_tasks(
|
||||||
|
prompts, max_tokens, temperature, top_p, top_k
|
||||||
|
)
|
||||||
|
n = len(prompts)
|
||||||
|
remaining = n
|
||||||
|
finished = [False] * n
|
||||||
|
|
||||||
|
def gen():
|
||||||
|
nonlocal remaining
|
||||||
|
try:
|
||||||
|
while remaining > 0:
|
||||||
|
items = result.pop_all()
|
||||||
|
for idx, token in items:
|
||||||
|
if token is STOP:
|
||||||
|
if not finished[idx]:
|
||||||
|
finished[idx] = True
|
||||||
|
remaining -= 1
|
||||||
|
else:
|
||||||
|
yield (idx, token) if is_batch else token
|
||||||
|
if remaining > 0:
|
||||||
|
result.wait(timeout=0.05)
|
||||||
|
finally:
|
||||||
|
for tid in task_ids:
|
||||||
|
self.scheduler.remove_task(tid)
|
||||||
|
|
||||||
|
return gen()
|
||||||
|
|
||||||
|
def _generate_non_streaming(
|
||||||
|
self,
|
||||||
|
prompts: List[str],
|
||||||
|
is_batch: bool,
|
||||||
|
max_tokens: Optional[int],
|
||||||
|
temperature: float,
|
||||||
|
top_p: float,
|
||||||
|
top_k: int,
|
||||||
|
) -> Union[str, List[str]]:
|
||||||
|
result, task_ids = self._submit_tasks(
|
||||||
|
prompts, max_tokens, temperature, top_p, top_k
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
result.wait_completion()
|
||||||
|
except TimeoutError:
|
||||||
|
for tid in task_ids:
|
||||||
|
self.scheduler.remove_task(tid)
|
||||||
|
raise
|
||||||
|
|
||||||
|
for tid in task_ids:
|
||||||
|
self.scheduler.remove_task(tid)
|
||||||
|
|
||||||
|
res = result.get_results()
|
||||||
|
return res if is_batch else res[0]
|
||||||
|
|
||||||
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
|
return self.scheduler.get_stats()
|
||||||
|
|
||||||
|
def shutdown(self) -> None:
|
||||||
|
self.scheduler.stop()
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
gc.collect()
|
||||||
@@ -0,0 +1,188 @@
|
|||||||
|
"""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):
|
||||||
|
tk = tk.to(logits.device, non_blocking=True).long().clamp(min=0)
|
||||||
|
max_k = int(tk.max().item())
|
||||||
|
if max_k <= 0:
|
||||||
|
return logits
|
||||||
|
max_k = min(max_k, logits.size(-1))
|
||||||
|
values, _ = torch.topk(logits, max_k, dim=-1)
|
||||||
|
per_row_k = tk.clamp(max=max_k)
|
||||||
|
thresholds = torch.full_like(logits[..., -1:], -float("inf"))
|
||||||
|
positive = per_row_k > 0
|
||||||
|
if positive.any():
|
||||||
|
row_idx = torch.arange(logits.size(0), device=logits.device)[positive]
|
||||||
|
thresholds[positive] = values[
|
||||||
|
row_idx, per_row_k[positive] - 1
|
||||||
|
].unsqueeze(-1)
|
||||||
|
logits[logits < thresholds] = filter_value
|
||||||
|
return logits
|
||||||
|
if tk > 0:
|
||||||
|
k = min(tk, logits.size(-1))
|
||||||
|
thresholds = torch.topk(logits, k, dim=-1)[0][..., -1:]
|
||||||
|
logits[logits < thresholds] = filter_value
|
||||||
|
return logits
|
||||||
|
|
||||||
|
|
||||||
|
class TopPStrategy(BaseSamplingStrategy):
|
||||||
|
"""Nucleus (top-p) filtering: keeps the smallest set of tokens whose
|
||||||
|
cumulative probability exceeds top_p.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
top_p: Scalar or ``[batch]`` tensor (1.0 disables).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, top_p: Union[float, Tensor] = 1.0):
|
||||||
|
self.top_p = top_p
|
||||||
|
|
||||||
|
def _apply(self, logits, 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,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,99 @@
|
|||||||
|
"""
|
||||||
|
AutoModel base class for model loading and saving.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Self, Union
|
||||||
|
|
||||||
|
import safetensors.torch as st
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
from astrai.config import ModelConfig
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
|
||||||
|
|
||||||
|
@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(BaseFactory["AutoModel"], nn.Module):
|
||||||
|
"""
|
||||||
|
Autoregressive language model base class.
|
||||||
|
Provides model loading/saving, registration, and generation.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, config: ModelConfig):
|
||||||
|
super().__init__()
|
||||||
|
self.config = config
|
||||||
|
|
||||||
|
@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 = AutoModel.get_component_class(model_type)
|
||||||
|
|
||||||
|
with _disable_random_init(enable=disable_random_init):
|
||||||
|
model = actual_cls(config)
|
||||||
|
|
||||||
|
# Load weights
|
||||||
|
weights_path = model_path / "model.safetensors"
|
||||||
|
if weights_path.exists():
|
||||||
|
state_dict = st.load_file(str(weights_path))
|
||||||
|
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,330 @@
|
|||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.inference.core.cache import KvcacheView
|
||||||
|
|
||||||
|
|
||||||
|
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
||||||
|
"""Repeat KV heads n_rep times for GQA."""
|
||||||
|
bs, slen, n_heads, head_dim = x.shape
|
||||||
|
if n_rep == 1:
|
||||||
|
return x
|
||||||
|
return (
|
||||||
|
x[:, :, :, None, :]
|
||||||
|
.expand(bs, slen, n_heads, n_rep, head_dim)
|
||||||
|
.reshape(bs, slen, n_heads * n_rep, head_dim)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def get_rotary_emb(
|
||||||
|
dim: int,
|
||||||
|
max_len: int,
|
||||||
|
base: float = 10000,
|
||||||
|
device: Optional[torch.device] = None,
|
||||||
|
) -> Tensor:
|
||||||
|
theta = base ** (-torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim)
|
||||||
|
t = torch.arange(0, max_len, dtype=torch.float64, device=device)
|
||||||
|
freqs = torch.outer(t, theta).float()
|
||||||
|
cos = torch.cos(freqs)
|
||||||
|
sin = torch.sin(freqs)
|
||||||
|
return torch.complex(cos, sin)
|
||||||
|
|
||||||
|
|
||||||
|
def apply_rotary_emb(x: torch.Tensor, freqs_cis: Tensor) -> Tensor:
|
||||||
|
dtype = x.dtype
|
||||||
|
x_ = x.float().reshape(*x.shape[:-1], -1, 2)
|
||||||
|
x_complex = torch.view_as_complex(x_)
|
||||||
|
freqs_cis = freqs_cis.unsqueeze(2)
|
||||||
|
x_rotated = x_complex * freqs_cis
|
||||||
|
x_out = torch.view_as_real(x_rotated).flatten(-2)
|
||||||
|
return x_out.to(dtype)
|
||||||
|
|
||||||
|
|
||||||
|
class RotaryEmbedding(nn.Module):
|
||||||
|
def __init__(self, dim: int, max_len: int, base: int = 10000):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
self.max_len = max_len
|
||||||
|
self.base = base
|
||||||
|
self._set_rotary_buffer(self.max_len)
|
||||||
|
|
||||||
|
def _set_rotary_buffer(self, max_len: int):
|
||||||
|
rotary_emb = get_rotary_emb(self.dim, max_len, self.base)
|
||||||
|
freqs_cis = torch.view_as_real(rotary_emb)
|
||||||
|
self.register_buffer("freqs_cis", freqs_cis, persistent=False)
|
||||||
|
|
||||||
|
def forward(self, x: Tensor, position_ids: Optional[Tensor] = None) -> Tensor:
|
||||||
|
if position_ids is None:
|
||||||
|
position_ids = (
|
||||||
|
torch.arange(x.size(1), device=x.device)
|
||||||
|
.unsqueeze(0)
|
||||||
|
.expand(x.size(0), -1)
|
||||||
|
)
|
||||||
|
position_freq_cis = self.freqs_cis[position_ids].float()
|
||||||
|
return torch.view_as_complex(position_freq_cis)
|
||||||
|
|
||||||
|
|
||||||
|
class Linear(nn.Module):
|
||||||
|
def __init__(self, in_dim: int, out_dim: int, bias: bool = False):
|
||||||
|
super().__init__()
|
||||||
|
self.weight = nn.Parameter(torch.empty((out_dim, in_dim)))
|
||||||
|
self.bias = nn.Parameter(torch.zeros(out_dim)) if bias else None
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
return F.linear(x, self.weight, self.bias)
|
||||||
|
|
||||||
|
|
||||||
|
class RMSNorm(nn.Module):
|
||||||
|
def __init__(self, dim, norm_eps):
|
||||||
|
super().__init__()
|
||||||
|
self.weight = nn.Parameter(torch.ones(dim))
|
||||||
|
self.normalized_shape = (dim,)
|
||||||
|
self.norm_eps = norm_eps
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
return F.rms_norm(x, self.normalized_shape, self.weight, self.norm_eps)
|
||||||
|
|
||||||
|
|
||||||
|
class MLP(nn.Module):
|
||||||
|
def __init__(self, dim: int, dim_feed_forward: int):
|
||||||
|
super().__init__()
|
||||||
|
self.up = Linear(dim, dim_feed_forward)
|
||||||
|
self.gate = Linear(dim, dim_feed_forward)
|
||||||
|
self.down = Linear(dim_feed_forward, dim)
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
gated = self.up(x) * F.silu(self.gate(x))
|
||||||
|
out = self.down(gated)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
class GQA(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim: int,
|
||||||
|
n_heads: int,
|
||||||
|
n_kv_heads: int,
|
||||||
|
use_qk_norm: bool,
|
||||||
|
norm_eps: float,
|
||||||
|
use_gated_attention: bool,
|
||||||
|
layer_id: int,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
assert dim % n_heads == 0
|
||||||
|
assert n_heads % n_kv_heads == 0
|
||||||
|
|
||||||
|
self.head_dim = dim // n_heads
|
||||||
|
self.layer_id = layer_id
|
||||||
|
self.dim = dim
|
||||||
|
self.n_heads = n_heads
|
||||||
|
self.n_kv_heads = n_kv_heads
|
||||||
|
self.n_rep = n_heads // n_kv_heads
|
||||||
|
self.use_qk_norm = use_qk_norm
|
||||||
|
self.use_gated_attention = use_gated_attention
|
||||||
|
|
||||||
|
self.q_proj = Linear(dim, n_heads * self.head_dim)
|
||||||
|
self.k_proj = Linear(dim, n_kv_heads * self.head_dim)
|
||||||
|
self.v_proj = Linear(dim, n_kv_heads * self.head_dim)
|
||||||
|
self.o_proj = Linear(dim, dim)
|
||||||
|
|
||||||
|
if self.use_qk_norm:
|
||||||
|
self.q_norm = RMSNorm(self.head_dim, norm_eps)
|
||||||
|
self.k_norm = RMSNorm(self.head_dim, norm_eps)
|
||||||
|
|
||||||
|
if self.use_gated_attention:
|
||||||
|
self.gate = Linear(dim, dim)
|
||||||
|
|
||||||
|
def _split_heads(self, x: Tensor, n_heads) -> Tensor:
|
||||||
|
batch_size, seq_len, _ = x.shape
|
||||||
|
x = x.reshape(batch_size, seq_len, n_heads, self.head_dim)
|
||||||
|
return x
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
x: Tensor,
|
||||||
|
rotary_emb: Tensor,
|
||||||
|
attn_mask: Tensor = None,
|
||||||
|
paged_cache: Optional[KvcacheView] = None,
|
||||||
|
) -> Tensor:
|
||||||
|
is_causal = attn_mask is None
|
||||||
|
|
||||||
|
# (bsz, seq_len, dim) -> (bsz, seq_len, n_heads, head_dim)
|
||||||
|
q = self._split_heads(self.q_proj(x), self.n_heads)
|
||||||
|
k = self._split_heads(self.k_proj(x), self.n_kv_heads)
|
||||||
|
v = self._split_heads(self.v_proj(x), self.n_kv_heads)
|
||||||
|
q, k = apply_rotary_emb(q, rotary_emb), apply_rotary_emb(k, rotary_emb)
|
||||||
|
|
||||||
|
if self.use_qk_norm:
|
||||||
|
q, k = self.q_norm(q), self.k_norm(k)
|
||||||
|
|
||||||
|
if paged_cache is not None:
|
||||||
|
paged_cache.write(self.layer_id, k, v)
|
||||||
|
k, v = paged_cache.gather(self.layer_id)
|
||||||
|
|
||||||
|
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)
|
||||||
|
|
||||||
|
# (bsz, seq_len, n_heads, head_dim) -> (bsz, n_heads, seq_len, head_dim)
|
||||||
|
q, k, v = q.permute(0, 2, 1, 3), k.permute(0, 2, 1, 3), v.permute(0, 2, 1, 3)
|
||||||
|
sdqa_out = (
|
||||||
|
F.scaled_dot_product_attention(q, k, v, attn_mask, is_causal=is_causal)
|
||||||
|
.permute(0, 2, 1, 3)
|
||||||
|
.contiguous()
|
||||||
|
.flatten(2)
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.use_gated_attention:
|
||||||
|
sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
|
||||||
|
|
||||||
|
out = self.o_proj(sdqa_out)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
class MLA(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim: int,
|
||||||
|
n_heads: int,
|
||||||
|
n_kv_heads: int,
|
||||||
|
kv_lora_rank: int,
|
||||||
|
qk_nope_head_dim: int,
|
||||||
|
qk_rope_head_dim: int,
|
||||||
|
norm_eps: float,
|
||||||
|
use_gated_attention: bool,
|
||||||
|
layer_id: int,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
self.n_heads = n_heads
|
||||||
|
self.n_kv_heads = n_kv_heads
|
||||||
|
self.kv_lora_rank = kv_lora_rank
|
||||||
|
self.qk_nope_head_dim = qk_nope_head_dim
|
||||||
|
self.qk_rope_head_dim = qk_rope_head_dim
|
||||||
|
self.head_dim = qk_nope_head_dim + qk_rope_head_dim
|
||||||
|
self.layer_id = layer_id
|
||||||
|
self.n_rep = n_heads // n_kv_heads
|
||||||
|
self.use_gated_attention = use_gated_attention
|
||||||
|
|
||||||
|
self.q_proj = Linear(dim, n_heads * self.head_dim, bias=False)
|
||||||
|
self.kv_a_proj = Linear(dim, kv_lora_rank, bias=False)
|
||||||
|
self.kv_norm = RMSNorm(kv_lora_rank, norm_eps)
|
||||||
|
|
||||||
|
# fused KV: (k_nope, k_rope, v)
|
||||||
|
self.kv_b_proj = Linear(
|
||||||
|
kv_lora_rank,
|
||||||
|
n_kv_heads * (self.head_dim + qk_rope_head_dim + self.head_dim),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.o_proj = Linear(dim, dim, bias=False)
|
||||||
|
|
||||||
|
if use_gated_attention:
|
||||||
|
self.gate = Linear(dim, dim, bias=False)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
x: Tensor,
|
||||||
|
rotary_emb: Tensor,
|
||||||
|
attn_mask: Tensor = None,
|
||||||
|
paged_cache: Optional[KvcacheView] = None,
|
||||||
|
) -> Tensor:
|
||||||
|
bsz, seq_len, _ = x.size()
|
||||||
|
is_causal = attn_mask is None
|
||||||
|
|
||||||
|
q = self.q_proj(x)
|
||||||
|
q = q.view(bsz, seq_len, self.n_heads, self.head_dim)
|
||||||
|
|
||||||
|
kv_compressed = self.kv_a_proj(x)
|
||||||
|
kv_compressed = self.kv_norm(kv_compressed)
|
||||||
|
|
||||||
|
kv = self.kv_b_proj(kv_compressed)
|
||||||
|
kv = kv.view(bsz, seq_len, self.n_kv_heads, -1)
|
||||||
|
|
||||||
|
k_nope, k_rope, v = torch.split(
|
||||||
|
kv, [self.qk_nope_head_dim, self.qk_rope_head_dim, self.head_dim], dim=-1
|
||||||
|
)
|
||||||
|
|
||||||
|
q_nope, q_rope = (
|
||||||
|
q[..., : self.qk_nope_head_dim],
|
||||||
|
q[..., self.qk_rope_head_dim :],
|
||||||
|
)
|
||||||
|
q_rope = apply_rotary_emb(q_rope, rotary_emb)
|
||||||
|
k_rope = apply_rotary_emb(k_rope, rotary_emb)
|
||||||
|
|
||||||
|
q = torch.cat([q_nope, q_rope], dim=-1)
|
||||||
|
k = torch.cat([k_nope, k_rope], dim=-1)
|
||||||
|
|
||||||
|
if paged_cache is not None:
|
||||||
|
paged_cache.write(self.layer_id, k, v)
|
||||||
|
k, v = paged_cache.gather(self.layer_id)
|
||||||
|
|
||||||
|
q = q.permute(0, 2, 1, 3)
|
||||||
|
k = k.permute(0, 2, 1, 3)
|
||||||
|
v = v.permute(0, 2, 1, 3)
|
||||||
|
|
||||||
|
attn_out = F.scaled_dot_product_attention(
|
||||||
|
q, k, v, attn_mask, is_causal=is_causal
|
||||||
|
)
|
||||||
|
attn_out = attn_out.permute(0, 2, 1, 3).contiguous().flatten(2)
|
||||||
|
|
||||||
|
if self.use_gated_attention:
|
||||||
|
attn_out = attn_out * F.sigmoid(self.gate(x))
|
||||||
|
|
||||||
|
out = self.o_proj(attn_out)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
class DecoderBlock(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim: int,
|
||||||
|
n_heads: int,
|
||||||
|
dim_ffn: int,
|
||||||
|
n_kv_heads: int,
|
||||||
|
norm_eps: int,
|
||||||
|
use_qk_norm: bool,
|
||||||
|
use_gated_attention: bool,
|
||||||
|
layer_id: int,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.attention = GQA(
|
||||||
|
dim,
|
||||||
|
n_heads,
|
||||||
|
n_kv_heads,
|
||||||
|
use_qk_norm,
|
||||||
|
norm_eps,
|
||||||
|
use_gated_attention,
|
||||||
|
layer_id,
|
||||||
|
)
|
||||||
|
self.input_norm = RMSNorm(dim, norm_eps)
|
||||||
|
self.mlp = MLP(dim, dim_ffn)
|
||||||
|
self.post_attention_norm = RMSNorm(dim, norm_eps)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
x: Tensor,
|
||||||
|
rotary_emb: Tensor,
|
||||||
|
attention_mask: Optional[Tensor] = None,
|
||||||
|
paged_cache: Optional[KvcacheView] = None,
|
||||||
|
) -> Tensor:
|
||||||
|
attn_output = self.attention(
|
||||||
|
self.input_norm(x),
|
||||||
|
rotary_emb,
|
||||||
|
attention_mask,
|
||||||
|
paged_cache,
|
||||||
|
)
|
||||||
|
x = attn_output + x
|
||||||
|
x = self.mlp(self.post_attention_norm(x)) + x
|
||||||
|
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
class Embedding(nn.Module):
|
||||||
|
def __init__(self, vocab_size: int, embedding_dim: int):
|
||||||
|
super().__init__()
|
||||||
|
self.weight = nn.Parameter(torch.empty((vocab_size, embedding_dim)))
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
return F.embedding(x, self.weight)
|
||||||
@@ -0,0 +1,140 @@
|
|||||||
|
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.core.cache import KvcacheView
|
||||||
|
from astrai.model.automodel import AutoModel
|
||||||
|
from astrai.model.module import (
|
||||||
|
DecoderBlock,
|
||||||
|
Embedding,
|
||||||
|
Linear,
|
||||||
|
RMSNorm,
|
||||||
|
RotaryEmbedding,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def process_attention_mask(
|
||||||
|
input_tensor: Tensor,
|
||||||
|
position_ids: Optional[Tensor],
|
||||||
|
input_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Optional[Tensor]:
|
||||||
|
if position_ids is None:
|
||||||
|
return None
|
||||||
|
if input_mask is not None and input_mask.dim() > 2:
|
||||||
|
return input_mask
|
||||||
|
|
||||||
|
device = input_tensor.device
|
||||||
|
dtype = input_tensor.dtype
|
||||||
|
B, S = input_tensor.size()[:2]
|
||||||
|
T = position_ids.max().item() + 1
|
||||||
|
|
||||||
|
if input_mask is None:
|
||||||
|
if position_ids.min().item() == 0 and is_causal:
|
||||||
|
return None
|
||||||
|
pad = torch.ones(B, T, dtype=torch.bool, device=device)
|
||||||
|
else:
|
||||||
|
pad = input_mask[:, :T].to(device=device, dtype=torch.bool)
|
||||||
|
|
||||||
|
attend = pad.view(B, 1, T).expand(B, S, T).clone()
|
||||||
|
if is_causal:
|
||||||
|
attend &= position_ids.unsqueeze(-1) >= torch.arange(T, device=device)
|
||||||
|
|
||||||
|
return torch.full(
|
||||||
|
(B, 1, S, T), -torch.finfo(dtype).max / 2, dtype=dtype, device=device
|
||||||
|
).masked_fill_(attend.unsqueeze(1), 0.0)
|
||||||
|
|
||||||
|
|
||||||
|
@AutoModel.register("transformer")
|
||||||
|
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[KvcacheView] = None,
|
||||||
|
position_ids: Optional[Tensor] = None,
|
||||||
|
) -> Tensor:
|
||||||
|
assert input_ids.ndim == 2
|
||||||
|
|
||||||
|
x = self.embed_tokens(input_ids)
|
||||||
|
rotary_emb = self.rotary_embedding(x, position_ids)
|
||||||
|
attn_mask = process_attention_mask(x, position_ids, input_mask, is_causal=True)
|
||||||
|
|
||||||
|
for layer in self.layers:
|
||||||
|
x = layer(x, rotary_emb, attn_mask, paged_cache)
|
||||||
|
|
||||||
|
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",
|
||||||
|
]
|
||||||
@@ -1,10 +1,10 @@
|
|||||||
|
from typing import Dict
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
import torch.distributed as dist
|
|
||||||
|
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
from typing import Dict
|
|
||||||
|
|
||||||
|
|
||||||
class ParallelModel(nn.Module):
|
class ParallelModel(nn.Module):
|
||||||
@@ -22,7 +22,7 @@ class RowParallelLinear(ParallelModel):
|
|||||||
in_features: int,
|
in_features: int,
|
||||||
out_features: int,
|
out_features: int,
|
||||||
bias: bool = True,
|
bias: bool = True,
|
||||||
reduce_results: bool = True
|
reduce_results: bool = True,
|
||||||
):
|
):
|
||||||
super().__init__(process_group)
|
super().__init__(process_group)
|
||||||
|
|
||||||
@@ -32,7 +32,9 @@ class RowParallelLinear(ParallelModel):
|
|||||||
self.reduce_results = reduce_results
|
self.reduce_results = reduce_results
|
||||||
|
|
||||||
if in_features % self.world_size != 0:
|
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}")
|
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.weight = nn.Parameter(torch.empty(out_features, self.in_features_per_rank))
|
||||||
self.bias = nn.Parameter(torch.zeros(out_features)) if bias else None
|
self.bias = nn.Parameter(torch.zeros(out_features)) if bias else None
|
||||||
@@ -49,8 +51,8 @@ class RowParallelLinear(ParallelModel):
|
|||||||
return output
|
return output
|
||||||
|
|
||||||
def load_state_dict(self, state_dict: Dict[str, Tensor]):
|
def load_state_dict(self, state_dict: Dict[str, Tensor]):
|
||||||
full_weight = state_dict.get('weight')
|
full_weight = state_dict.get("weight")
|
||||||
full_bias = state_dict.get('bias')
|
full_bias = state_dict.get("bias")
|
||||||
|
|
||||||
start_idx = self.rank * self.in_features_per_rank
|
start_idx = self.rank * self.in_features_per_rank
|
||||||
end_idx = start_idx + self.in_features_per_rank
|
end_idx = start_idx + self.in_features_per_rank
|
||||||
@@ -68,7 +70,7 @@ class ColumnParallelLinear(ParallelModel):
|
|||||||
in_features: int,
|
in_features: int,
|
||||||
out_features: int,
|
out_features: int,
|
||||||
bias: bool = True,
|
bias: bool = True,
|
||||||
gather_results: bool = True
|
gather_results: bool = True,
|
||||||
):
|
):
|
||||||
super().__init__(process_group)
|
super().__init__(process_group)
|
||||||
|
|
||||||
@@ -78,10 +80,16 @@ class ColumnParallelLinear(ParallelModel):
|
|||||||
self.gather_results = gather_results
|
self.gather_results = gather_results
|
||||||
|
|
||||||
if out_features % self.world_size != 0:
|
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}")
|
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.weight = nn.Parameter(
|
||||||
self.bias = nn.Parameter(torch.zeros(self.out_features_per_rank)) if bias else None
|
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:
|
def forward(self, input: Tensor) -> Tensor:
|
||||||
output = F.linear(input, self.weight, self.bias)
|
output = F.linear(input, self.weight, self.bias)
|
||||||
@@ -94,8 +102,8 @@ class ColumnParallelLinear(ParallelModel):
|
|||||||
return output
|
return output
|
||||||
|
|
||||||
def load_state_dict(self, state_dict: Dict[str, Tensor]):
|
def load_state_dict(self, state_dict: Dict[str, Tensor]):
|
||||||
full_weight = state_dict.get('weight')
|
full_weight = state_dict.get("weight")
|
||||||
full_bias = state_dict.get('bias')
|
full_bias = state_dict.get("bias")
|
||||||
|
|
||||||
start_idx = self.rank * self.out_features_per_rank
|
start_idx = self.rank * self.out_features_per_rank
|
||||||
end_idx = start_idx + self.out_features_per_rank
|
end_idx = start_idx + self.out_features_per_rank
|
||||||
@@ -1,28 +1,31 @@
|
|||||||
import os
|
import os
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from functools import wraps
|
||||||
|
from typing import Callable
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
import torch.multiprocessing as mp
|
import torch.multiprocessing as mp
|
||||||
|
|
||||||
from functools import wraps
|
|
||||||
from contextlib import contextmanager
|
|
||||||
from typing import Callable, List, Optional
|
|
||||||
|
|
||||||
|
|
||||||
def get_current_device():
|
def get_current_device():
|
||||||
return os.environ["LOCAL_DEVICE"]
|
return os.environ["LOCAL_DEVICE"]
|
||||||
|
|
||||||
|
|
||||||
def get_world_size() -> int:
|
def get_world_size() -> int:
|
||||||
if dist.is_available() and dist.is_initialized():
|
if dist.is_available() and dist.is_initialized():
|
||||||
return dist.get_world_size()
|
return dist.get_world_size()
|
||||||
else:
|
else:
|
||||||
return 1
|
return 1
|
||||||
|
|
||||||
|
|
||||||
def get_rank() -> int:
|
def get_rank() -> int:
|
||||||
if dist.is_available() and dist.is_initialized():
|
if dist.is_available() and dist.is_initialized():
|
||||||
return dist.get_rank()
|
return dist.get_rank()
|
||||||
else:
|
else:
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def setup_parallel(
|
def setup_parallel(
|
||||||
rank: int,
|
rank: int,
|
||||||
@@ -31,7 +34,6 @@ def setup_parallel(
|
|||||||
master_addr: str = "localhost",
|
master_addr: str = "localhost",
|
||||||
master_port: str = "29500",
|
master_port: str = "29500",
|
||||||
device_type: str = "cuda",
|
device_type: str = "cuda",
|
||||||
device_ids: Optional[List[int]] = None
|
|
||||||
):
|
):
|
||||||
|
|
||||||
if dist.is_available() and dist.is_initialized():
|
if dist.is_available() and dist.is_initialized():
|
||||||
@@ -42,30 +44,22 @@ def setup_parallel(
|
|||||||
yield None
|
yield None
|
||||||
return
|
return
|
||||||
|
|
||||||
if device_ids is None:
|
device_id = torch.device(device_type, rank)
|
||||||
device_ids = [i for i in range(world_size)]
|
|
||||||
|
|
||||||
rank = device_ids[rank % len(device_ids)]
|
os.environ["MASTER_ADDR"] = master_addr
|
||||||
device_id = torch.device(device_type, device_ids[rank])
|
os.environ["MASTER_PORT"] = master_port
|
||||||
|
os.environ["LOCAL_RANK"] = str(rank)
|
||||||
os.environ['MASTER_ADDR'] = master_addr
|
os.environ["WORLD_SIZE"] = str(world_size)
|
||||||
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)
|
os.environ["LOCAL_DEVICE"] = str(device_id)
|
||||||
|
|
||||||
dist.init_process_group(
|
dist.init_process_group(
|
||||||
rank=rank,
|
rank=rank, world_size=world_size, backend=backend, device_id=device_id
|
||||||
world_size=world_size,
|
|
||||||
backend=backend,
|
|
||||||
device_id=device_id
|
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if backend == "nccl" and torch.cuda.is_available():
|
if backend == "nccl" and torch.cuda.is_available():
|
||||||
torch.cuda.set_device(device_id)
|
torch.cuda.set_device(device_id)
|
||||||
elif backend == "ccl" and hasattr(torch, 'xpu') and torch.xpu.is_available():
|
elif backend == "ccl" and hasattr(torch, "xpu") and torch.xpu.is_available():
|
||||||
torch.xpu.set_device(device_id)
|
torch.xpu.set_device(device_id)
|
||||||
|
|
||||||
yield dist.group.WORLD
|
yield dist.group.WORLD
|
||||||
@@ -73,6 +67,7 @@ def setup_parallel(
|
|||||||
if dist.is_initialized():
|
if dist.is_initialized():
|
||||||
dist.destroy_process_group()
|
dist.destroy_process_group()
|
||||||
|
|
||||||
|
|
||||||
def only_on_rank(rank, sync=False):
|
def only_on_rank(rank, sync=False):
|
||||||
"""
|
"""
|
||||||
decorator to run a function only on a specific rank.
|
decorator to run a function only on a specific rank.
|
||||||
@@ -81,15 +76,20 @@ def only_on_rank(rank, sync=False):
|
|||||||
def decorator(func):
|
def decorator(func):
|
||||||
@wraps(func)
|
@wraps(func)
|
||||||
def wrapper(*args, **kwargs):
|
def wrapper(*args, **kwargs):
|
||||||
|
ret_args = None
|
||||||
if get_rank() == rank:
|
if get_rank() == rank:
|
||||||
return func(*args, **kwargs)
|
ret_args = func(*args, **kwargs)
|
||||||
if sync:
|
|
||||||
|
if sync and dist.is_available() and dist.is_initialized():
|
||||||
dist.barrier()
|
dist.barrier()
|
||||||
|
|
||||||
|
return ret_args
|
||||||
|
|
||||||
return wrapper
|
return wrapper
|
||||||
|
|
||||||
return decorator
|
return decorator
|
||||||
|
|
||||||
|
|
||||||
def wrapper_spawn_func(
|
def wrapper_spawn_func(
|
||||||
rank: int,
|
rank: int,
|
||||||
world_size: int,
|
world_size: int,
|
||||||
@@ -97,9 +97,8 @@ def wrapper_spawn_func(
|
|||||||
master_addr: str,
|
master_addr: str,
|
||||||
master_port: str,
|
master_port: str,
|
||||||
device_type: str,
|
device_type: str,
|
||||||
device_ids: List[int],
|
|
||||||
func: Callable,
|
func: Callable,
|
||||||
kwargs: dict
|
kwargs: dict,
|
||||||
):
|
):
|
||||||
try:
|
try:
|
||||||
with setup_parallel(
|
with setup_parallel(
|
||||||
@@ -109,7 +108,6 @@ def wrapper_spawn_func(
|
|||||||
master_addr=master_addr,
|
master_addr=master_addr,
|
||||||
master_port=master_port,
|
master_port=master_port,
|
||||||
device_type=device_type,
|
device_type=device_type,
|
||||||
device_ids=device_ids
|
|
||||||
):
|
):
|
||||||
func(**kwargs)
|
func(**kwargs)
|
||||||
|
|
||||||
@@ -117,6 +115,7 @@ def wrapper_spawn_func(
|
|||||||
print(f"Error in rank {rank}: {e}")
|
print(f"Error in rank {rank}: {e}")
|
||||||
raise
|
raise
|
||||||
|
|
||||||
|
|
||||||
def spawn_parallel_fn(
|
def spawn_parallel_fn(
|
||||||
func: Callable,
|
func: Callable,
|
||||||
world_size: int,
|
world_size: int,
|
||||||
@@ -124,28 +123,39 @@ def spawn_parallel_fn(
|
|||||||
master_addr: str = "localhost",
|
master_addr: str = "localhost",
|
||||||
master_port: str = "29500",
|
master_port: str = "29500",
|
||||||
device_type: str = "cuda",
|
device_type: str = "cuda",
|
||||||
device_ids: Optional[List[int]] = None,
|
**kwargs,
|
||||||
**kwargs
|
|
||||||
):
|
):
|
||||||
# clear environment variables
|
# clear environment variables
|
||||||
for key in ['MASTER_ADDR', 'MASTER_PORT', 'RANK', 'WORLD_SIZE', 'LOCAL_RANK', 'LOCAL_DEVICE']:
|
for key in [
|
||||||
|
"MASTER_ADDR",
|
||||||
|
"MASTER_PORT",
|
||||||
|
"RANK",
|
||||||
|
"WORLD_SIZE",
|
||||||
|
"LOCAL_RANK",
|
||||||
|
"LOCAL_DEVICE",
|
||||||
|
]:
|
||||||
if key in os.environ:
|
if key in os.environ:
|
||||||
del os.environ[key]
|
del os.environ[key]
|
||||||
|
|
||||||
if world_size == 1:
|
if world_size == 1:
|
||||||
device_ids = device_ids or [0]
|
device_id = torch.device(device_type, 0)
|
||||||
deice_id = torch.device(device_type, device_ids[0])
|
os.environ["LOCAL_RANK"] = "0"
|
||||||
os.environ["LOCAL_DEVICE"] = str(deice_id)
|
os.environ["WORLD_SIZE"] = "1"
|
||||||
|
os.environ["LOCAL_DEVICE"] = str(device_id)
|
||||||
|
|
||||||
func(**kwargs)
|
func(**kwargs)
|
||||||
return
|
return
|
||||||
|
|
||||||
wrapper_spawn_func_args = (world_size, backend, master_addr, master_port,
|
wrapper_spawn_func_args = (
|
||||||
device_type, device_ids, func, kwargs)
|
world_size,
|
||||||
|
backend,
|
||||||
|
master_addr,
|
||||||
|
master_port,
|
||||||
|
device_type,
|
||||||
|
func,
|
||||||
|
kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
mp.spawn(
|
mp.spawn(
|
||||||
wrapper_spawn_func,
|
wrapper_spawn_func, nprocs=world_size, args=wrapper_spawn_func_args, join=True
|
||||||
nprocs=world_size,
|
|
||||||
args=wrapper_spawn_func_args,
|
|
||||||
join=True
|
|
||||||
)
|
)
|
||||||
@@ -1,10 +1,12 @@
|
|||||||
import json
|
import json
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
|
import safetensors.torch as st
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
|
|
||||||
from pathlib import Path
|
from astrai.parallel.setup import get_rank
|
||||||
from typing import Dict, Any
|
|
||||||
from khaosz.parallel.setup import get_rank
|
|
||||||
|
|
||||||
|
|
||||||
class Checkpoint:
|
class Checkpoint:
|
||||||
@@ -13,10 +15,12 @@ class Checkpoint:
|
|||||||
state_dict: Dict[str, Any],
|
state_dict: Dict[str, Any],
|
||||||
epoch: int = 0,
|
epoch: int = 0,
|
||||||
iteration: int = 0,
|
iteration: int = 0,
|
||||||
|
extra: Optional[Dict[str, Any]] = None,
|
||||||
):
|
):
|
||||||
self.state_dict = state_dict
|
self.state_dict = state_dict
|
||||||
self.epoch = epoch
|
self.epoch = epoch
|
||||||
self.iteration = iteration
|
self.iteration = iteration
|
||||||
|
self.extra = extra or {}
|
||||||
|
|
||||||
def save(
|
def save(
|
||||||
self,
|
self,
|
||||||
@@ -35,8 +39,9 @@ class Checkpoint:
|
|||||||
with open(save_path / "meta.json", "w") as f:
|
with open(save_path / "meta.json", "w") as f:
|
||||||
json.dump(meta, f, indent=2)
|
json.dump(meta, f, indent=2)
|
||||||
|
|
||||||
with open(save_path / f"state_dict.pt", "wb") as f:
|
st.save_file(self.state_dict, save_path / "state_dict.safetensors")
|
||||||
torch.save(self.state_dict, f)
|
if self.extra:
|
||||||
|
torch.save(self.extra, save_path / "extra.pt")
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def load(
|
def load(
|
||||||
@@ -57,11 +62,16 @@ class Checkpoint:
|
|||||||
dist.broadcast_object_list(meta_list, src=0)
|
dist.broadcast_object_list(meta_list, src=0)
|
||||||
meta = meta_list[0]
|
meta = meta_list[0]
|
||||||
|
|
||||||
with open(save_path / f"state_dict.pt", "rb") as f:
|
state_dict = st.load_file(save_path / "state_dict.safetensors")
|
||||||
state_dict = torch.load(f)
|
|
||||||
|
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(
|
return cls(
|
||||||
state_dict=state_dict,
|
state_dict=state_dict,
|
||||||
epoch=meta["epoch"],
|
epoch=meta["epoch"],
|
||||||
iteration=meta["iteration"],
|
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",
|
||||||
|
]
|
||||||
@@ -0,0 +1,71 @@
|
|||||||
|
from typing import Any, Callable, Dict
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
|
||||||
|
def _grad_stat(
|
||||||
|
model: nn.Module, fn: Callable[[torch.Tensor], Any], default: Any
|
||||||
|
) -> dict:
|
||||||
|
results = {}
|
||||||
|
for name, param in model.named_parameters():
|
||||||
|
results[name] = default
|
||||||
|
if param.grad is not None:
|
||||||
|
results[name] = fn(param.grad.data)
|
||||||
|
return results
|
||||||
|
|
||||||
|
|
||||||
|
def grad_norm(model: nn.Module, norm_type: int = 2) -> Dict[str, float]:
|
||||||
|
return _grad_stat(model, lambda g: g.norm(norm_type).item(), 0.0)
|
||||||
|
|
||||||
|
|
||||||
|
def grad_std(model: nn.Module) -> Dict[str, float]:
|
||||||
|
return _grad_stat(model, lambda g: g.std().item(), 0.0)
|
||||||
|
|
||||||
|
|
||||||
|
def grad_max(model: nn.Module) -> Dict[str, float]:
|
||||||
|
return _grad_stat(model, lambda g: g.max().item(), -float("inf"))
|
||||||
|
|
||||||
|
|
||||||
|
def grad_min(model: nn.Module) -> Dict[str, float]:
|
||||||
|
return _grad_stat(model, lambda g: g.min().item(), float("inf"))
|
||||||
|
|
||||||
|
|
||||||
|
def grad_mean(model: nn.Module) -> Dict[str, float]:
|
||||||
|
return _grad_stat(model, lambda g: g.mean().item(), 0.0)
|
||||||
|
|
||||||
|
|
||||||
|
def grad_nan_num(model: nn.Module) -> Dict[str, int]:
|
||||||
|
return _grad_stat(model, lambda g: g.isnan().sum().item(), 0)
|
||||||
|
|
||||||
|
|
||||||
|
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
|
||||||
@@ -1,120 +1,127 @@
|
|||||||
import os
|
|
||||||
import json
|
import json
|
||||||
|
import os
|
||||||
import time
|
import time
|
||||||
import torch.nn as nn
|
|
||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from tqdm import tqdm
|
from typing import Callable, List, Optional, Protocol, runtime_checkable
|
||||||
from torch.nn.utils import clip_grad_norm_
|
|
||||||
from torch.optim.lr_scheduler import LRScheduler
|
|
||||||
from typing import Callable, List, Optional, Protocol
|
|
||||||
|
|
||||||
from khaosz.parallel import only_on_rank
|
import torch.nn as nn
|
||||||
from khaosz.trainer.metric_util import (
|
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_loss,
|
||||||
ctx_get_lr,
|
ctx_get_lr,
|
||||||
ctx_get_grad_max,
|
|
||||||
ctx_get_grad_min,
|
|
||||||
ctx_get_grad_norm,
|
|
||||||
ctx_get_grad_mean,
|
|
||||||
ctx_get_grad_std,
|
|
||||||
ctx_get_grad_nan_num
|
|
||||||
)
|
)
|
||||||
from khaosz.data.checkpoint import Checkpoint
|
from astrai.trainer.train_context import TrainContext
|
||||||
from khaosz.trainer.train_context import TrainContext
|
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
class TrainCallback(Protocol):
|
class TrainCallback(Protocol):
|
||||||
"""
|
"""
|
||||||
Callback interface for trainer.
|
Callback interface for trainer.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def on_train_begin(self, context: TrainContext):
|
def on_train_begin(self, context: TrainContext):
|
||||||
""" Called at the beginning of training. """
|
"""Called at the beginning of training."""
|
||||||
|
|
||||||
def on_train_end(self, context: TrainContext):
|
def on_train_end(self, context: TrainContext):
|
||||||
""" Called at the end of training. """
|
"""Called at the end of training."""
|
||||||
|
|
||||||
def on_epoch_begin(self, context: TrainContext):
|
def on_epoch_begin(self, context: TrainContext):
|
||||||
""" Called at the beginning of each epoch. """
|
"""Called at the beginning of each epoch."""
|
||||||
|
|
||||||
def on_epoch_end(self, context: TrainContext):
|
def on_epoch_end(self, context: TrainContext):
|
||||||
""" Called at the end of each epoch. """
|
"""Called at the end of each epoch."""
|
||||||
|
|
||||||
def on_step_begin(self, context: TrainContext):
|
def on_step_begin(self, context: TrainContext):
|
||||||
""" Called at the beginning of each step. """
|
"""Called at the beginning of each step."""
|
||||||
|
|
||||||
def on_step_end(self, context: TrainContext):
|
def on_step_end(self, context: TrainContext):
|
||||||
""" Called at the end of each step."""
|
"""Called at the end of each step."""
|
||||||
|
|
||||||
def on_batch_begin(self, context: TrainContext):
|
def on_batch_begin(self, context: TrainContext):
|
||||||
""" Called at the beginning of each batch. """
|
"""Called at the beginning of each batch."""
|
||||||
|
|
||||||
def on_batch_end(self, context: TrainContext):
|
def on_batch_end(self, context: TrainContext):
|
||||||
""" Called at the end of each batch. """
|
"""Called at the end of each batch."""
|
||||||
|
|
||||||
def on_error(self, context: TrainContext):
|
def on_error(self, context: TrainContext):
|
||||||
""" Called when an error occurs during training. """
|
"""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)
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
@CallbackFactory.register("gradient_clipping")
|
||||||
class GradientClippingCallback(TrainCallback):
|
class GradientClippingCallback(TrainCallback):
|
||||||
"""
|
"""
|
||||||
Gradient clipping callback for trainer.
|
Gradient clipping callback for trainer.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, max_grad_norm: float):
|
def __init__(self, max_grad_norm: float):
|
||||||
self.max_grad_norm = max_grad_norm
|
self.max_grad_norm = max_grad_norm
|
||||||
|
|
||||||
def on_step_begin(self, context: TrainContext):
|
def on_step_end(self, context: TrainContext):
|
||||||
_ = context
|
_ = context
|
||||||
clip_grad_norm_(context.model.parameters(), self.max_grad_norm)
|
clip_grad_norm_(context.model.parameters(), self.max_grad_norm)
|
||||||
|
|
||||||
|
|
||||||
class SchedulerCallback(TrainCallback):
|
@CallbackFactory.register("checkpoint")
|
||||||
"""
|
|
||||||
Scheduler callback for trainer.
|
|
||||||
"""
|
|
||||||
def __init__(self):
|
|
||||||
self.scheduler: LRScheduler = None
|
|
||||||
|
|
||||||
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"]
|
|
||||||
|
|
||||||
self.scheduler = context.scheduler
|
|
||||||
|
|
||||||
def on_batch_end(self, context: TrainContext):
|
|
||||||
_ = context
|
|
||||||
if self.scheduler:
|
|
||||||
self.scheduler.step()
|
|
||||||
|
|
||||||
|
|
||||||
class CheckpointCallback(TrainCallback):
|
class CheckpointCallback(TrainCallback):
|
||||||
"""
|
"""
|
||||||
Checkpoint callback for trainer.
|
Checkpoint callback for trainer.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
save_dir: str,
|
save_dir: str,
|
||||||
interval: int,
|
interval: int,
|
||||||
weight_only: bool = False,
|
weight_only: bool = False,
|
||||||
state_dict_fn: Optional[Callable[[nn.Module], dict]] = None
|
state_dict_fn: Optional[Callable[[nn.Module], dict]] = None,
|
||||||
|
save_extra_fn: Optional[Callable[["TrainContext"], dict]] = None,
|
||||||
):
|
):
|
||||||
self.save_dir = save_dir
|
self.save_dir = save_dir
|
||||||
self.interval = interval
|
self.interval = interval
|
||||||
self.weight_only = weight_only
|
self.weight_only = weight_only
|
||||||
self.state_dict_fn = state_dict_fn
|
self.state_dict_fn = state_dict_fn
|
||||||
|
self.save_extra_fn = save_extra_fn
|
||||||
self.last_ckpt_iter = 0
|
self.last_ckpt_iter = 0
|
||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
def _save_checkpoint(self, context: TrainContext):
|
def _save_checkpoint(self, context: TrainContext):
|
||||||
save_path = os.path.join(self.save_dir, f"epoch_{context.epoch}_iter_{context.iteration}")
|
save_path = os.path.join(
|
||||||
state_dict = self.state_dict_fn(context.model) if self.state_dict_fn else context.model.state_dict()
|
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(
|
context.checkpoint = Checkpoint(
|
||||||
state_dict=state_dict,
|
state_dict=state_dict,
|
||||||
epoch=context.epoch,
|
epoch=context.epoch,
|
||||||
iteration=context.iteration
|
iteration=context.iteration,
|
||||||
|
extra=extra,
|
||||||
)
|
)
|
||||||
|
|
||||||
context.checkpoint.save(save_path)
|
context.checkpoint.save(save_path)
|
||||||
@@ -132,10 +139,12 @@ class CheckpointCallback(TrainCallback):
|
|||||||
self._save_checkpoint(context)
|
self._save_checkpoint(context)
|
||||||
|
|
||||||
|
|
||||||
|
@CallbackFactory.register("progress_bar")
|
||||||
class ProgressBarCallback(TrainCallback):
|
class ProgressBarCallback(TrainCallback):
|
||||||
"""
|
"""
|
||||||
Progress bar callback for trainer.
|
Progress bar callback for trainer.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, num_epoch: int):
|
def __init__(self, num_epoch: int):
|
||||||
self.num_epoch = num_epoch
|
self.num_epoch = num_epoch
|
||||||
self.progress_bar: tqdm = None
|
self.progress_bar: tqdm = None
|
||||||
@@ -144,16 +153,18 @@ class ProgressBarCallback(TrainCallback):
|
|||||||
def on_epoch_begin(self, context: TrainContext):
|
def on_epoch_begin(self, context: TrainContext):
|
||||||
self.progress_bar = tqdm(
|
self.progress_bar = tqdm(
|
||||||
context.dataloader,
|
context.dataloader,
|
||||||
desc=f"Epoch {context.epoch+1}/{self.num_epoch}",
|
desc=f"Epoch {context.epoch + 1}/{self.num_epoch}",
|
||||||
dynamic_ncols=True
|
dynamic_ncols=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
def on_batch_end(self, context: TrainContext):
|
def on_batch_end(self, context: TrainContext):
|
||||||
self.progress_bar.set_postfix({
|
self.progress_bar.set_postfix(
|
||||||
|
{
|
||||||
"loss": f"{context.loss:.4f}",
|
"loss": f"{context.loss:.4f}",
|
||||||
"lr": f"{context.optimizer.param_groups[-1]['lr']:.2e}"
|
"lr": f"{context.optimizer.param_groups[-1]['lr']:.2e}",
|
||||||
})
|
}
|
||||||
|
)
|
||||||
self.progress_bar.update(1)
|
self.progress_bar.update(1)
|
||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
@@ -163,19 +174,19 @@ class ProgressBarCallback(TrainCallback):
|
|||||||
self.progress_bar.close()
|
self.progress_bar.close()
|
||||||
|
|
||||||
|
|
||||||
|
@CallbackFactory.register("metric_logger")
|
||||||
class MetricLoggerCallback(TrainCallback):
|
class MetricLoggerCallback(TrainCallback):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
log_dir:str,
|
log_dir: str,
|
||||||
save_interval:int,
|
save_interval: int,
|
||||||
log_interval:int=10,
|
log_interval: int = 10,
|
||||||
metrics:List[str]=None
|
metrics: List[str] = None,
|
||||||
):
|
):
|
||||||
self.step_num = 0
|
self.last_log_iter = 0
|
||||||
self.last_save_step = 0
|
|
||||||
self.save_interval = save_interval
|
self.save_interval = save_interval
|
||||||
self.log_interval = log_interval
|
self.log_interval = log_interval
|
||||||
self.metrics = metrics or ['loss', 'lr']
|
self.metrics = metrics or ["loss", "lr"]
|
||||||
|
|
||||||
self.log_dir = Path(log_dir) if log_dir else Path.cwd() / "logs"
|
self.log_dir = Path(log_dir) if log_dir else Path.cwd() / "logs"
|
||||||
self.log_dir.mkdir(parents=True, exist_ok=True)
|
self.log_dir.mkdir(parents=True, exist_ok=True)
|
||||||
@@ -183,22 +194,22 @@ class MetricLoggerCallback(TrainCallback):
|
|||||||
self.log_cache = []
|
self.log_cache = []
|
||||||
|
|
||||||
self._metric_funcs = {
|
self._metric_funcs = {
|
||||||
'loss': ctx_get_loss,
|
"loss": ctx_get_loss,
|
||||||
'lr': ctx_get_lr,
|
"lr": ctx_get_lr,
|
||||||
'grad_norm': ctx_get_grad_norm,
|
"grad_norm": ctx_get_grad_norm,
|
||||||
'grad_std': ctx_get_grad_std,
|
"grad_std": ctx_get_grad_std,
|
||||||
'grad_max': ctx_get_grad_max,
|
"grad_max": ctx_get_grad_max,
|
||||||
'grad_min': ctx_get_grad_min,
|
"grad_min": ctx_get_grad_min,
|
||||||
'grad_mean': ctx_get_grad_mean,
|
"grad_mean": ctx_get_grad_mean,
|
||||||
'grad_nan_num': ctx_get_grad_nan_num
|
"grad_nan_num": ctx_get_grad_nan_num,
|
||||||
}
|
}
|
||||||
|
|
||||||
def _get_log_data(self, context: TrainContext):
|
def _get_log_data(self, context: TrainContext):
|
||||||
return {
|
return {
|
||||||
"timestamp": time.strftime('%Y-%m-%d %H:%M:%S'),
|
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S"),
|
||||||
"epoch": context.epoch,
|
"epoch": context.epoch,
|
||||||
"iter": context.iteration,
|
"iter": context.iteration,
|
||||||
**{m: self._metric_funcs[m](context) for m in self.metrics}
|
**{m: self._metric_funcs[m](context) for m in self.metrics},
|
||||||
}
|
}
|
||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
@@ -209,24 +220,22 @@ class MetricLoggerCallback(TrainCallback):
|
|||||||
def _save_log(self, epoch, iter):
|
def _save_log(self, epoch, iter):
|
||||||
log_file = self.log_dir / f"epoch_{epoch}_iter_{iter}_metric.jsonl"
|
log_file = self.log_dir / f"epoch_{epoch}_iter_{iter}_metric.jsonl"
|
||||||
|
|
||||||
with open(log_file, 'w') as f:
|
with open(log_file, "w") as f:
|
||||||
for log in self.log_cache:
|
for log in self.log_cache:
|
||||||
f.write(json.dumps(log) + '\n')
|
f.write(json.dumps(log) + "\n")
|
||||||
|
|
||||||
def on_batch_end(self, context):
|
def on_batch_end(self, context):
|
||||||
if self.step_num % self.log_interval == 0:
|
if context.iteration % self.log_interval == 0:
|
||||||
log_data = self._get_log_data(context)
|
log_data = self._get_log_data(context)
|
||||||
self._add_log(log_data)
|
self._add_log(log_data)
|
||||||
|
|
||||||
if self.step_num - self.last_save_step >= self.save_interval:
|
if context.iteration - self.last_log_iter >= self.save_interval:
|
||||||
self._save_log(context.epoch, context.iteration)
|
self._save_log(context.epoch, context.iteration)
|
||||||
self.last_save_step = self.step_num
|
self.last_log_iter = context.iteration
|
||||||
|
|
||||||
self.step_num += 1
|
|
||||||
|
|
||||||
def on_train_end(self, context):
|
def on_train_end(self, context):
|
||||||
|
if context.iteration != self.last_log_iter:
|
||||||
self._save_log(context.epoch, context.iteration)
|
self._save_log(context.epoch, context.iteration)
|
||||||
|
|
||||||
def on_error(self, context):
|
def on_error(self, context):
|
||||||
self._save_log(context.epoch, context.iteration)
|
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,99 @@
|
|||||||
|
import logging
|
||||||
|
from itertools import batched
|
||||||
|
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),
|
||||||
|
]
|
||||||
|
|
||||||
|
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()
|
||||||
|
accumulation_steps = max(self.train_config.accumulation_steps, 1)
|
||||||
|
|
||||||
|
for epoch in range(context.epoch, self.train_config.n_epoch):
|
||||||
|
context.epoch = epoch
|
||||||
|
self._call_callbacks("on_epoch_begin", context)
|
||||||
|
|
||||||
|
for steps in batched(context.dataloader, accumulation_steps):
|
||||||
|
self._call_callbacks("on_step_begin", context)
|
||||||
|
|
||||||
|
step_batch_nums = len(steps)
|
||||||
|
for batch in steps:
|
||||||
|
self._call_callbacks("on_batch_begin", context)
|
||||||
|
loss = context.strategy(batch)
|
||||||
|
context.loss = loss.item()
|
||||||
|
context.iteration += 1
|
||||||
|
|
||||||
|
stand_loss = loss / step_batch_nums
|
||||||
|
stand_loss.backward()
|
||||||
|
self._call_callbacks("on_batch_end", context)
|
||||||
|
|
||||||
|
self._call_callbacks("on_step_end", context)
|
||||||
|
context.optimizer.step()
|
||||||
|
context.optimizer.zero_grad()
|
||||||
|
|
||||||
|
if context.scheduler:
|
||||||
|
context.scheduler.step()
|
||||||
|
|
||||||
|
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)
|
||||||
@@ -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("retrieve content:")
|
|
||||||
print("\n".join([f"{idx + 1}. " + text for idx, (text, _) in enumerate(retrieved)]))
|
|
||||||
|
|
||||||
print("\n\nretrive generate:")
|
|
||||||
print(retrive_response)
|
|
||||||
@@ -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,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
|
||||||
@@ -1,59 +0,0 @@
|
|||||||
__version__ = "1.3.2"
|
|
||||||
__author__ = "ViperEkura"
|
|
||||||
|
|
||||||
from khaosz.api import Khaosz
|
|
||||||
from khaosz.config import (
|
|
||||||
ModelConfig,
|
|
||||||
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",
|
|
||||||
"TrainConfig",
|
|
||||||
|
|
||||||
"DatasetLoader",
|
|
||||||
"BpeTokenizer",
|
|
||||||
|
|
||||||
"TextGenerator",
|
|
||||||
"ChatGenerator",
|
|
||||||
"StreamGenerator",
|
|
||||||
"BatchGenerator",
|
|
||||||
"RetrievalGenerator",
|
|
||||||
"EmbeddingEncoder",
|
|
||||||
|
|
||||||
"Trainer",
|
|
||||||
"StrategyFactory",
|
|
||||||
"SchedulerFactory"
|
|
||||||
]
|
|
||||||
-113
@@ -1,113 +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 ModelParameter
|
|
||||||
|
|
||||||
|
|
||||||
class Khaosz:
|
|
||||||
def __init__(self, model_dir: str):
|
|
||||||
self.parameter = ModelParameter()
|
|
||||||
self.parameter.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,16 +0,0 @@
|
|||||||
from khaosz.config.model_config import ModelConfig
|
|
||||||
from khaosz.config.param_config import BaseModelIO, ModelParameter
|
|
||||||
from khaosz.config.schedule_config import ScheduleConfig, CosineScheduleConfig, SGDRScheduleConfig
|
|
||||||
from khaosz.config.train_config import TrainConfig
|
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
"BaseModelIO",
|
|
||||||
"ModelParameter",
|
|
||||||
"ModelConfig",
|
|
||||||
"TrainConfig",
|
|
||||||
|
|
||||||
"ScheduleConfig",
|
|
||||||
"CosineScheduleConfig",
|
|
||||||
"SGDRScheduleConfig",
|
|
||||||
]
|
|
||||||
@@ -1,80 +0,0 @@
|
|||||||
import torch.nn as nn
|
|
||||||
import safetensors.torch as st
|
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
|
||||||
from typing import 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
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class BaseModelIO:
|
|
||||||
"""Base class for model I/O operations."""
|
|
||||||
|
|
||||||
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 _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 self.model is None:
|
|
||||||
self.model = Transformer(self.config)
|
|
||||||
|
|
||||||
if paths["model"].exists():
|
|
||||||
state_dict = st.load_file(str(paths["model"]))
|
|
||||||
self.model.load_state_dict(state_dict)
|
|
||||||
|
|
||||||
return self
|
|
||||||
|
|
||||||
def to(self, *args, **kwargs) -> "BaseModelIO":
|
|
||||||
"""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."""
|
|
||||||
|
|
||||||
def save(self, save_dir: Union[str, Path]):
|
|
||||||
self.save_components(save_dir)
|
|
||||||
|
|
||||||
def load(self, load_dir: Union[str, Path]) -> "ModelParameter":
|
|
||||||
return self.load_components(load_dir)
|
|
||||||
|
|
||||||
@@ -1,93 +0,0 @@
|
|||||||
from typing import Any, 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."}
|
|
||||||
)
|
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
|
||||||
self.schedule_type = "cosine"
|
|
||||||
self.validate()
|
|
||||||
|
|
||||||
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."}
|
|
||||||
)
|
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
|
||||||
self.schedule_type = "sgdr"
|
|
||||||
self.validate()
|
|
||||||
|
|
||||||
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,136 +0,0 @@
|
|||||||
import torch.nn as nn
|
|
||||||
from torch.utils.data import Dataset
|
|
||||||
from torch.optim import Optimizer
|
|
||||||
from torch.optim.lr_scheduler import LRScheduler
|
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
|
||||||
from typing import Callable, List, Optional
|
|
||||||
|
|
||||||
|
|
||||||
@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."}
|
|
||||||
)
|
|
||||||
checkpoint_dir: str = field(
|
|
||||||
default="./checkpoint",
|
|
||||||
metadata={"help": "Checkpoint directory."}
|
|
||||||
)
|
|
||||||
checkpoint_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_ids: Optional[List[int]] = field(
|
|
||||||
default=None,
|
|
||||||
metadata={"help": "Device ids for distributed training."}
|
|
||||||
)
|
|
||||||
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.")
|
|
||||||
|
|
||||||
|
|
||||||
@@ -1,24 +0,0 @@
|
|||||||
from khaosz.data.dataset import (
|
|
||||||
BaseDataset,
|
|
||||||
SeqDataset,
|
|
||||||
DpoDataset,
|
|
||||||
SftDataset,
|
|
||||||
PpoDataset,
|
|
||||||
MultiSegmentFetcher,
|
|
||||||
DatasetLoader
|
|
||||||
)
|
|
||||||
|
|
||||||
from khaosz.data.tokenizer import BpeTokenizer
|
|
||||||
from khaosz.data.sampler import ResumableDistributedSampler
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
"BaseDataset",
|
|
||||||
"SeqDataset",
|
|
||||||
"DpoDataset",
|
|
||||||
"SftDataset",
|
|
||||||
"PpoDataset",
|
|
||||||
"MultiSegmentFetcher",
|
|
||||||
"DatasetLoader",
|
|
||||||
"BpeTokenizer",
|
|
||||||
"ResumableDistributedSampler"
|
|
||||||
]
|
|
||||||
@@ -1,201 +0,0 @@
|
|||||||
import torch
|
|
||||||
import bisect
|
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
|
||||||
from torch import Tensor
|
|
||||||
from torch.utils.data import Dataset
|
|
||||||
from khaosz.data.file import load_h5
|
|
||||||
from typing import Callable, List, Dict, Literal, Optional, Union
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class BaseSegmentFetcher:
|
|
||||||
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:
|
|
||||||
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]))
|
|
||||||
data = self.segments[i][start:end]
|
|
||||||
result_segments.append(data)
|
|
||||||
|
|
||||||
return torch.cat(result_segments, dim=0)
|
|
||||||
|
|
||||||
|
|
||||||
class MultiSegmentFetcher:
|
|
||||||
def __init__(self, muti_segments: Dict):
|
|
||||||
self.muti_keys = list(muti_segments.keys())
|
|
||||||
self.muti_fetchers = {
|
|
||||||
key: BaseSegmentFetcher(segments)
|
|
||||||
for key, segments in muti_segments.items()
|
|
||||||
}
|
|
||||||
|
|
||||||
def __len__(self) -> int:
|
|
||||||
len_list = [len(seg) for seg in self.muti_fetchers.values()]
|
|
||||||
return min(len_list)
|
|
||||||
|
|
||||||
def key_fetch(self, begin_idx: int, end_idx: int, keys: Union[str, List[str]]) -> Dict:
|
|
||||||
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) -> Dict:
|
|
||||||
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 = {}
|
|
||||||
self.window_size = window_size
|
|
||||||
self.stride = stride
|
|
||||||
self.total_samples = None
|
|
||||||
|
|
||||||
def load(self, load_path: str):
|
|
||||||
self.segments = load_h5(load_path)
|
|
||||||
self.fetcher = MultiSegmentFetcher(self.segments)
|
|
||||||
self.total_samples = len(self.fetcher)
|
|
||||||
|
|
||||||
def get_index(self, index: int) -> int:
|
|
||||||
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]:
|
|
||||||
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 SeqDataset(BaseDataset):
|
|
||||||
def __init__(self, window_size: int, stride: int):
|
|
||||||
super().__init__(window_size, stride)
|
|
||||||
self.fetcher = MultiSegmentFetcher(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 = MultiSegmentFetcher(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 = MultiSegmentFetcher(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 = MultiSegmentFetcher(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: str,
|
|
||||||
window_size: int,
|
|
||||||
stride: Optional[int] = None,
|
|
||||||
) -> 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
|
|
||||||
@@ -1,42 +0,0 @@
|
|||||||
import os
|
|
||||||
import h5py
|
|
||||||
import torch
|
|
||||||
|
|
||||||
from pathlib import Path
|
|
||||||
from torch import Tensor
|
|
||||||
from typing import Dict, List
|
|
||||||
|
|
||||||
|
|
||||||
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
|
|
||||||
@@ -1,106 +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()
|
|
||||||
self._tokenizer = Tokenizer(model)
|
|
||||||
self._tokenizer.normalizer = normalizers.Sequence([
|
|
||||||
normalizers.NFC(),
|
|
||||||
normalizers.Strip()
|
|
||||||
])
|
|
||||||
|
|
||||||
self._tokenizer.pre_tokenizer = pre_tokenizers.Sequence([
|
|
||||||
pre_tokenizers.UnicodeScripts(),
|
|
||||||
pre_tokenizers.ByteLevel(add_prefix_space=False, use_regex=True)
|
|
||||||
])
|
|
||||||
|
|
||||||
self._tokenizer.decoder = decoders.ByteLevel()
|
|
||||||
self._tokenizer.post_processor = processors.ByteLevel(trim_offsets=True)
|
|
||||||
|
|
||||||
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, max_token_length=18) -> tuple:
|
|
||||||
assert reserved_token_size > len(self._special_tokens)
|
|
||||||
reserved_tokens = [f"<|reserve{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 // 6,
|
|
||||||
max_token_length=max_token_length,
|
|
||||||
special_tokens=self._control_tokens,
|
|
||||||
initial_alphabet=alphabet,
|
|
||||||
show_progress=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
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 +0,0 @@
|
|||||||
# init file
|
|
||||||
@@ -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.max_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.max_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_layers
|
|
||||||
self.max_len = config.max_len
|
|
||||||
self.num_heads = config.n_kv_heads
|
|
||||||
self.head_dim = config.dim //config.n_heads
|
|
||||||
|
|
||||||
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.max_len, self.num_layers, self.num_heads, self.head_dim),
|
|
||||||
device=self.device, dtype=self.dtype
|
|
||||||
)
|
|
||||||
v_cache = torch.zeros(
|
|
||||||
(self.batch_size, self.max_len, self.num_layers, 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,98 +0,0 @@
|
|||||||
import torch
|
|
||||||
from torch import Tensor
|
|
||||||
from functools import wraps
|
|
||||||
from inspect import signature
|
|
||||||
|
|
||||||
|
|
||||||
class CudaGraphWrapper:
|
|
||||||
def __init__(self, function, device="cuda", cast=False):
|
|
||||||
self.function = function
|
|
||||||
self.cast = cast
|
|
||||||
self.device = device
|
|
||||||
self.static_input = None
|
|
||||||
self.static_output = None
|
|
||||||
self.graph = None
|
|
||||||
self.signature = signature(function)
|
|
||||||
|
|
||||||
def _update_inplace(self, lhs, rhs):
|
|
||||||
if isinstance(lhs, Tensor) and isinstance(rhs, Tensor):
|
|
||||||
if lhs.shape != rhs.shape:
|
|
||||||
raise ValueError(
|
|
||||||
f"Tensor shape mismatch! "
|
|
||||||
f"Expected: {lhs.shape}, Got: {rhs.shape}. "
|
|
||||||
f"Function: {self.function}"
|
|
||||||
)
|
|
||||||
if self.cast:
|
|
||||||
if lhs.device != rhs.device:
|
|
||||||
rhs = rhs.to(device=lhs.device)
|
|
||||||
|
|
||||||
if lhs.dtype != rhs.dtype:
|
|
||||||
rhs = rhs.to(dtype=lhs.dtype)
|
|
||||||
else:
|
|
||||||
if lhs.device != rhs.device:
|
|
||||||
raise ValueError(
|
|
||||||
f"Tensor device mismatch! "
|
|
||||||
f"Expected: {lhs.device}, Got: {rhs.device}. "
|
|
||||||
f"Function: {self.function}"
|
|
||||||
)
|
|
||||||
if lhs.dtype != rhs.dtype:
|
|
||||||
raise ValueError(
|
|
||||||
f"Tensor dtype mismatch! "
|
|
||||||
f"Expected: {lhs.dtype}, Got: {rhs.dtype}. "
|
|
||||||
f"Function: {self.function}"
|
|
||||||
)
|
|
||||||
lhs.copy_(rhs)
|
|
||||||
elif isinstance(lhs, dict):
|
|
||||||
for k in lhs:
|
|
||||||
if k in rhs:
|
|
||||||
self._update_inplace(lhs[k], rhs[k])
|
|
||||||
elif isinstance(lhs, (list, tuple)):
|
|
||||||
for i in range(len(lhs)):
|
|
||||||
if i < len(rhs):
|
|
||||||
self._update_inplace(lhs[i], rhs[i])
|
|
||||||
elif isinstance(lhs, (int, float, bool, str, type(None))):
|
|
||||||
if lhs != rhs:
|
|
||||||
raise ValueError("Does not support changing control parameters.")
|
|
||||||
|
|
||||||
def _update_args(self, input_args, input_kwargs):
|
|
||||||
bound_args = self.signature.bind(*input_args, **input_kwargs)
|
|
||||||
bound_args.apply_defaults()
|
|
||||||
args_dict = bound_args.arguments
|
|
||||||
|
|
||||||
if self.static_input is None:
|
|
||||||
self.static_input = args_dict
|
|
||||||
else:
|
|
||||||
self._update_inplace(self.static_input, args_dict)
|
|
||||||
|
|
||||||
def run(self, *args, **kwargs):
|
|
||||||
self._update_args(args, kwargs)
|
|
||||||
|
|
||||||
if self.graph is None:
|
|
||||||
# warmup
|
|
||||||
_ = torch.matmul(
|
|
||||||
torch.randn(100, 100, device=self.device),
|
|
||||||
torch.randn(100, 100, device=self.device)
|
|
||||||
)
|
|
||||||
torch.cuda.synchronize()
|
|
||||||
|
|
||||||
# capture graph
|
|
||||||
self.graph = torch.cuda.CUDAGraph()
|
|
||||||
with torch.cuda.graph(self.graph):
|
|
||||||
self.static_output = self.function(**self.static_input)
|
|
||||||
|
|
||||||
self.graph.replay()
|
|
||||||
|
|
||||||
return self.static_output
|
|
||||||
|
|
||||||
|
|
||||||
def cuda_graph(device="cuda", cast=False):
|
|
||||||
def decorator(func):
|
|
||||||
wrapper = CudaGraphWrapper(func, device, cast)
|
|
||||||
|
|
||||||
@wraps(func)
|
|
||||||
def wrapped(*args, **kwargs):
|
|
||||||
return wrapper.run(*args, **kwargs)
|
|
||||||
|
|
||||||
return wrapped
|
|
||||||
|
|
||||||
return decorator
|
|
||||||
@@ -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.max_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.max_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,281 +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):
|
|
||||||
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:
|
|
||||||
rms = F.rms_norm(x.float(), self.normalized_shape, self.weight, self.norm_eps)
|
|
||||||
return rms.to(x.dtype)
|
|
||||||
|
|
||||||
|
|
||||||
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 Attention(nn.Module):
|
|
||||||
|
|
||||||
def forward(self, q: Tensor, k: Tensor, v: Tensor, mask: Optional[Tensor] = None, is_causal: bool= False):
|
|
||||||
# (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)
|
|
||||||
# (bsz, n_heads, seq_len, head_dim) - > (bsz, seq_len, n_heads*head_dim)
|
|
||||||
sdqa_out = F.scaled_dot_product_attention(q, k, v, mask, is_causal=is_causal).permute(0, 2, 1, 3).contiguous().flatten(2)
|
|
||||||
|
|
||||||
return sdqa_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.attention = 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,
|
|
||||||
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_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 kv_cache is not None:
|
|
||||||
k_cache, v_cache = kv_cache
|
|
||||||
|
|
||||||
# copy to cache
|
|
||||||
k_cache[:bsz, start_pos:start_pos + seq_len, self.layer_id] = k
|
|
||||||
v_cache[:bsz, start_pos:start_pos + seq_len, self.layer_id] = v
|
|
||||||
|
|
||||||
# get cache
|
|
||||||
k = k_cache[:bsz, :start_pos + seq_len, self.layer_id]
|
|
||||||
v = v_cache[:bsz, :start_pos + seq_len, self.layer_id]
|
|
||||||
|
|
||||||
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)
|
|
||||||
sdqa_out = self.attention(q, k, v, mask, is_causal=(mask == None))
|
|
||||||
|
|
||||||
if self.use_gated_attention:
|
|
||||||
sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
|
|
||||||
|
|
||||||
out = self.o_proj(sdqa_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,
|
|
||||||
kv_cache: Optional[Tuple[Tensor, Tensor]] = None,
|
|
||||||
start_pos: int = 0
|
|
||||||
) -> Tensor:
|
|
||||||
# attention
|
|
||||||
attn_output = self.attention(
|
|
||||||
self.input_norm(x),
|
|
||||||
rotary_emb,
|
|
||||||
attention_mask,
|
|
||||||
kv_cache,
|
|
||||||
start_pos
|
|
||||||
)
|
|
||||||
x = attn_output + x
|
|
||||||
|
|
||||||
# feed forward
|
|
||||||
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)
|
|
||||||
@@ -1,134 +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.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 == 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:
|
|
||||||
if lm_head_key not in state_dict and embed_key in state_dict:
|
|
||||||
# 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,27 +0,0 @@
|
|||||||
from khaosz.parallel.setup import (
|
|
||||||
get_world_size,
|
|
||||||
get_rank,
|
|
||||||
get_current_device,
|
|
||||||
|
|
||||||
only_on_rank,
|
|
||||||
setup_parallel,
|
|
||||||
spawn_parallel_fn
|
|
||||||
)
|
|
||||||
|
|
||||||
from khaosz.parallel.module import (
|
|
||||||
RowParallelLinear,
|
|
||||||
ColumnParallelLinear
|
|
||||||
)
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
"get_world_size",
|
|
||||||
"get_rank",
|
|
||||||
"get_current_device",
|
|
||||||
|
|
||||||
"only_on_rank",
|
|
||||||
"setup_parallel",
|
|
||||||
"spawn_parallel_fn",
|
|
||||||
|
|
||||||
"RowParallelLinear",
|
|
||||||
"ColumnParallelLinear"
|
|
||||||
]
|
|
||||||
@@ -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,
|
|
||||||
MetricLoggerCallback
|
|
||||||
)
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
# trainer
|
|
||||||
"Trainer",
|
|
||||||
|
|
||||||
# factory
|
|
||||||
"StrategyFactory",
|
|
||||||
"SchedulerFactory",
|
|
||||||
|
|
||||||
# callback
|
|
||||||
"TrainCallback",
|
|
||||||
"ProgressBarCallback",
|
|
||||||
"CheckpointCallback",
|
|
||||||
"TrainCallback",
|
|
||||||
"SchedulerCallback",
|
|
||||||
"MetricLoggerCallback"
|
|
||||||
]
|
|
||||||
@@ -1,89 +0,0 @@
|
|||||||
import torch.nn as nn
|
|
||||||
from typing import Dict
|
|
||||||
|
|
||||||
def grad_norm(model: nn.Module, norm_type: int = 2) -> Dict[str, float]:
|
|
||||||
""" Compute gradient norm for each parameter in the model. """
|
|
||||||
norms = {}
|
|
||||||
for name, param in model.named_parameters():
|
|
||||||
norms[name] = 0.0
|
|
||||||
if param.grad:
|
|
||||||
norm = param.grad.data.norm(norm_type).item()
|
|
||||||
norms[name] = norm
|
|
||||||
return norms
|
|
||||||
|
|
||||||
def grad_std(model: nn.Module) -> Dict[str, float]:
|
|
||||||
""" Compute standard deviation of gradients for each parameter. """
|
|
||||||
stds = {}
|
|
||||||
for name, param in model.named_parameters():
|
|
||||||
stds[name] = 0.0
|
|
||||||
if param.grad:
|
|
||||||
std = param.grad.data.std().item()
|
|
||||||
stds[name] = std
|
|
||||||
return stds
|
|
||||||
|
|
||||||
def grad_max(model: nn.Module) -> Dict[str, float]:
|
|
||||||
""" Find the maximum absolute gradient value for each parameter. """
|
|
||||||
max_vals = {}
|
|
||||||
for name, param in model.named_parameters():
|
|
||||||
max_vals[name] = -float('inf')
|
|
||||||
if param.grad:
|
|
||||||
max_val = param.grad.data.max().item()
|
|
||||||
max_vals[name] = max_val
|
|
||||||
|
|
||||||
return max_vals
|
|
||||||
|
|
||||||
def grad_min(model: nn.Module) -> Dict[str, float]:
|
|
||||||
""" Find the minimum absolute gradient value for each parameter. """
|
|
||||||
min_vals = {}
|
|
||||||
for name, param in model.named_parameters():
|
|
||||||
min_vals[name] = float('inf')
|
|
||||||
if param.grad:
|
|
||||||
min_val = param.grad.data.min().item()
|
|
||||||
min_vals[name] = min_val
|
|
||||||
|
|
||||||
return min_vals
|
|
||||||
|
|
||||||
def grad_mean(model: nn.Module) -> Dict[str, float]:
|
|
||||||
""" Compute mean of gradients for each parameter. """
|
|
||||||
means = {}
|
|
||||||
for name, param in model.named_parameters():
|
|
||||||
means[name] = 0.0
|
|
||||||
if param.grad:
|
|
||||||
mean = param.grad.data.mean().item()
|
|
||||||
means[name] = mean
|
|
||||||
|
|
||||||
return means
|
|
||||||
|
|
||||||
def grad_nan_num(model: nn.Module) -> Dict[str, int]:
|
|
||||||
""" Count the number of NaNs in gradients for each parameter. """
|
|
||||||
nan_nums = {}
|
|
||||||
for name, param in model.named_parameters():
|
|
||||||
nan_nums[name] = 0
|
|
||||||
if param.grad:
|
|
||||||
nan_num = param.grad.isnan().sum().item()
|
|
||||||
nan_nums[name] = nan_num
|
|
||||||
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)
|
|
||||||
@@ -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(optimizer, schedule_config: ScheduleConfig) -> BaseScheduler:
|
|
||||||
kwargs = schedule_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,137 +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, Callable, Dict, Union
|
|
||||||
from abc import ABC, abstractmethod
|
|
||||||
|
|
||||||
|
|
||||||
def get_logprobs(
|
|
||||||
model: Union[nn.Module, Callable[..., Dict[str, Tensor]]],
|
|
||||||
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.float(), 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: Dict[str, 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).float(),
|
|
||||||
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).float(),
|
|
||||||
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: Dict[str, 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 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,99 +0,0 @@
|
|||||||
import torch.nn as nn
|
|
||||||
from torch.optim import Optimizer
|
|
||||||
from torch.optim.lr_scheduler import LRScheduler
|
|
||||||
from torch.utils.data import DataLoader
|
|
||||||
|
|
||||||
from khaosz.data import ResumableDistributedSampler
|
|
||||||
from khaosz.data.checkpoint import Checkpoint
|
|
||||||
from khaosz.trainer.strategy import StrategyFactory, BaseStrategy
|
|
||||||
from khaosz.config.train_config import TrainConfig
|
|
||||||
from khaosz.parallel.setup import get_current_device, get_world_size, get_rank
|
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
|
||||||
from typing import Optional, Self
|
|
||||||
|
|
||||||
|
|
||||||
@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):
|
|
||||||
self.config = config
|
|
||||||
self._context = TrainContext(
|
|
||||||
model=config.model,
|
|
||||||
world_size=get_world_size(),
|
|
||||||
rank=get_rank(),
|
|
||||||
)
|
|
||||||
|
|
||||||
device = get_current_device()
|
|
||||||
self._context.model = self._context.model.to(device=device)
|
|
||||||
|
|
||||||
if self.config.nprocs > 1:
|
|
||||||
fn = self.config.parallel_wrapper
|
|
||||||
self._context.model = fn(self._context.model)
|
|
||||||
|
|
||||||
self._context.optimizer = self.config.optimizer_fn(self._context.model)
|
|
||||||
self._context.scheduler = self.config.scheduler_fn(self._context.optimizer)
|
|
||||||
|
|
||||||
def with_checkpoint(self, checkpoint: Optional[Checkpoint]) -> Self:
|
|
||||||
if checkpoint is None:
|
|
||||||
checkpoint = Checkpoint(
|
|
||||||
state_dict=self._context.model.state_dict(),
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
# resume from the assigned checkpoint or assigned iteration
|
|
||||||
self._context.epoch = max(checkpoint.epoch, self.config.start_epoch)
|
|
||||||
self._context.iteration = max(checkpoint.iteration, self.config.start_batch)
|
|
||||||
self._context.model.load_state_dict(checkpoint.state_dict)
|
|
||||||
|
|
||||||
self._context.checkpoint = checkpoint
|
|
||||||
return self
|
|
||||||
|
|
||||||
def with_dataloader(self) -> Self:
|
|
||||||
# fix: change batch level iteration to sample level offset
|
|
||||||
config = self.config
|
|
||||||
sampler_offset = self._context.iteration * config.batch_size
|
|
||||||
resumeable_sampler = ResumableDistributedSampler(
|
|
||||||
data_source=config.dataset,
|
|
||||||
start_epoch=self._context.epoch,
|
|
||||||
start_iter=sampler_offset,
|
|
||||||
seed=config.random_seed
|
|
||||||
)
|
|
||||||
|
|
||||||
dataloader = DataLoader(
|
|
||||||
config.dataset,
|
|
||||||
batch_size=config.batch_size,
|
|
||||||
sampler=resumeable_sampler,
|
|
||||||
num_workers=config.num_workers,
|
|
||||||
pin_memory=config.pin_memory,
|
|
||||||
prefetch_factor=config.prefetch_factor
|
|
||||||
)
|
|
||||||
self._context.dataloader = dataloader
|
|
||||||
return self
|
|
||||||
|
|
||||||
def with_strategy(self) -> Self:
|
|
||||||
self._context.strategy = StrategyFactory.load(
|
|
||||||
model=self.config.model,
|
|
||||||
train_type=self.config.strategy,
|
|
||||||
device=get_current_device(),
|
|
||||||
**self.config.extra_kwargs
|
|
||||||
)
|
|
||||||
return self
|
|
||||||
|
|
||||||
def build(self) -> TrainContext:
|
|
||||||
return self._context
|
|
||||||
@@ -1,102 +0,0 @@
|
|||||||
import logging
|
|
||||||
from typing import Optional, List
|
|
||||||
from khaosz.config import TrainConfig
|
|
||||||
from khaosz.trainer.train_callback import (
|
|
||||||
TrainCallback,
|
|
||||||
ProgressBarCallback,
|
|
||||||
CheckpointCallback,
|
|
||||||
MetricLoggerCallback,
|
|
||||||
GradientClippingCallback,
|
|
||||||
SchedulerCallback
|
|
||||||
)
|
|
||||||
from khaosz.trainer.train_context import TrainContext, TrainContextBuilder
|
|
||||||
from khaosz.data.checkpoint import Checkpoint
|
|
||||||
from khaosz.parallel.setup import spawn_parallel_fn
|
|
||||||
|
|
||||||
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]:
|
|
||||||
train_config = self.train_config
|
|
||||||
return [
|
|
||||||
ProgressBarCallback(train_config.n_epoch),
|
|
||||||
CheckpointCallback(train_config.checkpoint_dir, train_config.checkpoint_interval),
|
|
||||||
MetricLoggerCallback(train_config.checkpoint_dir, train_config.checkpoint_interval),
|
|
||||||
GradientClippingCallback(train_config.max_grad_norm),
|
|
||||||
SchedulerCallback(),
|
|
||||||
]
|
|
||||||
|
|
||||||
def _build_context(self, checkpoint: Optional[Checkpoint]) -> TrainContext:
|
|
||||||
return (TrainContextBuilder(self.train_config)
|
|
||||||
.with_checkpoint(checkpoint)
|
|
||||||
.with_dataloader()
|
|
||||||
.with_strategy()
|
|
||||||
.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,
|
|
||||||
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)
|
|
||||||
|
|
||||||
for batch in context.dataloader:
|
|
||||||
if context.iteration % self.train_config.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_batch = self.train_config.accumulation_steps * self.train_config.nprocs
|
|
||||||
stand_loss = loss / stand_batch
|
|
||||||
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)
|
|
||||||
@@ -1 +0,0 @@
|
|||||||
# init file
|
|
||||||
@@ -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)
|
|
||||||
+20
-4
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|||||||
|
|
||||||
[project]
|
[project]
|
||||||
dynamic = ["version"]
|
dynamic = ["version"]
|
||||||
name = "khaosz"
|
name = "astrai"
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
requires-python = ">=3.12"
|
requires-python = ">=3.12"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
@@ -15,7 +15,11 @@ dependencies = [
|
|||||||
"tqdm==4.67.1",
|
"tqdm==4.67.1",
|
||||||
"safetensors==0.5.3",
|
"safetensors==0.5.3",
|
||||||
"huggingface-hub==0.34.3",
|
"huggingface-hub==0.34.3",
|
||||||
"pytest==9.0.2"
|
"jinja2>=3.0.0",
|
||||||
|
"fastapi",
|
||||||
|
"uvicorn[standard]",
|
||||||
|
"httpx",
|
||||||
|
"requests",
|
||||||
]
|
]
|
||||||
keywords = ["nlp", "datasets", "language-models", "machine-learning"]
|
keywords = ["nlp", "datasets", "language-models", "machine-learning"]
|
||||||
license = { text = "GPL-3.0" }
|
license = { text = "GPL-3.0" }
|
||||||
@@ -24,7 +28,10 @@ classifiers = [
|
|||||||
"License :: OSI Approved :: GPL-3.0",
|
"License :: OSI Approved :: GPL-3.0",
|
||||||
"Operating System :: OS Independent",
|
"Operating System :: OS Independent",
|
||||||
]
|
]
|
||||||
urls = { Homepage = "https://github.com/khaosz/khaosz" }
|
urls = { Homepage = "https://github.com/ViperEkura/AstrAI" }
|
||||||
|
|
||||||
|
[project.optional-dependencies]
|
||||||
|
dev = ["pytest==9.0.2", "ruff"]
|
||||||
|
|
||||||
[tool.setuptools.packages.find]
|
[tool.setuptools.packages.find]
|
||||||
where = ["."]
|
where = ["."]
|
||||||
@@ -33,4 +40,13 @@ where = ["."]
|
|||||||
extra-index-url = "https://download.pytorch.org/whl/cu126"
|
extra-index-url = "https://download.pytorch.org/whl/cu126"
|
||||||
|
|
||||||
[tool.setuptools.dynamic]
|
[tool.setuptools.dynamic]
|
||||||
version = { attr = "khaosz.__version__" }
|
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"
|
||||||
@@ -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,56 @@
|
|||||||
|
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 = [
|
||||||
|
"你好",
|
||||||
|
"请问什么是人工智能",
|
||||||
|
"今天天气如何",
|
||||||
|
"我感到焦虑, 请问我应该怎么办",
|
||||||
|
"请问什么是显卡",
|
||||||
|
]
|
||||||
|
|
||||||
|
prompts = [
|
||||||
|
tokenizer.apply_chat_template(
|
||||||
|
[
|
||||||
|
{"role": "system", "content": "You are a helpful assistant."},
|
||||||
|
{"role": "user", "content": q},
|
||||||
|
],
|
||||||
|
tokenize=False,
|
||||||
|
)
|
||||||
|
for q in inputs
|
||||||
|
]
|
||||||
|
|
||||||
|
engine = InferenceEngine(
|
||||||
|
model=model,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
)
|
||||||
|
responses = engine.generate(
|
||||||
|
prompt=prompts,
|
||||||
|
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()
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user