Compare commits
74
Commits
bc7c82977e
..
v1.3.6
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
785d65436c | ||
|
|
64be81b7b3 | ||
|
|
45479b5731 | ||
|
|
e0a3337c22 | ||
|
|
812238060b | ||
|
|
14b0d56197 | ||
|
|
6c8533f1d2 | ||
|
|
2c2697390d | ||
|
|
7621f05d3f | ||
|
|
10ebd7211f | ||
|
|
42a391f0fb | ||
|
|
97c7ac0f4f | ||
|
|
8f1b32f2b6 | ||
|
|
c241a5dcef | ||
|
|
44dab27fdc | ||
|
|
a44fd22a99 | ||
|
|
8a11a7d444 | ||
|
|
1d54491809 | ||
|
|
ad9f4d9cf6 | ||
|
|
e1638a7ade | ||
|
|
f91bfee33e | ||
|
|
d7a7f570ed | ||
|
|
7dea929788 | ||
|
|
026d1fc33d | ||
|
|
7242eedbf4 | ||
|
|
04c0dc7a47 | ||
|
|
48a53121ba | ||
|
|
0ba8c70ce1 | ||
|
|
3d12a03909 | ||
|
|
c169659611 | ||
|
|
e12f1a7ee5 | ||
|
|
ef25efffa2 | ||
|
|
19532440b4 | ||
|
|
9096e413c3 | ||
|
|
9d5e9fa6c4 | ||
|
|
08dde46778 | ||
|
|
513f1f7826 | ||
|
|
e3382f6bb5 | ||
|
|
f0339022c1 | ||
|
|
d8da2cf17c | ||
|
|
205b40bd28 | ||
|
|
18fe6e9339 | ||
|
|
2196c34c52 | ||
|
|
466c2e1efd | ||
|
|
7e26d848ab | ||
|
|
ed95ef245c | ||
|
|
6d6ef99e66 | ||
|
|
a8e2a1ba45 | ||
|
|
6269bacfc3 | ||
|
|
c0effc9f5b | ||
|
|
df0845e916 | ||
|
|
7440e9c809 | ||
|
|
7d4029c2a4 | ||
|
|
0ca6c9e6eb | ||
|
|
6e49d27057 | ||
|
|
5203b7f53e | ||
|
|
5889179c54 | ||
|
|
38e18fdfd3 | ||
|
|
4753958f92 | ||
|
|
73d6cc0f26 | ||
|
|
317ed90bac | ||
|
|
951df8155c | ||
|
|
a58fab8d6e | ||
|
|
a3c8296135 | ||
|
|
c95ace41aa | ||
|
|
3da428e0e4 | ||
|
|
133a9de98f | ||
|
|
523eacf5fe | ||
|
|
cffedaad5e | ||
|
|
3583c46b66 | ||
|
|
ca4e6b907c | ||
|
|
db99d8b254 | ||
|
|
b98c9cefdc | ||
|
|
283bcaf2ff |
@@ -2,7 +2,7 @@
|
|||||||
name: Bug report
|
name: Bug report
|
||||||
about: Create a report to help us improve
|
about: Create a report to help us improve
|
||||||
title: "[BUG]"
|
title: "[BUG]"
|
||||||
labels: enhancement
|
labels: bug
|
||||||
assignees: ''
|
assignees: ''
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|||||||
@@ -16,9 +16,9 @@ Please delete options that are not relevant.
|
|||||||
Please describe the tests that you ran to verify your changes. Provide instructions so we can reproduce.
|
Please describe the tests that you ran to verify your changes. Provide instructions so we can reproduce.
|
||||||
|
|
||||||
## Checklist:
|
## Checklist:
|
||||||
- [ ] My code follows the style guidelines of this project (run `ruff format .` and `ruff check --fix .`)
|
- [ ] My code follows the style guidelines of this project (run `ruff format .` and `ruff check . --select I`)
|
||||||
- [ ] I have performed a self-review of my own code
|
- [ ] I have performed a self-review of my own code
|
||||||
- [ ] I have commented my code, particularly in hard-to-understand areas
|
- [ ] Code is self-documenting (no unnecessary comments)
|
||||||
- [ ] I have made corresponding changes to the documentation
|
- [ ] I have made corresponding changes to the documentation
|
||||||
- [ ] My changes generate no new warnings
|
- [ ] My changes generate no new warnings
|
||||||
- [ ] I have added tests that prove my fix is effective or that my feature works
|
- [ ] I have added tests that prove my fix is effective or that my feature works
|
||||||
|
|||||||
+80
-48
@@ -1,68 +1,100 @@
|
|||||||
# Contributing to AstrAI
|
# Contributing to AstrAI
|
||||||
|
|
||||||
Thank you for your interest in contributing to AstrAI! This document provides guidelines and steps for contributing.
|
Thank you for your interest in contributing! This document provides step-by-step guidelines.
|
||||||
|
|
||||||
## How to Contribute
|
## Quick Start
|
||||||
|
|
||||||
### Reporting Issues
|
```bash
|
||||||
If you encounter a bug or have a feature request, please open an issue on GitHub. Include as much detail as possible:
|
git clone https://github.com/your-username/AstrAI.git
|
||||||
- A clear description of the problem or request.
|
cd AstrAI
|
||||||
- Steps to reproduce (for bugs).
|
pip install -e ".[dev]" # install with dev dependencies (pytest, ruff)
|
||||||
- Your environment (Python version, OS, etc.).
|
```
|
||||||
|
|
||||||
### Submitting Changes
|
## Before You Commit
|
||||||
1. **Fork** the repository.
|
|
||||||
2. **Clone** your fork:
|
|
||||||
```bash
|
|
||||||
git clone https://github.com/your-username/AstrAI.git
|
|
||||||
cd AstrAI
|
|
||||||
```
|
|
||||||
3. **Create a feature branch**:
|
|
||||||
```bash
|
|
||||||
git checkout -b feature/your-feature-name
|
|
||||||
```
|
|
||||||
4. **Make your changes**. Follow the code style guidelines below.
|
|
||||||
5. **Commit your changes** with a descriptive commit message:
|
|
||||||
```bash
|
|
||||||
git commit -m "Add: brief description of the change"
|
|
||||||
```
|
|
||||||
6. **Push** to your fork:
|
|
||||||
```bash
|
|
||||||
git push origin feature/your-feature-name
|
|
||||||
```
|
|
||||||
7. **Open a Pull Request** (PR) against the `main` branch of the upstream repository.
|
|
||||||
|
|
||||||
## Code Style
|
Run the following checks **in order** — CI will reject if any fail.
|
||||||
|
|
||||||
AstrAI uses [Ruff](https://docs.astral.sh/ruff/) for code formatting and linting. Please ensure your code is formatted before submitting.
|
### 1. Format
|
||||||
|
|
||||||
- Run Ruff to format and lint:
|
```bash
|
||||||
```bash
|
ruff format .
|
||||||
ruff format .
|
```
|
||||||
ruff check --fix .
|
|
||||||
```
|
|
||||||
- The project uses **double quotes** for strings and **4‑space indentation** (as configured in `pyproject.toml`).
|
|
||||||
|
|
||||||
## Testing
|
> **Note**: `ruff format` may rename parameters (e.g. `mask` → `attn_mask`).
|
||||||
|
> Always review the diff after formatting.
|
||||||
|
|
||||||
If you add or modify functionality, please include appropriate tests.
|
### 2. Import sorting
|
||||||
|
|
||||||
- Run the test suite with:
|
```bash
|
||||||
```bash
|
ruff check . --select I
|
||||||
pytest
|
```
|
||||||
```
|
|
||||||
- Ensure all tests pass before submitting your PR.
|
If this fails, **manually fix** import ordering (ruff does not auto-fix in this project's CI):
|
||||||
|
|
||||||
|
```bash
|
||||||
|
ruff check . --select I --fix .
|
||||||
|
ruff format . # re-format after fix
|
||||||
|
```
|
||||||
|
|
||||||
|
### 3. Run tests
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python -u -m pytest tests/ -v
|
||||||
|
```
|
||||||
|
|
||||||
|
> Failed tests may leave orphan tempdirs under `%TEMP%`. Clean them manually if needed.
|
||||||
|
|
||||||
|
### 4. (Optional) Full pre-commit check
|
||||||
|
|
||||||
|
If you have Git Bash available:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
bash scripts/pre_commit.sh
|
||||||
|
```
|
||||||
|
|
||||||
|
This runs format check, import sort check, and tests in one go.
|
||||||
|
|
||||||
|
## Commit Style
|
||||||
|
|
||||||
|
```
|
||||||
|
fix/feat/chore/docs/refactor/perf/test/style/ci/build/revert : short description (~50 chars)
|
||||||
|
|
||||||
|
- bullet point body (each ~60 chars)
|
||||||
|
```
|
||||||
|
|
||||||
|
- **Type** must be one of: `fix`, `feat`, `chore`, `docs`, `refactor`, `perf`, `test`, `style`, `ci`, `build`, `revert`.
|
||||||
|
- **Subject line** ends with no period.
|
||||||
|
- **Body** uses bullet points starting with `-`.
|
||||||
|
- No `(scope)` parentheses.
|
||||||
|
|
||||||
|
## Common Issues
|
||||||
|
|
||||||
|
| Problem | Cause | Fix |
|
||||||
|
|---------|-------|-----|
|
||||||
|
| `ruff check --select I` fails | Wrong import order | `ruff check . --select I --fix .` then `ruff format .` |
|
||||||
|
| `ruff format` changed many files | Not formatted before commit | Review diff carefully before staging |
|
||||||
|
| Pre-commit hook rejects | Tests or lint failed | Fix individually, do not `--no-verify` |
|
||||||
|
| Tests fail with tempdir left | Test crash | Clean `%TEMP%` manually |
|
||||||
|
|
||||||
|
## Submitting Changes
|
||||||
|
|
||||||
|
1. Fork the repo.
|
||||||
|
2. Create a feature branch: `git checkout -b feat/my-feature`
|
||||||
|
3. Make changes following the steps above.
|
||||||
|
4. Commit with the commit style above.
|
||||||
|
5. Push: `git push origin feat/my-feature`
|
||||||
|
6. Open a Pull Request against `main`.
|
||||||
|
|
||||||
## Code Review
|
## Code Review
|
||||||
|
|
||||||
All submissions will be reviewed. We may request changes or discuss alternatives. Please be responsive to feedback.
|
- All PRs are reviewed. We may request changes.
|
||||||
|
- CI runs `ruff format --check .` then `ruff check . --select I` (no `--fix` in CI).
|
||||||
|
- Ensure all tests pass.
|
||||||
|
|
||||||
## License
|
## License
|
||||||
|
|
||||||
By contributing, you agree that your contributions will be licensed under the same [GPL-3.0 License](LICENSE) that covers the project.
|
By contributing, you agree that your contributions will be licensed under the [GPL-3.0 License](LICENSE).
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
If you have any questions, feel free to ask in the [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions) or open an issue.
|
Questions? Ask in [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions) or open an issue.
|
||||||
|
|
||||||
Happy contributing!
|
|
||||||
|
|||||||
+5
-4
@@ -1,7 +1,7 @@
|
|||||||
# AstrAI Dockerfile - Multi-stage Build (Optimized)
|
# AstrAI Dockerfile - Multi-stage Build (Optimized)
|
||||||
|
|
||||||
# Build stage - use base image with minimal build tools
|
# Build stage - use base image with minimal build tools
|
||||||
FROM nvidia/cuda:12.6.0-base-ubuntu24.04 AS builder
|
FROM ubuntu:24.04 AS builder
|
||||||
|
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
@@ -18,7 +18,7 @@ RUN apt-get update && DEBIAN_FRONTEND=noninteractive apt-get install -y --no-ins
|
|||||||
RUN python3.12 -m venv --copies /opt/venv
|
RUN python3.12 -m venv --copies /opt/venv
|
||||||
ENV PATH="/opt/venv/bin:$PATH"
|
ENV PATH="/opt/venv/bin:$PATH"
|
||||||
|
|
||||||
# Copy source code and install dependencies
|
# Copy source code and install (deps read from pyproject.toml)
|
||||||
COPY astrai/ ./astrai/
|
COPY astrai/ ./astrai/
|
||||||
COPY pyproject.toml .
|
COPY pyproject.toml .
|
||||||
RUN pip install --no-cache-dir --upgrade pip \
|
RUN pip install --no-cache-dir --upgrade pip \
|
||||||
@@ -26,13 +26,14 @@ RUN pip install --no-cache-dir --upgrade pip \
|
|||||||
--extra-index-url https://download.pytorch.org/whl/cu126
|
--extra-index-url https://download.pytorch.org/whl/cu126
|
||||||
|
|
||||||
# Production stage
|
# Production stage
|
||||||
FROM nvidia/cuda:12.6.0-base-ubuntu24.04 AS production
|
FROM ubuntu:24.04 AS production
|
||||||
|
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
# Install Python 3.12 runtime
|
# Install Python 3.12 runtime and healthcheck dependency
|
||||||
RUN apt-get update && DEBIAN_FRONTEND=noninteractive apt-get install -y --no-install-recommends \
|
RUN apt-get update && DEBIAN_FRONTEND=noninteractive apt-get install -y --no-install-recommends \
|
||||||
python3.12 \
|
python3.12 \
|
||||||
|
curl \
|
||||||
&& rm -rf /var/lib/apt/lists/*
|
&& rm -rf /var/lib/apt/lists/*
|
||||||
|
|
||||||
# Copy virtual environment from builder
|
# Copy virtual environment from builder
|
||||||
|
|||||||
@@ -46,7 +46,7 @@
|
|||||||
- 💡 **Easy to Use**: Simple API with comprehensive examples and demos.
|
- 💡 **Easy to Use**: Simple API with comprehensive examples and demos.
|
||||||
- 📦 **Lightweight**: Minimal dependencies, easy to deploy.
|
- 📦 **Lightweight**: Minimal dependencies, easy to deploy.
|
||||||
- 🔬 **Research‑Friendly**: Modular design, easy to experiment with new ideas.
|
- 🔬 **Research‑Friendly**: Modular design, easy to experiment with new ideas.
|
||||||
- 🤗 **HuggingFace Integration**: Compatible with HuggingFace models and datasets.
|
- 🤗 **HuggingFace-Style API**: AutoModel/AutoTokenizer APIs inspired by HuggingFace for easy model and tokenizer loading.
|
||||||
- 🔌 **Dual API Compatibility**: Supports both OpenAI and Anthropic chat completion APIs out of the box.
|
- 🔌 **Dual API Compatibility**: Supports both OpenAI and Anthropic chat completion APIs out of the box.
|
||||||
|
|
||||||
### Quick Start
|
### Quick Start
|
||||||
@@ -65,49 +65,51 @@ For development dependencies:
|
|||||||
pip install -e ".[dev]"
|
pip install -e ".[dev]"
|
||||||
```
|
```
|
||||||
|
|
||||||
|
#### Download Pre-trained Model
|
||||||
|
|
||||||
|
Download pre-trained model weights (1B bilingual checkpoint) to `params/`:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/demo/download.py
|
||||||
|
```
|
||||||
|
|
||||||
|
Or download manually from [HuggingFace](https://huggingface.co/ViperEk/KHAOSZ) into `params/`.
|
||||||
|
|
||||||
#### Train a Model
|
#### Train a Model
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python scripts/tools/train.py --train_type=seq --data_root_path=/path/to/dataset --param_path=/path/to/model
|
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||||
|
|
||||||
|
nohup python scripts/tools/train.py \
|
||||||
|
--nprocs=4 \
|
||||||
|
--train_type=seq \
|
||||||
|
--data_root_path=/path/to/dataset \
|
||||||
|
--param_path=/path/to/model \
|
||||||
|
--batch_per_device=4 \
|
||||||
|
--grad_accum_steps=8 \
|
||||||
|
--warmup_ratio=0.05 \
|
||||||
|
--max_lr=1e-4 \
|
||||||
|
--max_grad_norm=1.0 \
|
||||||
|
--adamw_beta1=0.9 \
|
||||||
|
--adamw_beta2=0.95 \
|
||||||
|
--adamw_weight_decay=0.01 \
|
||||||
|
--window_size=2048 \
|
||||||
|
--ckpt_interval=10000 \
|
||||||
|
--ckpt_dir=./checkpoint \
|
||||||
|
--random_seed=3407 \
|
||||||
|
--label_smoothing=0.05 \
|
||||||
|
> out.log 2> err.log &
|
||||||
```
|
```
|
||||||
|
|
||||||
| Parameter | Description | Default |
|
Full reference at [Parameter Guide](assets/docs/params.md).
|
||||||
|-----------|-------------|---------|
|
|
||||||
| `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`) | required |
|
|
||||||
| `--data_root_path` | Dataset root directory | required |
|
|
||||||
| `--param_path` | Model / checkpoint path | required |
|
|
||||||
| `--n_epoch` | Training epochs | 1 |
|
|
||||||
| `--batch_size` | Batch size | 1 |
|
|
||||||
| `--accumulation_steps` | Gradient accumulation steps | 1 |
|
|
||||||
| `--warmup_steps` | LR warmup steps | 1000 |
|
|
||||||
| `--max_lr` | Peak learning rate (cosine decay) | 3e-4 |
|
|
||||||
| `--max_grad_norm` | Max gradient norm for clipping | 1.0 |
|
|
||||||
| `--adamw_beta1` | AdamW beta1 | 0.9 |
|
|
||||||
| `--adamw_beta2` | AdamW beta2 | 0.95 |
|
|
||||||
| `--adamw_weight_decay` | AdamW weight decay | 0.01 |
|
|
||||||
| `--random_seed` | Random seed | 3407 |
|
|
||||||
| `--num_workers` | DataLoader workers | 4 |
|
|
||||||
| `--window_size` | Max input sequence length | auto |
|
|
||||||
| `--stride` | Sequence stride | auto |
|
|
||||||
| `--label_smoothing` | Label smoothing for cross entropy | 0.1 |
|
|
||||||
| `--dpo_beta` | DPO beta | 0.1 |
|
|
||||||
| `--grpo_clip_eps` | GRPO clip epsilon | 0.2 |
|
|
||||||
| `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 |
|
|
||||||
| `--group_size` | GRPO group size | 4 |
|
|
||||||
| `--grpo_sync_interval` | GRPO ref model sync interval (steps) | 200 |
|
|
||||||
| `--ckpt_interval` | Checkpoint interval (iters) | 5000 |
|
|
||||||
| `--ckpt_dir` | Checkpoint directory | checkpoint |
|
|
||||||
| `--start_epoch` | Start epoch (for resume) | 0 |
|
|
||||||
| `--start_batch` | Start batch (for resume) | 0 |
|
|
||||||
| `--nprocs` | Number of GPUs | 1 |
|
|
||||||
| `--device_type` | Device type | cuda |
|
|
||||||
|
|
||||||
Full reference at [Parameter Guide](./assets/docs/params.md#training-parameters).
|
|
||||||
|
|
||||||
#### Generate Text
|
#### Generate Text
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python scripts/tools/generate.py --param_path=/path/to/param_path
|
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
|
||||||
@@ -140,8 +142,6 @@ docker compose --profile cpu up -d
|
|||||||
|
|
||||||
> **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`.
|
> **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`.
|
||||||
|
|
||||||
> **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`.
|
|
||||||
|
|
||||||
#### Start HTTP Server
|
#### Start HTTP Server
|
||||||
|
|
||||||
Start the inference server with OpenAI and Anthropic-compatible HTTP API:
|
Start the inference server with OpenAI and Anthropic-compatible HTTP API:
|
||||||
@@ -213,16 +213,17 @@ python scripts/demo/generate_batch.py
|
|||||||
python scripts/demo/generate_ar.py
|
python scripts/demo/generate_ar.py
|
||||||
```
|
```
|
||||||
|
|
||||||
Watch a video walkthrough on [bilibili](https://www.bilibili.com/video/BV1z5RPYHEkd).
|
Watch a video walkthrough on [bilibili](https://www.bilibili.com/video/BV1fuLB6yEj6).
|
||||||
|
|
||||||
### Documentation
|
### Documentation
|
||||||
|
|
||||||
| Document | Description |
|
| Document | Description |
|
||||||
|----------|-------------|
|
|----------|-------------|
|
||||||
| [Parameter Guide](./assets/docs/params.md) | Training & inference parameters |
|
| [Parameter Guide](./assets/docs/params.md) | Training & inference parameters |
|
||||||
| [Design Document](./assets/docs/design.md) | Framework architecture & module design |
|
| [Architecture](./assets/docs/architecture.md) | System architecture, class diagram & design patterns |
|
||||||
| [Data Flow](./assets/docs/dataflow.md) | Data processing pipeline details |
|
| [Training](./assets/docs/training.md) | Training loop, strategies & formulas |
|
||||||
| [Model Introduction](./assets/docs/introduction.md) | Model architecture & technical details |
|
| [Inference](./assets/docs/inference.md) | KVCache, continuous batching, sampling & HTTP API |
|
||||||
|
| [Data Flow](./assets/docs/dataflow.md) | Data pipeline, storage backends & dataset architecture |
|
||||||
|
|
||||||
### Contributing
|
### Contributing
|
||||||
|
|
||||||
|
|||||||
+42
-39
@@ -52,7 +52,7 @@
|
|||||||
- 💡 **易用**: 简洁的 API 与丰富的示例、演示。
|
- 💡 **易用**: 简洁的 API 与丰富的示例、演示。
|
||||||
- 📦 **轻量**: 依赖少,部署简单。
|
- 📦 **轻量**: 依赖少,部署简单。
|
||||||
- 🔬 **研究友好**: 模块化设计,便于实验新想法。
|
- 🔬 **研究友好**: 模块化设计,便于实验新想法。
|
||||||
- 🤗 **HuggingFace 集成**: 兼容 HuggingFace 模型与数据集。
|
- 🤗 **HuggingFace 风格 API**: 类 HuggingFace 的 AutoModel/AutoTokenizer 接口,方便加载模型和分词器。
|
||||||
- 🔌 **双 API 兼容**: 同时支持 OpenAI 和 Anthropic 聊天补全 API,开箱即用。
|
- 🔌 **双 API 兼容**: 同时支持 OpenAI 和 Anthropic 聊天补全 API,开箱即用。
|
||||||
|
|
||||||
### 快速开始
|
### 快速开始
|
||||||
@@ -71,49 +71,51 @@ pip install -e .
|
|||||||
pip install -e ".[dev]"
|
pip install -e ".[dev]"
|
||||||
```
|
```
|
||||||
|
|
||||||
|
#### 下载预训练模型
|
||||||
|
|
||||||
|
下载预训练模型权重(1B 双语检查点)到 `params/` 目录:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/demo/download.py
|
||||||
|
```
|
||||||
|
|
||||||
|
或从 [HuggingFace](https://huggingface.co/ViperEk/KHAOSZ) 手动下载放入 `params/`。
|
||||||
|
|
||||||
#### 训练模型
|
#### 训练模型
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python scripts/tools/train.py --train_type=seq --data_root_path=/path/to/dataset --param_path=/path/to/model
|
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||||
|
|
||||||
|
nohup python scripts/tools/train.py \
|
||||||
|
--nprocs=4 \
|
||||||
|
--train_type=seq \
|
||||||
|
--data_root_path=/path/to/dataset \
|
||||||
|
--param_path=/path/to/model \
|
||||||
|
--batch_per_device=4 \
|
||||||
|
--grad_accum_steps=8 \
|
||||||
|
--warmup_ratio=0.05 \
|
||||||
|
--max_lr=1e-4 \
|
||||||
|
--max_grad_norm=1.0 \
|
||||||
|
--adamw_beta1=0.9 \
|
||||||
|
--adamw_beta2=0.95 \
|
||||||
|
--adamw_weight_decay=0.01 \
|
||||||
|
--window_size=2048 \
|
||||||
|
--ckpt_interval=10000 \
|
||||||
|
--ckpt_dir=./checkpoint \
|
||||||
|
--random_seed=3407 \
|
||||||
|
--label_smoothing=0.05 \
|
||||||
|
> out.log 2> err.log &
|
||||||
```
|
```
|
||||||
|
|
||||||
| 参数 | 说明 | 默认值 |
|
完整参数列表见[参数说明](./params.md)。
|
||||||
|------|------|--------|
|
|
||||||
| `--train_type` | 训练类型(`seq`, `sft`, `dpo`, `grpo`) | 必填 |
|
|
||||||
| `--data_root_path` | 数据集根目录 | 必填 |
|
|
||||||
| `--param_path` | 模型参数或断点路径 | 必填 |
|
|
||||||
| `--n_epoch` | 训练轮数 | 1 |
|
|
||||||
| `--batch_size` | 批次大小 | 1 |
|
|
||||||
| `--accumulation_steps` | 梯度累积步数 | 1 |
|
|
||||||
| `--warmup_steps` | 预热步数 | 1000 |
|
|
||||||
| `--max_lr` | 峰值学习率(余弦衰减) | 3e-4 |
|
|
||||||
| `--max_grad_norm` | 梯度裁剪最大值 | 1.0 |
|
|
||||||
| `--adamw_beta1` | AdamW beta1 | 0.9 |
|
|
||||||
| `--adamw_beta2` | AdamW beta2 | 0.95 |
|
|
||||||
| `--adamw_weight_decay` | AdamW 权重衰减 | 0.01 |
|
|
||||||
| `--random_seed` | 随机种子 | 3407 |
|
|
||||||
| `--num_workers` | 数据加载线程数 | 4 |
|
|
||||||
| `--window_size` | 最大输入序列长度 | auto |
|
|
||||||
| `--stride` | 序列步长 | auto |
|
|
||||||
| `--label_smoothing` | 交叉熵标签平滑 | 0.1 |
|
|
||||||
| `--dpo_beta` | DPO beta | 0.1 |
|
|
||||||
| `--grpo_clip_eps` | GRPO 裁剪 epsilon | 0.2 |
|
|
||||||
| `--grpo_kl_coef` | GRPO KL 惩罚系数 | 0.01 |
|
|
||||||
| `--group_size` | GRPO 组大小 | 4 |
|
|
||||||
| `--grpo_sync_interval` | GRPO ref_model 同步间隔(步) | 200 |
|
|
||||||
| `--ckpt_interval` | 检查点间隔(迭代步) | 5000 |
|
|
||||||
| `--ckpt_dir` | 检查点保存目录 | checkpoint |
|
|
||||||
| `--start_epoch` | 起始轮次(用于断点续训) | 0 |
|
|
||||||
| `--start_batch` | 起始批次(用于断点续训) | 0 |
|
|
||||||
| `--nprocs` | GPU 数量 | 1 |
|
|
||||||
| `--device_type` | 设备类型 | cuda |
|
|
||||||
|
|
||||||
完整参数列表见[参数说明](./params.md#training-parameters)。
|
|
||||||
|
|
||||||
#### 文本生成
|
#### 文本生成
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python scripts/tools/generate.py --param_path=/path/to/param_path
|
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
|
||||||
@@ -217,16 +219,17 @@ python scripts/demo/generate_batch.py
|
|||||||
python scripts/demo/generate_ar.py
|
python scripts/demo/generate_ar.py
|
||||||
```
|
```
|
||||||
|
|
||||||
观看 [bilibili](https://www.bilibili.com/video/BV1z5RPYHEkd) 上的视频演示。
|
观看 [bilibili](https://www.bilibili.com/video/BV1fuLB6yEj6) 上的视频演示。
|
||||||
|
|
||||||
### 文档
|
### 文档
|
||||||
|
|
||||||
| 文档 | 说明 |
|
| 文档 | 说明 |
|
||||||
|------|------|
|
|------|------|
|
||||||
| [参数说明](./params.md) | 训练与推理参数配置 |
|
| [参数说明](./params.md) | 训练与推理参数配置 |
|
||||||
| [设计文档](./design.md) | 系统架构与模块设计 |
|
| [架构文档](./architecture.md) | 系统架构、类图与设计模式 |
|
||||||
| [数据流程](./dataflow.md) | 数据处理管道详解 |
|
| [训练文档](./training.md) | 训练循环、策略与公式 |
|
||||||
| [模型介绍](./introduction.md) | 模型架构与技术细节 |
|
| [推理文档](./inference.md) | KVCache、连续批处理、采样与 HTTP API |
|
||||||
|
| [数据流程](./dataflow.md) | 数据管道、存储后端与数据集架构 |
|
||||||
|
|
||||||
### 贡献
|
### 贡献
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,999 @@
|
|||||||
|
# AstrAI Architecture
|
||||||
|
|
||||||
|
## Class Diagram
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
classDiagram
|
||||||
|
namespace config {
|
||||||
|
class BaseConfig {
|
||||||
|
+to_dict() Dict
|
||||||
|
+from_dict(d) Self
|
||||||
|
}
|
||||||
|
|
||||||
|
class BaseModelConfig {
|
||||||
|
+Optional[str] model_type
|
||||||
|
+from_file(config_path) Self
|
||||||
|
+to_file(config_path)
|
||||||
|
}
|
||||||
|
|
||||||
|
class AutoRegressiveLMConfig {
|
||||||
|
+int vocab_size
|
||||||
|
+int dim
|
||||||
|
+int n_layers
|
||||||
|
+float norm_eps
|
||||||
|
+int dim_ffn
|
||||||
|
+bool tie_weight
|
||||||
|
+int max_len
|
||||||
|
+float rope_theta
|
||||||
|
+str attn_type
|
||||||
|
+int n_heads
|
||||||
|
+int n_kv_heads
|
||||||
|
+bool use_qk_norm
|
||||||
|
+bool use_gated_attention
|
||||||
|
+Optional[int] kv_lora_rank
|
||||||
|
+Optional[int] qk_nope_head_dim
|
||||||
|
+Optional[int] qk_rope_head_dim
|
||||||
|
+str ffn_type
|
||||||
|
+int n_routed_experts
|
||||||
|
+int n_shared_experts
|
||||||
|
+int n_activated_experts
|
||||||
|
+Optional[str] topk_method
|
||||||
|
}
|
||||||
|
|
||||||
|
class EncoderConfig {
|
||||||
|
+int vocab_size
|
||||||
|
+int dim
|
||||||
|
+int n_layers
|
||||||
|
+float norm_eps
|
||||||
|
+int dim_ffn
|
||||||
|
+int max_len
|
||||||
|
+float rope_theta
|
||||||
|
+int n_heads
|
||||||
|
+int n_kv_heads
|
||||||
|
+bool use_qk_norm
|
||||||
|
+bool use_gated_attention
|
||||||
|
+Optional[str] pooling_type
|
||||||
|
+Optional[bool] normalize_embeddings
|
||||||
|
}
|
||||||
|
|
||||||
|
class ConfigFactory {
|
||||||
|
+Registry _registry
|
||||||
|
+register(name) decorator
|
||||||
|
+load(raw) BaseConfig
|
||||||
|
}
|
||||||
|
|
||||||
|
class TrainConfig {
|
||||||
|
+nn.Module model
|
||||||
|
+str strategy
|
||||||
|
+Dataset dataset
|
||||||
|
+Callable optimizer_fn
|
||||||
|
+Callable scheduler_fn
|
||||||
|
+int n_epoch
|
||||||
|
+int batch_per_device
|
||||||
|
+int grad_accum_steps
|
||||||
|
+float max_grad_norm
|
||||||
|
+list gradient_checkpointing_modules
|
||||||
|
+int start_epoch
|
||||||
|
+int start_batch
|
||||||
|
+str ckpt_dir
|
||||||
|
+int ckpt_interval
|
||||||
|
+str log_dir
|
||||||
|
+int log_interval
|
||||||
|
+List[str] metrics
|
||||||
|
+int random_seed
|
||||||
|
+int num_workers
|
||||||
|
+Optional[int] prefetch_factor
|
||||||
|
+bool pin_memory
|
||||||
|
+int nprocs
|
||||||
|
+str backend
|
||||||
|
+str master_addr
|
||||||
|
+str master_port
|
||||||
|
+Callable parallel_wrapper
|
||||||
|
+Callable state_dict_fn
|
||||||
|
+str start_method
|
||||||
|
+str device_type
|
||||||
|
+Optional[Dataset] val_dataset
|
||||||
|
+int val_step
|
||||||
|
+dict extra_kwargs
|
||||||
|
+validate()
|
||||||
|
}
|
||||||
|
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace dataset {
|
||||||
|
class BaseDataset {
|
||||||
|
+int window_size
|
||||||
|
+int stride
|
||||||
|
+Optional[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 StorageFactory {
|
||||||
|
+Registry _registry
|
||||||
|
+register(name) decorator
|
||||||
|
+create(storage_type) BaseStorage
|
||||||
|
}
|
||||||
|
|
||||||
|
class DatasetFactory {
|
||||||
|
+Registry _registry
|
||||||
|
+register(name) decorator
|
||||||
|
+create(train_type, window_size, stride) BaseDataset
|
||||||
|
+load(train_type, load_path, window_size, stride, storage_type, tokenizer) BaseDataset
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace serialization {
|
||||||
|
class Checkpoint {
|
||||||
|
+dict state_dict
|
||||||
|
+int epoch
|
||||||
|
+int iteration
|
||||||
|
+dict extra
|
||||||
|
+dict meta
|
||||||
|
+save(save_dir)
|
||||||
|
+load(save_dir) Checkpoint
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace model {
|
||||||
|
class AutoModel {
|
||||||
|
+BaseModelConfig config
|
||||||
|
+Registry _registry
|
||||||
|
+register(model_type) decorator
|
||||||
|
+get_component_class(model_type) Type
|
||||||
|
+from_pretrained(path, disable_random_init, strict) nn.Module
|
||||||
|
+save_pretrained(save_directory)
|
||||||
|
+to(*args, **kwargs) Self
|
||||||
|
}
|
||||||
|
|
||||||
|
class AutoRegressiveLM {
|
||||||
|
+AutoRegressiveLMConfig config
|
||||||
|
+RotaryEmbedding rotary_embedding
|
||||||
|
+Embedding embed_tokens
|
||||||
|
+ModuleList layers
|
||||||
|
+RMSNorm norm
|
||||||
|
+Linear lm_head
|
||||||
|
+forward(input_ids, input_mask, paged_cache, position_ids) Dict[str, Tensor]
|
||||||
|
+load_state_dict(state_dict)
|
||||||
|
+state_dict()
|
||||||
|
}
|
||||||
|
|
||||||
|
class EmbeddingEncoder {
|
||||||
|
+EncoderConfig config
|
||||||
|
+RotaryEmbedding rotary_embedding
|
||||||
|
+Embedding embed_tokens
|
||||||
|
+ModuleList layers
|
||||||
|
+RMSNorm norm
|
||||||
|
+str pooling_type
|
||||||
|
+bool normalize_embeddings
|
||||||
|
+forward(input_ids, input_mask, position_ids) Tensor
|
||||||
|
+load_state_dict(state_dict)
|
||||||
|
}
|
||||||
|
|
||||||
|
class DecoderBlock {
|
||||||
|
+nn.Module attention # GQA or MLA via AttnFactory
|
||||||
|
+RMSNorm input_norm
|
||||||
|
+nn.Module mlp # MLP or DeepSeekMoE via FFNFactory
|
||||||
|
+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
|
||||||
|
+int n_rep
|
||||||
|
+int layer_id
|
||||||
|
+bool use_qk_norm
|
||||||
|
+bool use_gated_attention
|
||||||
|
+Linear q_proj, k_proj, v_proj, o_proj
|
||||||
|
+Linear gate # only if use_gated_attention
|
||||||
|
+RMSNorm q_norm, k_norm # only if use_qk_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
|
||||||
|
+int n_rep
|
||||||
|
+int layer_id
|
||||||
|
+bool use_gated_attention
|
||||||
|
+Linear q_proj, kv_a_proj, kv_b_proj
|
||||||
|
+Linear o_proj
|
||||||
|
+Linear gate # only if use_gated_attention
|
||||||
|
+RMSNorm kv_norm
|
||||||
|
+forward(x, rotary_emb, attn_mask, paged_cache) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class MLP {
|
||||||
|
+Linear up, gate, down
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class DeepSeekMoE {
|
||||||
|
+int dim
|
||||||
|
+int n_routed_experts
|
||||||
|
+int n_shared_experts
|
||||||
|
+int n_activated_experts
|
||||||
|
+str topk_method
|
||||||
|
+Linear router
|
||||||
|
+ModuleList shared_experts
|
||||||
|
+ModuleList routed_experts
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class AttnFactory {
|
||||||
|
+create(attn_type, **kwargs) nn.Module
|
||||||
|
}
|
||||||
|
|
||||||
|
class FFNFactory {
|
||||||
|
+create(ffn_type, dim, dim_ffn, **kwargs) nn.Module
|
||||||
|
}
|
||||||
|
|
||||||
|
class RMSNorm {
|
||||||
|
+Parameter weight
|
||||||
|
+float norm_eps
|
||||||
|
+tuple normalized_shape
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class Linear {
|
||||||
|
+Parameter weight
|
||||||
|
+Optional[Parameter] bias # only if bias=True
|
||||||
|
+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)
|
||||||
|
+_get_default_callbacks() List[TrainCallback]
|
||||||
|
}
|
||||||
|
|
||||||
|
class TrainContext {
|
||||||
|
+nn.Module model
|
||||||
|
+BaseStrategy strategy
|
||||||
|
+DataLoader dataloader
|
||||||
|
+Optimizer optimizer
|
||||||
|
+LRScheduler scheduler
|
||||||
|
+Checkpoint checkpoint
|
||||||
|
+TrainConfig config
|
||||||
|
+int epoch
|
||||||
|
+int iteration
|
||||||
|
+float loss
|
||||||
|
+DataLoader val_dataloader
|
||||||
|
+float val_loss
|
||||||
|
+int world_size
|
||||||
|
+int rank
|
||||||
|
+dict kwargs
|
||||||
|
}
|
||||||
|
|
||||||
|
class TrainContextBuilder {
|
||||||
|
+TrainConfig config
|
||||||
|
+with_checkpoint(checkpoint) TrainContextBuilder
|
||||||
|
+build() TrainContext
|
||||||
|
}
|
||||||
|
|
||||||
|
class BaseStrategy {
|
||||||
|
+Union[Callable, nn.Module] model
|
||||||
|
+str device
|
||||||
|
+compute_loss(batch) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class StrategyFactory {
|
||||||
|
+Registry _registry
|
||||||
|
+register(name) decorator
|
||||||
|
+create(train_type, model, 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
|
||||||
|
+sync_ref_model()
|
||||||
|
}
|
||||||
|
|
||||||
|
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 {
|
||||||
|
<<protocol>>
|
||||||
|
+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 GradientCheckpointingCallback {
|
||||||
|
+tuple modules
|
||||||
|
+on_train_begin(context)
|
||||||
|
+on_train_end(context)
|
||||||
|
}
|
||||||
|
|
||||||
|
class CheckpointCallback {
|
||||||
|
+str save_dir
|
||||||
|
+int interval
|
||||||
|
+bool weight_only
|
||||||
|
+Callable state_dict_fn
|
||||||
|
+Callable save_extra_fn
|
||||||
|
+Callable load_extra_fn
|
||||||
|
+_save_checkpoint(context)
|
||||||
|
+on_train_begin(context)
|
||||||
|
+on_batch_end(context)
|
||||||
|
+on_train_end(context)
|
||||||
|
+on_error(context)
|
||||||
|
+save_extra(context)$
|
||||||
|
+load_extra(extra, context)$
|
||||||
|
}
|
||||||
|
|
||||||
|
class ProgressBarCallback {
|
||||||
|
+int num_epoch
|
||||||
|
+int log_interval
|
||||||
|
+IO file
|
||||||
|
+on_epoch_begin(context)
|
||||||
|
+on_batch_end(context)
|
||||||
|
+on_epoch_end(context)
|
||||||
|
}
|
||||||
|
|
||||||
|
class MetricLoggerCallback {
|
||||||
|
+str log_dir
|
||||||
|
+int save_interval
|
||||||
|
+int log_interval
|
||||||
|
+List[str] metrics
|
||||||
|
+on_batch_end(context)
|
||||||
|
+on_train_end(context)
|
||||||
|
+on_error(context)
|
||||||
|
}
|
||||||
|
|
||||||
|
class ValidationCallback {
|
||||||
|
+_run_validation(context)
|
||||||
|
+on_step_end(context)
|
||||||
|
}
|
||||||
|
|
||||||
|
class CallbackFactory {
|
||||||
|
+Registry _registry
|
||||||
|
+register(name) decorator
|
||||||
|
+create(name, **kwargs) TrainCallback
|
||||||
|
}
|
||||||
|
|
||||||
|
class Muon {
|
||||||
|
+float lr
|
||||||
|
+float momentum
|
||||||
|
+float weight_decay
|
||||||
|
+int ns_steps
|
||||||
|
+step(closure) Optional[float]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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 Executor {
|
||||||
|
+AutoModel model
|
||||||
|
+AutoTokenizer tokenizer
|
||||||
|
+KVCache page_cache
|
||||||
|
+execute_prefill(tasks, prompt_len, start_pos)
|
||||||
|
+execute_decode(tasks) List[int]
|
||||||
|
}
|
||||||
|
|
||||||
|
class InferenceScheduler {
|
||||||
|
+KVCache _page_cache
|
||||||
|
+Executor _executor
|
||||||
|
+TaskManager _task_mgr
|
||||||
|
+bool _running
|
||||||
|
+Thread _loop_thread
|
||||||
|
+int max_seq_len
|
||||||
|
+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
|
||||||
|
+List[int] _refs
|
||||||
|
+OrderedDict _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 page_size
|
||||||
|
+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 TaskManager {
|
||||||
|
+AutoTokenizer tokenizer
|
||||||
|
+Deque waiting_queue
|
||||||
|
+List active_tasks
|
||||||
|
+add_task(prompt, **kwargs) str
|
||||||
|
+remove_task(task_id) List[Task]
|
||||||
|
+remove_finished_tasks(stop_ids) List[Task]
|
||||||
|
+pull_candidates(n) List[Task]
|
||||||
|
+activate(task)
|
||||||
|
+return_to_waiting(tasks)
|
||||||
|
+get_active_tasks() List[Task]
|
||||||
|
}
|
||||||
|
|
||||||
|
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[BaseSamplingStrategy] 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[Tuple[int, str]]
|
||||||
|
+wait(timeout) bool
|
||||||
|
+wait_completion(timeout)
|
||||||
|
}
|
||||||
|
|
||||||
|
class ChatMessage {
|
||||||
|
+str role
|
||||||
|
+str content
|
||||||
|
}
|
||||||
|
|
||||||
|
class ChatCompletionRequest {
|
||||||
|
+str model
|
||||||
|
+List[ChatMessage] messages
|
||||||
|
+Optional[float] temperature
|
||||||
|
+Optional[float] top_p
|
||||||
|
+Optional[int] top_k
|
||||||
|
+Optional[int] max_tokens
|
||||||
|
+Optional[bool] stream
|
||||||
|
+Optional[Union[str, List[str]]] stop
|
||||||
|
+Optional[int] n
|
||||||
|
+Optional[float] presence_penalty
|
||||||
|
+Optional[float] frequency_penalty
|
||||||
|
+Optional[Dict[int, float]] logit_bias
|
||||||
|
+Optional[str] user
|
||||||
|
}
|
||||||
|
|
||||||
|
class AnthropicMessage {
|
||||||
|
+str role
|
||||||
|
+Union[str, List[Dict]] content
|
||||||
|
}
|
||||||
|
|
||||||
|
class MessagesRequest {
|
||||||
|
+str model
|
||||||
|
+List[AnthropicMessage] messages
|
||||||
|
+Optional[str] system
|
||||||
|
+float temperature
|
||||||
|
+float top_p
|
||||||
|
+int top_k
|
||||||
|
+int max_tokens
|
||||||
|
+bool stream
|
||||||
|
+Optional[List[str]] stop_sequences
|
||||||
|
}
|
||||||
|
|
||||||
|
class ProtocolHandler {
|
||||||
|
<<abstract>>
|
||||||
|
+request
|
||||||
|
+engine
|
||||||
|
+build_prompt() str
|
||||||
|
+create_response_id() str
|
||||||
|
+get_stop_sequences() List[str]
|
||||||
|
+create_stop_checker() StopChecker
|
||||||
|
+on_token(ctx, token, stop_checker) Optional[str]
|
||||||
|
+format_stream_start(ctx) List[str]
|
||||||
|
+format_stream_token(ctx, token) str
|
||||||
|
+format_stream_end(ctx) List[str]
|
||||||
|
+format_non_stream_response(ctx, content) Dict
|
||||||
|
+handle() Union[StreamingResponse, Dict]
|
||||||
|
}
|
||||||
|
|
||||||
|
class OpenAIHandler {
|
||||||
|
+build_prompt() str
|
||||||
|
+create_response_id() str
|
||||||
|
}
|
||||||
|
|
||||||
|
class AnthropicHandler {
|
||||||
|
+build_prompt() str
|
||||||
|
+create_response_id() str
|
||||||
|
+on_token(ctx, token, stop_checker) Optional[str]
|
||||||
|
}
|
||||||
|
|
||||||
|
class StopChecker {
|
||||||
|
+has_sequences (property) bool
|
||||||
|
+check(text) Optional[str]
|
||||||
|
+trim(text, matched) str
|
||||||
|
}
|
||||||
|
|
||||||
|
class StreamContext {
|
||||||
|
+str resp_id
|
||||||
|
+int created
|
||||||
|
+str model
|
||||||
|
+int prompt_tokens
|
||||||
|
+int completion_tokens
|
||||||
|
+str accumulated
|
||||||
|
+Optional[str] stop_matched
|
||||||
|
+str last_yield_trimmed
|
||||||
|
}
|
||||||
|
|
||||||
|
class app {
|
||||||
|
<<singleton>>
|
||||||
|
+FastAPI app
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace parallel {
|
||||||
|
class Functions {
|
||||||
|
<<module>>
|
||||||
|
+spawn_parallel_fn(func, world_size, backend, master_addr, master_port, device_type, start_method, **kwargs)
|
||||||
|
+setup_parallel(rank, world_size, backend, master_addr, master_port, device_type)
|
||||||
|
+get_current_device() str
|
||||||
|
+get_world_size() int
|
||||||
|
+get_rank() int
|
||||||
|
+only_on_rank(rank, sync) decorator
|
||||||
|
}
|
||||||
|
|
||||||
|
class ParallelModel {
|
||||||
|
+dist.ProcessGroup process_group
|
||||||
|
+int rank
|
||||||
|
+int world_size
|
||||||
|
}
|
||||||
|
|
||||||
|
class ColumnParallelLinear {
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class RowParallelLinear {
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
%% Relationships — UML notation: <|-- generalization, *-- composition, o-- aggregation, --> association, ..> dependency
|
||||||
|
|
||||||
|
%% --- Generalization (inheritance) ---
|
||||||
|
BaseStrategy <|-- SEQStrategy
|
||||||
|
BaseStrategy <|-- SFTStrategy
|
||||||
|
BaseStrategy <|-- DPOStrategy
|
||||||
|
BaseStrategy <|-- GRPOStrategy
|
||||||
|
BaseScheduler <|-- CosineScheduler
|
||||||
|
BaseScheduler <|-- SGDRScheduler
|
||||||
|
TrainCallback <|-- GradientClippingCallback
|
||||||
|
TrainCallback <|-- GradientCheckpointingCallback
|
||||||
|
TrainCallback <|-- CheckpointCallback
|
||||||
|
TrainCallback <|-- ProgressBarCallback
|
||||||
|
TrainCallback <|-- MetricLoggerCallback
|
||||||
|
BaseDataset <|-- SEQDataset
|
||||||
|
BaseDataset <|-- SFTDataset
|
||||||
|
BaseDataset <|-- DPODataset
|
||||||
|
BaseDataset <|-- GRPODataset
|
||||||
|
BaseStorage <|-- H5Storage
|
||||||
|
BaseStorage <|-- JSONStorage
|
||||||
|
BaseSamplingStrategy <|-- TemperatureStrategy
|
||||||
|
BaseSamplingStrategy <|-- TopKStrategy
|
||||||
|
BaseSamplingStrategy <|-- TopPStrategy
|
||||||
|
BaseSamplingStrategy <|-- SamplingPipeline
|
||||||
|
ParallelModel <|-- RowParallelLinear
|
||||||
|
ParallelModel <|-- ColumnParallelLinear
|
||||||
|
AutoModel <|-- AutoRegressiveLM
|
||||||
|
AutoModel <|-- EmbeddingEncoder
|
||||||
|
BaseConfig <|-- BaseModelConfig
|
||||||
|
BaseConfig <|-- TrainConfig
|
||||||
|
BaseModelConfig <|-- AutoRegressiveLMConfig
|
||||||
|
BaseModelConfig <|-- EncoderConfig
|
||||||
|
BaseFactory <|-- AutoModel
|
||||||
|
BaseFactory <|-- AttnFactory
|
||||||
|
BaseFactory <|-- FFNFactory
|
||||||
|
BaseFactory <|-- DatasetFactory
|
||||||
|
BaseFactory <|-- StrategyFactory
|
||||||
|
BaseFactory <|-- SchedulerFactory
|
||||||
|
BaseFactory <|-- CallbackFactory
|
||||||
|
BaseFactory <|-- StorageFactory
|
||||||
|
BaseFactory <|-- ConfigFactory
|
||||||
|
TrainCallback <|-- ValidationCallback
|
||||||
|
ProtocolHandler <|-- OpenAIHandler
|
||||||
|
ProtocolHandler <|-- AnthropicHandler
|
||||||
|
|
||||||
|
%% --- Composition (strong ownership, part destroyed with whole) ---
|
||||||
|
KVCache *-- PagePool
|
||||||
|
KVCache *-- Storage
|
||||||
|
KVCache *-- TaskTable
|
||||||
|
PagePool *-- Allocator
|
||||||
|
PagePool *-- PrefixCache
|
||||||
|
InferenceEngine *-- InferenceScheduler
|
||||||
|
InferenceScheduler *-- KVCache
|
||||||
|
InferenceScheduler *-- Executor
|
||||||
|
InferenceScheduler *-- TaskManager
|
||||||
|
AutoRegressiveLM *-- DecoderBlock
|
||||||
|
AutoRegressiveLM *-- RotaryEmbedding
|
||||||
|
AutoRegressiveLM *-- Embedding
|
||||||
|
EmbeddingEncoder *-- DecoderBlock
|
||||||
|
EmbeddingEncoder *-- RotaryEmbedding
|
||||||
|
EmbeddingEncoder *-- Embedding
|
||||||
|
DecoderBlock *-- RMSNorm
|
||||||
|
ChatCompletionRequest *-- ChatMessage
|
||||||
|
MessagesRequest *-- AnthropicMessage
|
||||||
|
AutoTokenizer *-- ChatTemplate
|
||||||
|
BaseFactory *-- Registry
|
||||||
|
|
||||||
|
%% --- Aggregation (weak ownership) ---
|
||||||
|
AutoModel o-- BaseModelConfig
|
||||||
|
Trainer o-- TrainCallback
|
||||||
|
TrainContext o-- BaseStrategy
|
||||||
|
TrainContext o-- BaseScheduler
|
||||||
|
TrainContext o-- Checkpoint
|
||||||
|
KvcacheView o-- Storage
|
||||||
|
SamplingPipeline o-- BaseSamplingStrategy
|
||||||
|
BaseDataset o-- BaseStorage
|
||||||
|
|
||||||
|
%% --- Dependency (uses temporarily) ---
|
||||||
|
TrainConfig ..> BaseStrategy : selects
|
||||||
|
StrategyFactory ..> BaseStrategy : creates
|
||||||
|
SchedulerFactory ..> BaseScheduler : creates
|
||||||
|
DatasetFactory ..> BaseDataset : creates
|
||||||
|
CallbackFactory ..> TrainCallback : creates
|
||||||
|
AttnFactory ..> GQA : creates
|
||||||
|
AttnFactory ..> MLA : creates
|
||||||
|
FFNFactory ..> MLP : creates
|
||||||
|
FFNFactory ..> DeepSeekMoE : creates
|
||||||
|
DecoderBlock ..> AttnFactory : uses
|
||||||
|
DecoderBlock ..> FFNFactory : uses
|
||||||
|
StorageFactory ..> H5Storage : creates
|
||||||
|
StorageFactory ..> JSONStorage : creates
|
||||||
|
ConfigFactory ..> AutoRegressiveLMConfig : creates
|
||||||
|
ConfigFactory ..> EncoderConfig : creates
|
||||||
|
Trainer ..> TrainContextBuilder : uses
|
||||||
|
TrainContextBuilder ..> TrainContext : creates
|
||||||
|
Trainer ..> Functions : spawns
|
||||||
|
TrainContextBuilder ..> StrategyFactory : uses
|
||||||
|
TrainContextBuilder ..> ResumableDistributedSampler : creates
|
||||||
|
Checkpoint ..> Checkpoint : serializes
|
||||||
|
CheckpointCallback ..> Checkpoint : creates
|
||||||
|
KVCache ..> KvcacheView : binds
|
||||||
|
InferenceEngine ..> GenerationRequest : uses
|
||||||
|
InferenceEngine ..> GenerateResult : creates
|
||||||
|
OpenAIHandler ..> ChatCompletionRequest : receives
|
||||||
|
AnthropicHandler ..> MessagesRequest : receives
|
||||||
|
ProtocolHandler ..> StopChecker : creates
|
||||||
|
ProtocolHandler ..> StreamContext : creates
|
||||||
|
|
||||||
|
%% --- Association (general usage) ---
|
||||||
|
Trainer --> TrainConfig
|
||||||
|
DPOStrategy --> AutoModel
|
||||||
|
GRPOStrategy --> AutoModel
|
||||||
|
InferenceScheduler --> Task
|
||||||
|
InferenceScheduler --> TaskStatus
|
||||||
|
Task --> TaskStatus
|
||||||
|
InferenceEngine --> AutoModel
|
||||||
|
Executor --> AutoModel
|
||||||
|
Executor --> AutoTokenizer
|
||||||
|
TaskManager --> AutoTokenizer
|
||||||
|
MultiSegmentFetcher --> BaseSegmentFetcher
|
||||||
|
ResumableDistributedSampler --> BaseDataset
|
||||||
|
|
||||||
|
```
|
||||||
|
|
||||||
|
|
||||||
|
## Module Overview
|
||||||
|
|
||||||
|
| Module | Components | Description |
|
||||||
|
|--------|------------|-------------|
|
||||||
|
| **astrai.config** | BaseConfig, BaseModelConfig, AutoRegressiveLMConfig, EncoderConfig, ConfigFactory, TrainConfig | Configuration management (to_dict/from_dict, to_file/from_file) |
|
||||||
|
| **astrai.dataset** | BaseDataset–GRPODataset, BaseStorage–JSONStorage, StorageFactory, BaseSegmentFetcher, MultiSegmentFetcher, ResumableDistributedSampler, DatasetFactory | Dataset loading and management |
|
||||||
|
| **astrai.serialization** | Checkpoint | Model serialization |
|
||||||
|
| **astrai.model** | AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, RotaryEmbedding, Embedding | Neural network model |
|
||||||
|
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
||||||
|
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–SGDRScheduler, SchedulerFactory, TrainCallback(Protocol)–ValidationCallback, CallbackFactory, Muon | Training workflow |
|
||||||
|
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, KVCache–KvcacheView, Allocator–Storage, Task, TaskManager, TaskStatus, GenerationRequest, BaseSamplingStrategy–SamplingPipeline, ProtocolHandler–AnthropicHandler, ChatMessage–MessagesRequest, app | Inference service |
|
||||||
|
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, ParallelModel, RowParallelLinear, ColumnParallelLinear | Distributed parallel |
|
||||||
|
| **astrai.factory** | Registry, BaseFactory[T] | Component registration |
|
||||||
|
|
||||||
|
## Design Patterns
|
||||||
|
|
||||||
|
| Pattern | Classes | Purpose |
|
||||||
|
|---------|---------|---------|
|
||||||
|
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StorageFactory`, `ConfigFactory` | Decorator-based component creation |
|
||||||
|
| **Registry** | `BaseFactory`, `Registry` | Component registration with category/priority |
|
||||||
|
| **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching |
|
||||||
|
| **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations |
|
||||||
|
| **Template Method** | `ProtocolHandler`, `OpenAIHandler`, `AnthropicHandler` | HTTP API handler with format hooks |
|
||||||
|
| **Builder** | `TrainContextBuilder` | Chain-building training context |
|
||||||
|
| **Observer** | `TrainCallback`, callback implementations | Training process monitoring |
|
||||||
|
| **Context** | `TrainContext` | Unified training state bag |
|
||||||
|
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction |
|
||||||
|
| **Storage** | `BaseStorage`, `H5Storage`, `JSONStorage` | Format-agnostic data access |
|
||||||
|
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
|
||||||
|
| **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
|
||||||
|
|
||||||
|
## Core Relationships
|
||||||
|
|
||||||
|
1. **Config → Training**: `TrainConfig` holds model, dataset, optimizer_fn, scheduler_fn
|
||||||
|
2. **Training Flow**: `Trainer` → `TrainContextBuilder` → `TrainContext`, uses `BaseStrategy` for loss
|
||||||
|
3. **Strategy Selection**: `StrategyFactory` creates strategy by `train_type`
|
||||||
|
4. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `AutoRegressiveLM`, backed by `KVCache` + `SamplingPipeline`
|
||||||
|
5. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
|
||||||
|
6. **Dataset Loading**: `DatasetFactory` creates datasets, `BaseStorage` (H5Storage/JSONStorage) loads via `BaseSegmentFetcher` + `MultiSegmentFetcher`
|
||||||
|
7. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only)
|
||||||
|
8. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`
|
||||||
|
9. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
|
||||||
|
|
||||||
|
> Document Update Time: 2026-05-17
|
||||||
+36
-258
@@ -1,279 +1,57 @@
|
|||||||
# AstrAI Data Flow Documentation
|
# Data Flow
|
||||||
|
|
||||||
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.
|
This document describes the data pipeline: from raw text to model input tensors.
|
||||||
|
|
||||||
## Overview
|
## Overview
|
||||||
|
|
||||||
AstrAI adopts a modular design with the following main components:
|
```
|
||||||
- **Dataset Module** (`astrai/dataset/`): Dataset, sampler, serialization tools
|
Raw Text → AutoTokenizer → Token IDs → .h5/.json → Dataset → Sampler → DataLoader → Training/Inference
|
||||||
- **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/`): Model, training, scheduler, and other configurations
|
|
||||||
- **Factory Module** (`astrai/factory/`): Registry, BaseFactory for component registration
|
|
||||||
- **Parallel Module** (`astrai/parallel/`): Distributed training support
|
|
||||||
- **Serialization** (`astrai/serialization.py`): HDF5 data loading, checkpoint management
|
|
||||||
|
|
||||||
The data flow can generally be divided into two main lines: **Training Data Flow** and **Inference Data Flow**.
|
|
||||||
|
|
||||||
## Data Flow Diagram
|
|
||||||
|
|
||||||
```mermaid
|
|
||||||
flowchart LR
|
|
||||||
subgraph A[Data Preparation]
|
|
||||||
direction TB
|
|
||||||
A1[Raw Text] --> A2[AutoTokenizer]
|
|
||||||
A2 --> A3[Serialize to .h5 files]
|
|
||||||
A3 --> A4[BaseDataset]
|
|
||||||
A4 --> A5[ResumableDistributedSampler]
|
|
||||||
A5 --> A6[PyTorch DataLoader]
|
|
||||||
end
|
|
||||||
|
|
||||||
subgraph B[Training]
|
|
||||||
direction TB
|
|
||||||
B1[Batch Data] --> B2[TrainContextBuilder]
|
|
||||||
B2 --> B3[TrainContext]
|
|
||||||
B3 --> B4[BaseStrategy]
|
|
||||||
B4 --> B5[Transformer]
|
|
||||||
B5 --> B6[Compute Loss]
|
|
||||||
B6 --> B7[Backward]
|
|
||||||
B7 --> B8[Optimizer]
|
|
||||||
B8 --> B9[LRScheduler]
|
|
||||||
B9 --> B10[CheckpointCallback]
|
|
||||||
end
|
|
||||||
|
|
||||||
subgraph C[Inference]
|
|
||||||
direction TB
|
|
||||||
C1[Checkpoint] --> C2[AutoModel]
|
|
||||||
C2 --> C3[Transformer + Tokenizer]
|
|
||||||
C3 --> C4[GenerationRequest + apply_chat_template]
|
|
||||||
C4 --> C5[InferenceEngine]
|
|
||||||
C5 --> C6[InferenceScheduler]
|
|
||||||
C6 --> C7[sample]
|
|
||||||
C7 --> C8[Transformer Forward]
|
|
||||||
C8 --> C9[Paged KV Cache]
|
|
||||||
C9 --> C10{End Condition?}
|
|
||||||
C10 -->|No| C8
|
|
||||||
C10 -->|Yes| C11[Output Text]
|
|
||||||
end
|
|
||||||
|
|
||||||
A --> B
|
|
||||||
B --> C
|
|
||||||
```
|
```
|
||||||
|
|
||||||
## Detailed Module Descriptions
|
## Data Preparation
|
||||||
|
|
||||||
### 1. Serialization (`astrai/serialization.py`)
|
Raw text is tokenized via `AutoTokenizer.encode()` and saved as HDF5 (`.h5`) or JSON (`.json`/`.jsonl`) files with keyed tensor groups.
|
||||||
|
|
||||||
- **`save_h5`**: Saves multiple tensors by groups as HDF5 files (`.h5`), each key corresponds to a list of tensors
|
Storage format is auto-detected by `detect_format()`; backends are dispatched via registry:
|
||||||
- **`load_h5`**: Loads `.h5` files, returns `Dict[str, List[Tensor]]`, supports shared memory (`share_memory=True`)
|
|
||||||
- **`Checkpoint` class**: Encapsulates model state dict, training epoch, iteration count; supports safetensors format for saving and loading
|
|
||||||
|
|
||||||
### 2. Dataset Module
|
```
|
||||||
|
StorageFactory.create("h5") → H5Storage
|
||||||
|
StorageFactory.create("json") → JSONStorage
|
||||||
|
```
|
||||||
|
|
||||||
#### 2.1 Dataset (`dataset.py`)
|
Both support shared memory via `.share_memory_()`.
|
||||||
- **`BaseDataset`**: Abstract base class, defines common logic for window sampling, stride, etc.
|
|
||||||
- **`BaseSegmentFetcher`** and **`MultiSegmentFetcher`**: Efficiently fetch data from specified index ranges in multiple segments
|
|
||||||
- **`DatasetFactory`**: Factory pattern, supports dynamic registration of dataset types (`seq`, `sft`, `dpo`, `grpo`)
|
|
||||||
- After dataset loading, multiple data keys (such as `"sequence"`, `"mask"`) are managed through `MultiSegmentFetcher`
|
|
||||||
|
|
||||||
#### 2.2 Sampler (`sampler.py`)
|
## Data Keys by Training Type
|
||||||
- **`ResumableDistributedSampler`**: Resumable sampler supporting distributed training
|
|
||||||
- Records current epoch and iteration position, enabling training resume from breakpoints
|
|
||||||
- Supports shuffle and drop_last options
|
|
||||||
|
|
||||||
### 3. Model Module
|
| Type | Storage Keys |
|
||||||
|
|------|-------------|
|
||||||
|
| `seq` | `sequence` (→ input_ids, target_ids via offset-by-1) |
|
||||||
|
| `sft` | `sequence`, `loss_mask` |
|
||||||
|
| `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` |
|
||||||
|
| `grpo` | `prompts`, `responses`, `masks`, `rewards` |
|
||||||
|
|
||||||
#### 3.1 Transformer / AutoModel (`transformer.py`, `automodel.py`)
|
## Dataset Architecture
|
||||||
- **`AutoModel`**: Base class for autoregressive language models with `from_pretrained()` and `save_pretrained()` methods
|
|
||||||
- **`Transformer`**: Core autoregressive decoder architecture (registered via `@AutoModel.register('transformer')`)
|
|
||||||
- Contains embedding layer, multi-layer `DecoderBlock`, RMSNorm, and linear output head
|
|
||||||
- Supports weight tying (`tie_weight=True`) to reduce parameter count
|
|
||||||
- Uses Rotary Position Embedding (RoPE) to inject position information
|
|
||||||
- Supports loading from safetensors format with automatic model type detection from `config.json`
|
|
||||||
|
|
||||||
#### 3.2 Submodules (`module.py`)
|
```
|
||||||
- **`RotaryEmbedding`**: Generates RoPE cos/sin cache
|
DatasetFactory.load(train_type, path, window_size, stride)
|
||||||
- **`DecoderBlock`**: Contains multi-head attention (supports GQA and MLA), feedforward network (FFN), residual connections
|
→ StorageFactory.create(detect_format(path))
|
||||||
- **`GQA`**: Grouped Query Attention implementation
|
→ MultiSegmentFetcher(BaseSegmentFetcher per key)
|
||||||
- **`MLA`**: Multi-Latent Attention implementation (like Qwen2-VL)
|
→ BaseDataset.__getitem__(idx)
|
||||||
- **`MLP`**: Feed-forward network with SiLU activation and gated mechanism
|
→ sliding window [begin, end) via get_index(idx)
|
||||||
- **`RMSNorm`**: Layer normalization variant
|
```
|
||||||
- **`Linear`**, **`Embedding`**: Custom linear layer and embedding layer, supporting parallelism wrappers
|
|
||||||
|
|
||||||
### 4. Training Module
|
`window_size` = max input length, `stride` = step between consecutive samples.
|
||||||
|
|
||||||
#### 4.1 Training Context (`train_context.py`)
|
## Sampler
|
||||||
- **`TrainContext`**: Data class encapsulating all components needed for training (model, optimizer, data loader, strategy, etc.)
|
|
||||||
- **`TrainContextBuilder`**: Builder pattern, progressively assembles training context, supports resume from checkpoint
|
|
||||||
|
|
||||||
#### 4.2 Trainer (`trainer.py`)
|
`ResumableDistributedSampler` supports checkpoint-aware distributed sampling:
|
||||||
- **`Trainer`**: Main training loop, manages callbacks (progress bar, checkpoint, metric logging, gradient clipping, scheduler)
|
|
||||||
- Supports distributed training (launches multi-process via `spawn_parallel_fn`)
|
|
||||||
- Training steps include:
|
|
||||||
1. `on_train_begin` → 2. `on_epoch_begin` → 3. `on_batch_begin` → 4. Forward/loss calculation → 5. `on_batch_end` → 6. Gradient accumulation → 7. `on_step_begin` → 8. Optimizer update → 9. `on_step_end` → 10. `on_epoch_end`
|
|
||||||
|
|
||||||
#### 4.3 Strategy (`strategy.py`)
|
- Tracks `start_epoch` / `start_iter` for resume
|
||||||
- **`BaseStrategy`**: Defines training strategy interface
|
- Shuffle via `torch.Generator(seed + epoch)`
|
||||||
- **`SEQStrategy`**: Standard next-token prediction training
|
- Per-replica index slicing for DDP
|
||||||
- **`SFTStrategy`**: Supervised Fine-tuning with loss masking
|
|
||||||
- **`DPOStrategy`**: Direct Preference Optimization
|
|
||||||
- **`GRPOStrategy`**: Group Relative Policy Optimization
|
|
||||||
- Strategy receives batch data, executes model forward pass, loss calculation, returns loss tensor
|
|
||||||
- Created dynamically by `StrategyFactory` according to configuration
|
|
||||||
|
|
||||||
#### 4.4 Scheduler (`schedule.py`)
|
## DataLoader
|
||||||
- **`BaseScheduler`**: Abstract base class defining learning rate scheduling interface
|
|
||||||
- **`CosineScheduler`**: Cosine decay scheduler with warmup
|
|
||||||
- **`SGDRScheduler`**: Stochastic Gradient Descent with Warm Restarts
|
|
||||||
- **`SchedulerFactory`**: Factory pattern, supports registration of various schedulers
|
|
||||||
- Scheduler is automatically created according to configuration and bound to optimizer
|
|
||||||
|
|
||||||
#### 4.5 Callbacks (`train_callback.py`)
|
Standard PyTorch `DataLoader` with configurable `batch_size`, `num_workers`, `pin_memory`, `prefetch_factor`. Sampler produces indices; dataloader fetches tensor batches via `__getitem__`.
|
||||||
- **`TrainCallback`**: Protocol interface for trainer callbacks
|
|
||||||
- **`CheckpointCallback`**: Saves model checkpoints at configurable intervals
|
|
||||||
- **`ProgressBarCallback`**: Displays training progress
|
|
||||||
- **`MetricLoggerCallback`**: Logs training metrics to JSON files
|
|
||||||
- **`GradientClippingCallback`**: Clips gradient norms
|
|
||||||
- **`SchedulerCallback`**: Steps learning rate scheduler
|
|
||||||
|
|
||||||
#### 4.6 Metric Utility (`metric_util.py`)
|
> Document Update Time: 2026-05-17
|
||||||
- **`MetricTracker`**: Tracks and aggregates training metrics across epochs
|
|
||||||
- **`get_learning_rate`**: Utility to extract current learning rates from optimizer param groups
|
|
||||||
|
|
||||||
### 5. Factory Module
|
|
||||||
|
|
||||||
#### 5.1 Registry and BaseFactory (`factory.py`)
|
|
||||||
- **`Registry`**: Flexible registry for component classes with category and priority support
|
|
||||||
- **`BaseFactory`**: Generic factory class for component registration and creation
|
|
||||||
- Supports decorator-based registration pattern for extensible components
|
|
||||||
- Provides methods for registration, retrieval, and listing with filtering
|
|
||||||
|
|
||||||
### 6. Parallel Module
|
|
||||||
|
|
||||||
#### 6.1 Setup (`setup.py`)
|
|
||||||
- **`spawn_parallel_fn`**: Spawns multiple processes for distributed training using PyTorch multiprocessing
|
|
||||||
- **`setup_parallel`**: Context manager for initializing distributed process group (NCCL/CCL backend)
|
|
||||||
- **`only_on_rank`**: Decorator to execute functions only on specific ranks
|
|
||||||
- **`get_rank`**: Returns current process rank in distributed group
|
|
||||||
- **`get_world_size`**: Returns total number of processes in distributed group
|
|
||||||
- **`get_current_device`**: Returns current device from environment
|
|
||||||
|
|
||||||
#### 6.2 Parallel Layers (`module.py`)
|
|
||||||
- **`ParallelModel`**: Base class for parallel models with process group
|
|
||||||
- **`ColumnParallelLinear`**: Column-parallel linear layer with input splitting and output gathering
|
|
||||||
- **`RowParallelLinear`**: Row-parallel linear layer with output reduction
|
|
||||||
|
|
||||||
### 7. Inference Module
|
|
||||||
|
|
||||||
#### 7.1 Inference Engine (`engine.py`)
|
|
||||||
- **`InferenceEngine`**: Unified inference interface, supports streaming, async streaming, and non-streaming generation
|
|
||||||
- **`InferenceScheduler`**: Continuous batching scheduler with paged KV cache
|
|
||||||
- **`GenerationRequest`**: Encapsulates generation parameters (top_k, top_p, temperature, max_len, messages, etc.)
|
|
||||||
- **`GenerationParams`**: Immutable value object for sampling hyperparameters
|
|
||||||
- **`messages` format**: List of message dictionaries with `role` (system/user/assistant) and `content`
|
|
||||||
- **`apply_chat_template`** (from `tokenizer.py`): Converts messages into prompt string using ChatML format
|
|
||||||
- Provides streaming (`stream=True`), async streaming (`generate_async`), and non-streaming (`stream=False`) generation interfaces
|
|
||||||
- Supports continuous batching with `max_batch_size` and `max_seq_len` parameters
|
|
||||||
- Uses separate model and tokenizer initialization for flexibility
|
|
||||||
|
|
||||||
#### 7.2 Cache (`cache.py`)
|
|
||||||
- **`PagedCache`**: Page-based KV cache with page-table-indirected read/write; uses bitmask for O(1) page allocation/deallocation
|
|
||||||
- **`CacheView`**: Per-batch view bundling a `PagedCache` with its page table for attention layer access
|
|
||||||
|
|
||||||
#### 7.3 Scheduler (`scheduler.py`)
|
|
||||||
- **`Task`**: Individual generation task with state management (PENDING, RUNNING, FINISHED, ABORTED)
|
|
||||||
- **`TaskStatus`**: Task state enumeration
|
|
||||||
- **`sample`** (from `sampling.py`): Applies temperature, top-k, top-p sampling to logits via composable `SamplingPipeline`
|
|
||||||
- Uses `PagedCache` for paged KV cache management with page table indirection
|
|
||||||
- Continuous batching: new requests can join at any time, completed requests release pages immediately
|
|
||||||
|
|
||||||
#### 7.4 Server (`server.py`)
|
|
||||||
- FastAPI-based HTTP inference server
|
|
||||||
- OpenAI-compatible `/v1/chat/completions` endpoint
|
|
||||||
- Health check and statistics endpoints
|
|
||||||
- Supports both streaming and non-streaming responses
|
|
||||||
|
|
||||||
### 8. Tokenizer Module
|
|
||||||
|
|
||||||
#### 8.1 Tokenizer (`tokenizer.py`)
|
|
||||||
- Implemented based on HuggingFace tokenizers library (Byte-Level BPE)
|
|
||||||
- **`AutoTokenizer`**: Auto-loading tokenizer class
|
|
||||||
- Supports special tokens: `<|begin▁of▁sentence|>`, `<|end▁of▁sentence|>`, `<|▁pad▁|>`, `<|im▁start|>`, `<|im▁end|>`
|
|
||||||
- Provides `encode`/`decode` methods for mutual conversion between text and token IDs
|
|
||||||
- Uses `AutoTokenizer` for loading pre-trained tokenizers
|
|
||||||
|
|
||||||
#### 8.2 Chat Template (`chat_template.py`)
|
|
||||||
- **`ChatTemplate`**: Jinja2-based chat template with rendering support
|
|
||||||
- Handles multi-role message formatting (system, user, assistant)
|
|
||||||
- Supports dynamic prompts and generation prompts
|
|
||||||
|
|
||||||
## Training Data Flow - Detailed Steps
|
|
||||||
|
|
||||||
1. **Data Preparation**
|
|
||||||
- Raw text is converted to token ID sequences through AutoTokenizer
|
|
||||||
- Token ID sequences (possibly with masks, labels, etc.) are saved by groups as `.h5` files
|
|
||||||
- Files can contain multiple segments, each segment corresponds to a tensor
|
|
||||||
|
|
||||||
2. **Dataset Loading**
|
|
||||||
- `BaseDataset`'s `load` method calls `load_h5`, obtaining `segments` dictionary
|
|
||||||
- Create `MultiSegmentFetcher` to manage data for multiple keys
|
|
||||||
- Calculate total sample count, and determine start/end indices for each sample based on window size and stride
|
|
||||||
|
|
||||||
3. **Sampling and Batch Loading**
|
|
||||||
- `ResumableDistributedSampler` generates index sequence based on current epoch and iteration position
|
|
||||||
- PyTorch `DataLoader` uses sampler to get indices, calls dataset's `__getitem__` to get actual data
|
|
||||||
- Batch data shape is `[batch_size, window_size]` (or varies according to specific dataset type)
|
|
||||||
|
|
||||||
4. **Strategy Forward and Loss Calculation**
|
|
||||||
- Batch data is passed to strategy (such as `SEQStrategy`)
|
|
||||||
- Strategy internally calls `Transformer` model, obtaining logits
|
|
||||||
- Calculate cross-entropy loss (or DPO loss, etc.) according to task type
|
|
||||||
- Return loss tensor
|
|
||||||
|
|
||||||
5. **Backpropagation and Optimization**
|
|
||||||
- Loss is normalized by dividing by accumulation steps, then `loss.backward()` is executed
|
|
||||||
- After accumulating `accumulation_steps` batches, optimizer `step()` and `zero_grad()` are executed
|
|
||||||
- Learning rate scheduler updates learning rate after each step
|
|
||||||
|
|
||||||
6. **Checkpoint Saving**
|
|
||||||
- `CheckpointCallback` saves checkpoints at set intervals
|
|
||||||
- Checkpoints contain model state dict, current epoch, iteration, and other metadata
|
|
||||||
- Saved in safetensors format, ensuring safety and efficiency
|
|
||||||
|
|
||||||
## Inference Data Flow - Detailed Steps
|
|
||||||
|
|
||||||
1. **Model Loading**
|
|
||||||
- Load `Transformer` model from checkpoint via `AutoModel.from_pretrained()`
|
|
||||||
- Set model to evaluation mode (`model.eval()`), enable inference mode (`torch.inference_mode`)
|
|
||||||
|
|
||||||
2. **Prompt Construction and Encoding**
|
|
||||||
- User messages (list of dict with role and content) are converted to ChatML format string through `apply_chat_template` method in tokenizer
|
|
||||||
- Tokenizer encodes prompt string to token ID sequence `input_ids`
|
|
||||||
- For batch generation, use `pad_sequence` for padding
|
|
||||||
|
|
||||||
3. **Autoregressive Generation Loop**
|
|
||||||
- Scheduler allocates pages via `PagedCache.alloc_n()` for each task's prompt
|
|
||||||
- Prefill phase: runs full prompt through model with `PagedCache.bind()` to fill initial KV cache pages
|
|
||||||
- Decode phase: loops until generating `max_len` tokens or encountering stop token:
|
|
||||||
- Input last token ID to model, obtain `logits`
|
|
||||||
- Apply `sample()` (temperature, top-k, top-p) to `logits`
|
|
||||||
- Sample next token ID from the processed distribution
|
|
||||||
- Write new KV entries into paged cache; allocate additional pages as needed
|
|
||||||
- For streaming generation, yield each token to caller immediately via `stream_callback`
|
|
||||||
|
|
||||||
4. **Decoding and Output**
|
|
||||||
- Decode generated token ID sequence to text through tokenizer
|
|
||||||
- Remove special tokens, return plain text response
|
|
||||||
|
|
||||||
## Checkpoint and Serialization
|
|
||||||
|
|
||||||
- **Training Checkpoint**: Saves model parameters, optimizer state, scheduler state, current epoch and iteration
|
|
||||||
- **Model Parameters**: Supports safetensors format, automatically handles special logic like weight tying during loading
|
|
||||||
- **Dataset Serialization**: HDF5 format supports efficient random access and shared memory, suitable for large-scale pre-training data
|
|
||||||
|
|
||||||
## Summary
|
|
||||||
|
|
||||||
The data flow design of AstrAI reflects the characteristics of modularity, extensibility, and resumability. The training data flow supports large-scale distributed training through chunk loading, resumable sampling, gradient accumulation, and other mechanisms; the inference data flow achieves efficient text generation using paged KV cache, continuous batching, and composable sampling strategies. Clear interfaces between modules facilitate customization and extension.
|
|
||||||
|
|
||||||
> Document Update Time: 2026-04-09
|
|
||||||
|
|||||||
@@ -1,736 +0,0 @@
|
|||||||
## 1. Why I Created This Project
|
|
||||||
|
|
||||||
There are many large language models on the market today, such as GPT, LLaMA, and others, with tens of billions or even hundreds of billions of parameters. But honestly, these models have extremely high hardware requirements, making them inaccessible for ordinary developers. I thought: **Can we create a model that is both useful and can run on ordinary computers?** This is also what most people currently hope for - a locally deployable AI project that achieves complete privatization while maintaining some level of intelligence.
|
|
||||||
|
|
||||||
Thus, the AstrAI project was born - 1B parameters, Chinese-English bilingual, supporting dialogue, text generation, and the training code is open source!
|
|
||||||
|
|
||||||
## 2. System Architecture
|
|
||||||
|
|
||||||
```mermaid
|
|
||||||
classDiagram
|
|
||||||
namespace config {
|
|
||||||
class ModelConfig {
|
|
||||||
+int vocab_size
|
|
||||||
+int dim
|
|
||||||
+int n_layers
|
|
||||||
+float norm_eps
|
|
||||||
+int dim_ffn
|
|
||||||
+bool tie_weight
|
|
||||||
+int max_len
|
|
||||||
+float rope_theta
|
|
||||||
+int n_heads
|
|
||||||
+int n_kv_heads
|
|
||||||
+bool use_qk_norm
|
|
||||||
+bool use_gated_attention
|
|
||||||
+load(config_path) ModelConfig
|
|
||||||
+save(config_path)
|
|
||||||
}
|
|
||||||
|
|
||||||
class TrainConfig {
|
|
||||||
+nn.Module model
|
|
||||||
+str strategy
|
|
||||||
+Dataset dataset
|
|
||||||
+Callable optimizer_fn
|
|
||||||
+Callable scheduler_fn
|
|
||||||
+int n_epoch
|
|
||||||
+int batch_size
|
|
||||||
+int accumulation_steps
|
|
||||||
+float max_grad_norm
|
|
||||||
+int start_epoch
|
|
||||||
+int start_batch
|
|
||||||
+str ckpt_dir
|
|
||||||
+int ckpt_interval
|
|
||||||
+int random_seed
|
|
||||||
+int num_workers
|
|
||||||
+int prefetch_factor
|
|
||||||
+bool pin_memory
|
|
||||||
+int nprocs
|
|
||||||
+str backend
|
|
||||||
+str master_addr
|
|
||||||
+str master_port
|
|
||||||
+Callable parallel_wrapper
|
|
||||||
+Callable state_dict_fn
|
|
||||||
+List[int] device_ids
|
|
||||||
+str device_type
|
|
||||||
+dict extra_kwargs
|
|
||||||
+validate()
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
namespace dataset {
|
|
||||||
class BaseDataset {
|
|
||||||
+int window_size
|
|
||||||
+int stride
|
|
||||||
+MultiSegmentFetcher fetcher
|
|
||||||
+load(load_path)
|
|
||||||
+__getitem__(index)
|
|
||||||
+__len__()
|
|
||||||
}
|
|
||||||
|
|
||||||
class SEQDataset {
|
|
||||||
+__getitem__(index) Dict
|
|
||||||
}
|
|
||||||
|
|
||||||
class SFTDataset {
|
|
||||||
+__getitem__(index) Dict
|
|
||||||
}
|
|
||||||
|
|
||||||
class DPODataset {
|
|
||||||
+__getitem__(index) Dict
|
|
||||||
}
|
|
||||||
|
|
||||||
class GRPODataset {
|
|
||||||
+__getitem__(index) Dict
|
|
||||||
}
|
|
||||||
|
|
||||||
class BaseSegmentFetcher {
|
|
||||||
+List[Tensor] segments
|
|
||||||
+List[int] cum_lengths
|
|
||||||
+int total_length
|
|
||||||
+fetch_data(begin_idx, end_idx) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class MultiSegmentFetcher {
|
|
||||||
+Dict multi_fetchers
|
|
||||||
+List multi_keys
|
|
||||||
+key_fetch(begin_idx, end_idx, keys) Dict
|
|
||||||
+fetch_data(begin_idx, end_idx) Dict
|
|
||||||
}
|
|
||||||
|
|
||||||
class ResumableDistributedSampler {
|
|
||||||
+int start_epoch
|
|
||||||
+int start_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
|
|
||||||
+Dict _registry
|
|
||||||
+register(model_type) decorator
|
|
||||||
+get_model_class(model_type) Type
|
|
||||||
+from_pretrained(path, disable_random_init) nn.Module
|
|
||||||
+save_pretrained(save_directory)
|
|
||||||
+to(*args, **kwargs) Self
|
|
||||||
}
|
|
||||||
|
|
||||||
class Transformer {
|
|
||||||
+ModelConfig config
|
|
||||||
+RotaryEmbedding rotary_embedding
|
|
||||||
+Embedding embed_tokens
|
|
||||||
+ModuleList layers
|
|
||||||
+RMSNorm norm
|
|
||||||
+Linear lm_head
|
|
||||||
+forward(input_ids, input_mask, persistent_key_values, start_pos) Dict
|
|
||||||
+load_state_dict(state_dict)
|
|
||||||
+state_dict()
|
|
||||||
}
|
|
||||||
|
|
||||||
class DecoderBlock {
|
|
||||||
+GQA attention
|
|
||||||
+RMSNorm input_norm
|
|
||||||
+MLP mlp
|
|
||||||
+RMSNorm post_attention_norm
|
|
||||||
+forward(x, rotary_emb, attention_mask, kv_cache, start_pos) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class GQA {
|
|
||||||
+int n_heads
|
|
||||||
+int n_kv_heads
|
|
||||||
+int head_dim
|
|
||||||
+Linear q_proj, k_proj, v_proj, o_proj
|
|
||||||
+RMSNorm q_norm, k_norm
|
|
||||||
+forward(x, rotary_emb, mask, kv_cache, start_pos) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class MLA {
|
|
||||||
+int n_heads
|
|
||||||
+int n_kv_heads
|
|
||||||
+int head_dim
|
|
||||||
+Linear q_a_proj, q_b_proj, q_c_proj
|
|
||||||
+Linear kv_a_proj, kv_b_proj, kv_c_proj
|
|
||||||
+Linear o_proj
|
|
||||||
+RMSNorm q_norm, k_norm
|
|
||||||
+forward(x, rotary_emb, mask, kv_cache, start_pos) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class MLP {
|
|
||||||
+Linear up, gate, down
|
|
||||||
+forward(x) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class RMSNorm {
|
|
||||||
+Parameter weight
|
|
||||||
+float norm_eps
|
|
||||||
+forward(x) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class Linear {
|
|
||||||
+Parameter weight
|
|
||||||
+Parameter bias
|
|
||||||
+forward(x) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class RotaryEmbedding {
|
|
||||||
+int dim
|
|
||||||
+int max_len
|
|
||||||
+float base
|
|
||||||
+forward(x, start_pos) Tuple[Tensor, Tensor]
|
|
||||||
}
|
|
||||||
|
|
||||||
class Embedding {
|
|
||||||
+Parameter weight
|
|
||||||
+forward(x) Tensor
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
namespace tokenize {
|
|
||||||
class AutoTokenizer {
|
|
||||||
+List[str] stop_ids
|
|
||||||
+int bos_id
|
|
||||||
+int eos_id
|
|
||||||
+int pad_id
|
|
||||||
+vocab_size int
|
|
||||||
+encode(tokens, out_ids, add_special_tokens) List[int]
|
|
||||||
+decode(tokens, skip_special_tokens) str
|
|
||||||
+apply_chat_template(messages, tokenize) Union[str, List[int]]
|
|
||||||
+set_chat_template(template)
|
|
||||||
+load(path)
|
|
||||||
+from_pretrained(path) AutoTokenizer
|
|
||||||
+save_pretrained(save_path)
|
|
||||||
}
|
|
||||||
|
|
||||||
class ChatTemplate {
|
|
||||||
+String template_str
|
|
||||||
+render(messages, add_generation_prompt) 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
|
|
||||||
+with_dataloader() TrainContextBuilder
|
|
||||||
+with_strategy() TrainContextBuilder
|
|
||||||
+build() TrainContext
|
|
||||||
}
|
|
||||||
|
|
||||||
class BaseStrategy {
|
|
||||||
+nn.Module model
|
|
||||||
+str device
|
|
||||||
+compute_loss(batch) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class StrategyFactory {
|
|
||||||
+Registry _registry
|
|
||||||
+register(name) decorator
|
|
||||||
+create(model, train_type, device, **kwargs) BaseStrategy
|
|
||||||
}
|
|
||||||
|
|
||||||
class SEQStrategy {
|
|
||||||
+float label_smoothing
|
|
||||||
+compute_loss(batch) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class SFTStrategy {
|
|
||||||
+float label_smoothing
|
|
||||||
+compute_loss(batch) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class DPOStrategy {
|
|
||||||
+nn.Module ref_model
|
|
||||||
+float beta
|
|
||||||
+str reduction
|
|
||||||
+compute_loss(batch) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class GRPOStrategy {
|
|
||||||
+nn.Module ref_model
|
|
||||||
+float clip_eps
|
|
||||||
+float kl_coef
|
|
||||||
+int group_size
|
|
||||||
+compute_loss(batch) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class BaseScheduler {
|
|
||||||
+get_lr() List[float]
|
|
||||||
+step()
|
|
||||||
}
|
|
||||||
|
|
||||||
class SchedulerFactory {
|
|
||||||
+Registry _registry
|
|
||||||
+register(name) decorator
|
|
||||||
+create(optimizer, schedule_type, **kwargs) BaseScheduler
|
|
||||||
}
|
|
||||||
|
|
||||||
class CosineScheduler {
|
|
||||||
+int warmup_steps
|
|
||||||
+int lr_decay_steps
|
|
||||||
+float min_rate
|
|
||||||
}
|
|
||||||
|
|
||||||
class SGDRScheduler {
|
|
||||||
+int warmup_steps
|
|
||||||
+int cycle_length
|
|
||||||
+float min_rate
|
|
||||||
+int t_mult
|
|
||||||
}
|
|
||||||
|
|
||||||
class TrainCallback {
|
|
||||||
+on_train_begin(context)
|
|
||||||
+on_train_end(context)
|
|
||||||
+on_epoch_begin(context)
|
|
||||||
+on_epoch_end(context)
|
|
||||||
+on_step_begin(context)
|
|
||||||
+on_step_end(context)
|
|
||||||
+on_batch_begin(context)
|
|
||||||
+on_batch_end(context)
|
|
||||||
+on_error(context)
|
|
||||||
}
|
|
||||||
|
|
||||||
class GradientClippingCallback {
|
|
||||||
+float max_grad_norm
|
|
||||||
+on_step_begin(context)
|
|
||||||
}
|
|
||||||
|
|
||||||
class SchedulerCallback {
|
|
||||||
+on_train_begin(context)
|
|
||||||
+on_batch_end(context)
|
|
||||||
}
|
|
||||||
|
|
||||||
class CheckpointCallback {
|
|
||||||
+str save_dir
|
|
||||||
+int interval
|
|
||||||
+_save_checkpoint(context)
|
|
||||||
+on_batch_end(context)
|
|
||||||
+on_train_end(context)
|
|
||||||
+on_error(context)
|
|
||||||
}
|
|
||||||
|
|
||||||
class ProgressBarCallback {
|
|
||||||
+int num_epoch
|
|
||||||
+on_epoch_begin(context)
|
|
||||||
+on_batch_end(context)
|
|
||||||
+on_epoch_end(context)
|
|
||||||
}
|
|
||||||
|
|
||||||
class MetricLoggerCallback {
|
|
||||||
+str log_dir
|
|
||||||
+int save_interval
|
|
||||||
+on_batch_end(context)
|
|
||||||
+on_train_end(context)
|
|
||||||
}
|
|
||||||
|
|
||||||
class CallbackFactory {
|
|
||||||
+Registry _registry
|
|
||||||
+register(name) decorator
|
|
||||||
+create(name, **kwargs) TrainCallback
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
namespace inference {
|
|
||||||
class InferenceEngine {
|
|
||||||
+nn.Module model
|
|
||||||
+AutoTokenizer tokenizer
|
|
||||||
+InferenceScheduler scheduler
|
|
||||||
+int max_batch_size
|
|
||||||
+Optional int max_seq_len
|
|
||||||
+generate(prompt, stream, max_tokens, temperature, top_p, top_k) Union[Generator, str, List[str]]
|
|
||||||
+generate_with_request(request) Union[Generator, str, List[str]]
|
|
||||||
+generate_async(prompt, max_tokens, temperature, top_p, top_k) AsyncGenerator
|
|
||||||
+get_stats() Dict
|
|
||||||
+shutdown()
|
|
||||||
}
|
|
||||||
|
|
||||||
class InferenceScheduler {
|
|
||||||
+nn.Module model
|
|
||||||
+AutoTokenizer tokenizer
|
|
||||||
+PagedCache page_cache
|
|
||||||
+int max_batch_size
|
|
||||||
+int max_seq_len
|
|
||||||
+int max_prompt_len
|
|
||||||
+int page_size
|
|
||||||
+List waiting_queue
|
|
||||||
+List active_tasks
|
|
||||||
+add_task(prompt, max_tokens, temperature, top_p, top_k, stream_callback) str
|
|
||||||
+remove_task(task_id)
|
|
||||||
+start()
|
|
||||||
+stop()
|
|
||||||
+get_stats() Dict
|
|
||||||
}
|
|
||||||
|
|
||||||
class PagedCache {
|
|
||||||
+int page_size
|
|
||||||
+int _free_mask
|
|
||||||
+List[int] _refs
|
|
||||||
+Tensor k_cache
|
|
||||||
+Tensor v_cache
|
|
||||||
+alloc() int
|
|
||||||
+alloc_n(n) List[int]
|
|
||||||
+free(idx)
|
|
||||||
+bind(page_table, total_len) CacheView
|
|
||||||
+write(layer_id, page_table, start_pos, k, v)
|
|
||||||
+gather(layer_id, page_table) Tuple[Tensor, Tensor]
|
|
||||||
}
|
|
||||||
|
|
||||||
class CacheView {
|
|
||||||
+PagedCache _cache
|
|
||||||
+Tensor _page_table
|
|
||||||
+int _total_len
|
|
||||||
+write(layer_id, start_pos, k, v)
|
|
||||||
+gather(layer_id) Tuple[Tensor, Tensor]
|
|
||||||
}
|
|
||||||
|
|
||||||
class Task {
|
|
||||||
+str task_id
|
|
||||||
+List prompt_ids
|
|
||||||
+int max_tokens
|
|
||||||
+float temperature
|
|
||||||
+float top_p
|
|
||||||
+int top_k
|
|
||||||
+TaskStatus status
|
|
||||||
+List output_ids
|
|
||||||
+int input_tokens
|
|
||||||
+int output_tokens
|
|
||||||
+List[int] page_table
|
|
||||||
+int n_pages
|
|
||||||
+float arrival_time
|
|
||||||
+float finish_time
|
|
||||||
+Callable stream_callback
|
|
||||||
+next_pos() int
|
|
||||||
+is_finished(stop_ids) bool
|
|
||||||
}
|
|
||||||
|
|
||||||
class TaskStatus {
|
|
||||||
<<enumeration>>
|
|
||||||
PENDING
|
|
||||||
RUNNING
|
|
||||||
FINISHED
|
|
||||||
ABORTED
|
|
||||||
}
|
|
||||||
|
|
||||||
class GenerationRequest {
|
|
||||||
+List[Dict] messages
|
|
||||||
+GenerationParams params
|
|
||||||
+bool stream
|
|
||||||
}
|
|
||||||
|
|
||||||
class GenerationParams {
|
|
||||||
<<value object>>
|
|
||||||
+int top_k
|
|
||||||
+float top_p
|
|
||||||
+float temperature
|
|
||||||
+int max_tokens
|
|
||||||
}
|
|
||||||
|
|
||||||
class BaseSamplingStrategy {
|
|
||||||
<<abstract>>
|
|
||||||
+apply(logits, filter_value) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class TemperatureStrategy {
|
|
||||||
+float temperature
|
|
||||||
+apply(logits, filter_value) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class TopKStrategy {
|
|
||||||
+int top_k
|
|
||||||
+apply(logits, filter_value) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class TopPStrategy {
|
|
||||||
+float top_p
|
|
||||||
+apply(logits, filter_value) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class SamplingPipeline {
|
|
||||||
+List strategies
|
|
||||||
+apply(logits, filter_value) Tensor
|
|
||||||
+sample(logits, filter_value) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class Server {
|
|
||||||
+start()
|
|
||||||
+predict(request)
|
|
||||||
}
|
|
||||||
|
|
||||||
class _Result {
|
|
||||||
+List[str] tokens
|
|
||||||
+List[str] results
|
|
||||||
+List[bool] done_flags
|
|
||||||
+append(token, idx)
|
|
||||||
+get_results() List[str]
|
|
||||||
+pop_all() List[str]
|
|
||||||
+wait(timeout) bool
|
|
||||||
}
|
|
||||||
|
|
||||||
class ChatMessage {
|
|
||||||
+str role
|
|
||||||
+str content
|
|
||||||
}
|
|
||||||
|
|
||||||
class ChatCompletionRequest {
|
|
||||||
+List[ChatMessage] messages
|
|
||||||
+float temperature
|
|
||||||
+float top_p
|
|
||||||
+int top_k
|
|
||||||
+int max_tokens
|
|
||||||
+bool stream
|
|
||||||
+Optional[str] stop
|
|
||||||
+Optional[int] n
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
namespace parallel {
|
|
||||||
class ParallelSetup {
|
|
||||||
+spawn_parallel_fn(fn, nprocs)
|
|
||||||
+setup_parallel(rank, world_size, backend, master_addr, master_port, device_type, device_ids)
|
|
||||||
}
|
|
||||||
|
|
||||||
class ParallelModel {
|
|
||||||
+dist.ProcessGroup process_group
|
|
||||||
+int rank
|
|
||||||
+int world_size
|
|
||||||
}
|
|
||||||
|
|
||||||
class ColumnParallelLinear {
|
|
||||||
+forward(x) Tensor
|
|
||||||
}
|
|
||||||
|
|
||||||
class RowParallelLinear {
|
|
||||||
+forward(x) Tensor
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
%% Relationships
|
|
||||||
TrainConfig --> ModelConfig : uses
|
|
||||||
TrainConfig --> BaseDataset : uses
|
|
||||||
TrainConfig --> StrategyFactory : selects
|
|
||||||
StrategyFactory ..> BaseStrategy : creates
|
|
||||||
BaseStrategy <|-- SEQStrategy
|
|
||||||
BaseStrategy <|-- SFTStrategy
|
|
||||||
BaseStrategy <|-- DPOStrategy
|
|
||||||
BaseStrategy <|-- GRPOStrategy
|
|
||||||
DPOStrategy --> Transformer : uses
|
|
||||||
GRPOStrategy --> Transformer : uses
|
|
||||||
Trainer --> TrainConfig : configures
|
|
||||||
Trainer --> TrainContextBuilder : builds
|
|
||||||
Trainer --> TrainCallback : manages
|
|
||||||
TrainContextBuilder --> TrainContext : creates
|
|
||||||
Checkpoint ..> Checkpoint : saves/loads
|
|
||||||
TrainContext --> Checkpoint : manages
|
|
||||||
TrainContext --> BaseStrategy : uses
|
|
||||||
TrainContext --> BaseScheduler : uses
|
|
||||||
SchedulerFactory ..> BaseScheduler : creates
|
|
||||||
BaseScheduler <|-- CosineScheduler
|
|
||||||
BaseScheduler <|-- SGDRScheduler
|
|
||||||
CallbackFactory ..> TrainCallback : creates
|
|
||||||
TrainCallback <|-- GradientClippingCallback
|
|
||||||
TrainCallback <|-- SchedulerCallback
|
|
||||||
TrainCallback <|-- CheckpointCallback
|
|
||||||
TrainCallback <|-- ProgressBarCallback
|
|
||||||
TrainCallback <|-- MetricLoggerCallback
|
|
||||||
InferenceEngine --> InferenceScheduler : uses
|
|
||||||
InferenceEngine --> GenerationRequest : uses
|
|
||||||
GenerationRequest --> GenerationParams : contains
|
|
||||||
InferenceScheduler --> Task : manages
|
|
||||||
Task --> TaskStatus : uses
|
|
||||||
InferenceScheduler --> TaskStatus : uses
|
|
||||||
InferenceScheduler --> PagedCache : uses
|
|
||||||
InferenceScheduler --> Transformer : uses
|
|
||||||
InferenceEngine --> Transformer : uses
|
|
||||||
InferenceEngine --> _Result : uses
|
|
||||||
BaseSamplingStrategy <|-- TemperatureStrategy
|
|
||||||
BaseSamplingStrategy <|-- TopKStrategy
|
|
||||||
BaseSamplingStrategy <|-- TopPStrategy
|
|
||||||
SamplingPipeline --> BaseSamplingStrategy : composes
|
|
||||||
Server --> InferenceEngine : uses
|
|
||||||
Server --> ChatMessage : uses
|
|
||||||
Server --> ChatCompletionRequest : uses
|
|
||||||
ParallelSetup --> Trainer : enables
|
|
||||||
BaseDataset <|-- SEQDataset
|
|
||||||
BaseDataset <|-- SFTDataset
|
|
||||||
BaseDataset <|-- DPODataset
|
|
||||||
BaseDataset <|-- GRPODataset
|
|
||||||
DatasetFactory ..> BaseDataset : creates
|
|
||||||
BaseSegmentFetcher --> MultiSegmentFetcher : used by
|
|
||||||
MultiSegmentFetcher --> BaseDataset : used by
|
|
||||||
AutoModel <|-- Transformer
|
|
||||||
AutoModel --> ModelConfig : contains
|
|
||||||
Transformer --> DecoderBlock : uses
|
|
||||||
Transformer --> RotaryEmbedding : uses
|
|
||||||
Transformer --> Embedding : uses
|
|
||||||
DecoderBlock --> GQA : uses
|
|
||||||
DecoderBlock --> MLA : uses
|
|
||||||
DecoderBlock --> MLP : uses
|
|
||||||
DecoderBlock --> RMSNorm : uses
|
|
||||||
TrainContextBuilder --> ResumableDistributedSampler : creates
|
|
||||||
ResumableDistributedSampler --> BaseDataset : samples
|
|
||||||
ParallelModel <|-- RowParallelLinear
|
|
||||||
ParallelModel <|-- ColumnParallelLinear
|
|
||||||
AutoTokenizer --> ChatTemplate : uses
|
|
||||||
TrainConfig --> DatasetFactory : selects
|
|
||||||
TrainConfig --> SchedulerFactory : selects
|
|
||||||
TrainConfig --> CallbackFactory : selects
|
|
||||||
AutoModel ..> AutoTokenizer : loads with
|
|
||||||
BaseFactory <|-- DatasetFactory
|
|
||||||
BaseFactory <|-- StrategyFactory
|
|
||||||
BaseFactory <|-- SchedulerFactory
|
|
||||||
BaseFactory <|-- CallbackFactory
|
|
||||||
```
|
|
||||||
|
|
||||||
### Module Overview
|
|
||||||
|
|
||||||
| Module | Components | Description |
|
|
||||||
|--------|------------|-------------|
|
|
||||||
| **astrai.config** | ModelConfig, TrainConfig | Configuration management |
|
|
||||||
| **astrai.dataset** | BaseDataset, SEQDataset, SFTDataset, DPODataset, GRPODataset, BaseSegmentFetcher, MultiSegmentFetcher, ResumableDistributedSampler, DatasetFactory | Dataset loading and management |
|
|
||||||
| **astrai.serialization** | Checkpoint, save_h5, load_h5 | Model serialization and checkpoint management |
|
|
||||||
| **astrai.model** | AutoModel, Transformer, DecoderBlock, GQA, MLA, MLP, RMSNorm, Linear, RotaryEmbedding, Embedding | Neural network model |
|
|
||||||
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
|
|
||||||
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy, StrategyFactory, BaseScheduler, SchedulerFactory, TrainCallback, CallbackFactory | Training workflow management |
|
|
||||||
| **astrai.inference** | InferenceEngine, InferenceScheduler, PagedCache, CacheView, Task, TaskStatus, GenerationParams, GenerationRequest, BaseSamplingStrategy, TemperatureStrategy, TopKStrategy, TopPStrategy, SamplingPipeline, ChatMessage, ChatCompletionRequest | Inference service with continuous batching and paged KV cache |
|
|
||||||
| **astrai.parallel** | ParallelSetup, 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) |
|
|
||||||
| **Singleton** | `TrainContext` | Training process global state management |
|
|
||||||
| **Registry** | `BaseFactory`, `Registry` | Generic component registration with category and priority support |
|
|
||||||
| **Object Pool** | `PagedCache` | Page-based KV cache with O(1) alloc/free via bitmask |
|
|
||||||
| **Strategy (Sampling)** | `BaseSamplingStrategy`, `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations with temperature, top-k, top-p |
|
|
||||||
| **Producer-Consumer** | `InferenceScheduler`, `Task`, `waiting_queue`, `active_tasks` | Continuous batching with dynamic task queue management |
|
|
||||||
| **Event-Driven** | `threading.Event`, `_task_event` | Non-blocking wait mechanism for task scheduling using Python's `threading` module |
|
|
||||||
| **AutoModel Registry** | `AutoModel`, `Transformer` | Model type registration and dynamic loading via decorator pattern |
|
|
||||||
| **Generator Pattern** | `_Result`, `GenerationRequest` | Event-based result notification for streaming/non-streaming generation |
|
|
||||||
|
|
||||||
### Core Relationships
|
|
||||||
|
|
||||||
1. **Configuration → Training**: `TrainConfig` contains `ModelConfig`, holds model, dataset, optimizer and other references
|
|
||||||
2. **Training Flow**: `Trainer` → `TrainContextBuilder` → `TrainContext`, uses `BaseStrategy` to compute loss
|
|
||||||
3. **Strategy Selection**: `StrategyFactory` creates corresponding strategy instance based on `train_type`
|
|
||||||
4. **Inference Flow**: `Server` → `InferenceEngine` → `InferenceScheduler` → `Transformer`, uses `PagedCache` for paged KV cache management and `SamplingPipeline` for efficient continuous batching with streaming/non-streaming
|
|
||||||
5. **Distributed Support**: `ParallelSetup` provides 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}
|
|
||||||
$$
|
|
||||||
|
|
||||||
In this implementation, an off-policy approach is used ($\pi_\theta = \pi_{\text{ref}}$), and the policy loss simplifies to:
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{policy}} = -\mathbb{E}[A]
|
|
||||||
$$
|
|
||||||
|
|
||||||
The KL divergence term uses mean squared error approximation:
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{KL} = \lambda \cdot \mathbb{E} \left[ (\log \pi_\theta - \log \pi_{\text{ref}})^2 \right]
|
|
||||||
$$
|
|
||||||
|
|
||||||
The final loss is the sum of both: $L = L_{\text{policy}} + L_{KL}$
|
|
||||||
|
|
||||||
Through the above three-stage progressive training, the model completes its evolution from a general language foundation to a specialized, highly-aligned dialogue intelligence.
|
|
||||||
|
|
||||||
> Document Update Time: 2026-04-09
|
|
||||||
@@ -0,0 +1,140 @@
|
|||||||
|
# Inference
|
||||||
|
|
||||||
|
## KV Cache
|
||||||
|
|
||||||
|
At decode time, only the last query token matters. All previous K/V are cached to avoid recomputation:
|
||||||
|
|
||||||
|
$$
|
||||||
|
o_n = \sum_j \text{softmax}\left(\frac{q_n k_j}{\sqrt{d_k}}\right) v_j
|
||||||
|
$$
|
||||||
|
|
||||||
|
RoPE is applied **before** KV cache write, not after — otherwise position encoding drift occurs.
|
||||||
|
|
||||||
|
## KVCache System
|
||||||
|
|
||||||
|
Six classes working together:
|
||||||
|
|
||||||
|
```
|
||||||
|
KVCache (facade)
|
||||||
|
├── Allocator bitmask-based page allocator + ref-count + LRU eviction
|
||||||
|
├── PrefixCache hash-based prefix matching (page_hash via rolling hash)
|
||||||
|
├── PagePool orchestrates Allocator + PrefixCache
|
||||||
|
├── TaskTable maps task_id → page_table + cached token count
|
||||||
|
├── Storage k_cache / v_cache tensors (n_layers × n_pages × page_size × n_kv_heads × head_dim)
|
||||||
|
└── KvcacheView bundles Storage + page_table + total_len for attention layers
|
||||||
|
```
|
||||||
|
|
||||||
|
`KVCache.bind(page_table, total_len)` returns a `KvcacheView` used by attention layers via `write()` / `gather()`.
|
||||||
|
|
||||||
|
## Continuous Batching
|
||||||
|
|
||||||
|
`InferenceScheduler` runs a daemon thread with a 4-phase loop:
|
||||||
|
|
||||||
|
```
|
||||||
|
1. Cleanup → Remove finished tasks, free KV pages
|
||||||
|
2. Refill → Pop from waiting_queue, task_alloc pages, activate
|
||||||
|
3. Prefill → Group by (prompt_len, start_pos), run full forward
|
||||||
|
4. Decode → Pick largest same-position group, single-token forward
|
||||||
|
```
|
||||||
|
|
||||||
|
## Sampling (Strategy Pattern)
|
||||||
|
|
||||||
|
```
|
||||||
|
BaseSamplingStrategy → TemperatureStrategy → TopKStrategy → TopPStrategy
|
||||||
|
```
|
||||||
|
|
||||||
|
`SamplingPipeline` composes them: Temperature → Top-K → Top-P → softmax → multinomial.
|
||||||
|
`sample()` is a convenience shortcut for one-shot usage.
|
||||||
|
|
||||||
|
## Protocol Handlers (Template Method)
|
||||||
|
|
||||||
|
```python
|
||||||
|
class ProtocolHandler(ABC):
|
||||||
|
def handle(self):
|
||||||
|
ctx = StreamContext(...)
|
||||||
|
agen = engine.generate_async(prompt, ...)
|
||||||
|
if stream: self._handle_stream(agen, ctx)
|
||||||
|
else: self._handle_non_stream(agen, ctx)
|
||||||
|
```
|
||||||
|
|
||||||
|
Subclass hooks: `build_prompt()`, `create_response_id()`, `format_stream_start/token/end()`, `format_non_stream_response()`.
|
||||||
|
|
||||||
|
`OpenAIHandler` → `/v1/chat/completions`, `AnthropicHandler` → `/v1/messages`.
|
||||||
|
|
||||||
|
## Engine & GenerateResult
|
||||||
|
|
||||||
|
```
|
||||||
|
InferenceEngine
|
||||||
|
├── generate(prompt, stream, ...) → str | List[str] | Generator
|
||||||
|
├── generate_with_request(req) → same
|
||||||
|
└── generate_async(prompt, ...) → AsyncGenerator
|
||||||
|
```
|
||||||
|
|
||||||
|
`GenerateResult` uses `Condition` for non-streaming (`wait_completion()`) and `Event` for streaming (`wait()`). Stream callback is `cb(token)`.
|
||||||
|
|
||||||
|
## HTTP API
|
||||||
|
|
||||||
|
```
|
||||||
|
POST /v1/chat/completions OpenAI
|
||||||
|
POST /v1/messages Anthropic
|
||||||
|
GET /health {"status":"ok","model_loaded":true}
|
||||||
|
GET /stats scheduler statistics
|
||||||
|
```
|
||||||
|
|
||||||
|
### OpenAI
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{"messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||||
|
```
|
||||||
|
|
||||||
|
Response:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"id": "chatcmpl-abc123",
|
||||||
|
"object": "chat.completion",
|
||||||
|
"choices": [{"message": {"role": "assistant", "content": "Hello!"}, "finish_reason": "stop"}],
|
||||||
|
"usage": {"prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Streaming SSE: `data: {"choices":[{"delta":{"role":"assistant"}}]}` → token chunks → `data: [DONE]`
|
||||||
|
|
||||||
|
### Anthropic
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{"model":"astrai","system":"You are helpful.","messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||||
|
```
|
||||||
|
|
||||||
|
Supports `stop_sequences` and streaming via `event: content_block_delta`.
|
||||||
|
|
||||||
|
### GenerationRequest Parameters
|
||||||
|
|
||||||
|
| Param | Type | Default | Description |
|
||||||
|
|-------|------|---------|-------------|
|
||||||
|
| `messages` | List[dict] | required | Chat messages (role, content) |
|
||||||
|
| `temperature` | float | 1.0 | Sampling temperature (0.0–2.0) |
|
||||||
|
| `top_p` | float | 1.0 | Nucleus threshold |
|
||||||
|
| `top_k` | int | 50 | Top-k count |
|
||||||
|
| `max_tokens` | int | None | Max generation length |
|
||||||
|
| `stream` | bool | False | Stream output |
|
||||||
|
|
||||||
|
## Engine API
|
||||||
|
|
||||||
|
```python
|
||||||
|
# Non-streaming
|
||||||
|
engine.generate("Hello", stream=False) # -> str
|
||||||
|
engine.generate(["A", "B"], stream=False) # -> List[str]
|
||||||
|
|
||||||
|
# Streaming
|
||||||
|
engine.generate("Hello", stream=True) # -> Generator[str]
|
||||||
|
engine.generate(["A", "B"], stream=True) # -> Generator[Tuple[int, str]]
|
||||||
|
|
||||||
|
# Async
|
||||||
|
await engine.generate_async("Hello", ...) # -> AsyncGenerator[str]
|
||||||
|
```
|
||||||
|
|
||||||
|
> Document Update Time: 2026-05-17
|
||||||
@@ -1,334 +0,0 @@
|
|||||||
## Model Introduction
|
|
||||||
|
|
||||||
### 1. Model Architecture
|
|
||||||
|
|
||||||
This model uses the Transformer architecture with GQA mechanism (q_head=24, kv_head=4), which saves KV cache memory compared to traditional MHA. The model is built by stacking 32 layers of Transformer blocks, with 1.0 billion parameters. Transformer is an autoregressive model that calculates the relationship between all previous tokens to obtain the probability distribution of the next token.
|
|
||||||
|
|
||||||
The model now uses the **AutoModel** base class for flexible loading and saving:
|
|
||||||
|
|
||||||
```python
|
|
||||||
from astrai.model import AutoModel
|
|
||||||
|
|
||||||
# Load model from checkpoint
|
|
||||||
model = AutoModel.from_pretrained("path/to/model")
|
|
||||||
|
|
||||||
# Save model to new directory
|
|
||||||
model.save_pretrained("path/to/save")
|
|
||||||
```
|
|
||||||
|
|
||||||
The Transformer model is registered via `@AutoModel.register('transformer')` decorator, allowing easy extension for new model types.
|
|
||||||
|
|
||||||
```mermaid
|
|
||||||
flowchart TB
|
|
||||||
subgraph Layers["Transformer Layers"]
|
|
||||||
direction TB
|
|
||||||
A[Input Embedding] --> B[Transformer Block\nLayer 1]
|
|
||||||
B --> C[Transformer Block\nLayer ...]
|
|
||||||
C --> D[Transformer Block\nLayer 32]
|
|
||||||
D --> E[RMSNorm]
|
|
||||||
E --> F[Linear]
|
|
||||||
F --> G[SoftMax]
|
|
||||||
end
|
|
||||||
|
|
||||||
subgraph TransformerBlock["Transformer Block"]
|
|
||||||
direction TB
|
|
||||||
H[x] --> I[RMSNorm]
|
|
||||||
I --> J[Linear → Q/K/V]
|
|
||||||
J --> K[Q]
|
|
||||||
J --> L[K]
|
|
||||||
J --> M[V]
|
|
||||||
K --> N[RoPE]
|
|
||||||
L --> O[RoPE]
|
|
||||||
N --> P["Q @ K^T / sqrt(d)"]
|
|
||||||
O --> P
|
|
||||||
P --> Q[Masked SoftMax]
|
|
||||||
Q --> R[S @ V]
|
|
||||||
M --> R
|
|
||||||
R --> S[Linear]
|
|
||||||
S --> T[+]
|
|
||||||
H --> T
|
|
||||||
T --> U[RMSNorm]
|
|
||||||
U --> V[Linear]
|
|
||||||
V --> W[SiLU]
|
|
||||||
V --> X[×]
|
|
||||||
W --> X
|
|
||||||
X --> Y[Linear]
|
|
||||||
Y --> Z[+]
|
|
||||||
T --> Z
|
|
||||||
Z --> AA[x']
|
|
||||||
end
|
|
||||||
|
|
||||||
classDef main fill:#e6f3ff,stroke:#0066cc;
|
|
||||||
classDef block fill:#fff2e6,stroke:#cc6600;
|
|
||||||
class Layers main;
|
|
||||||
class TransformerBlock block;
|
|
||||||
```
|
|
||||||
|
|
||||||
What is an autoregressive model? After splitting a sentence into tokens, the model predicts the probability distribution of the next token. This means the model calculates the probability of the next possible token and its corresponding probability based on the given context (the sequence of tokens that have already appeared).
|
|
||||||
|
|
||||||
#### 1. Autoregression
|
|
||||||
|
|
||||||
In autoregressive modeling, when a sentence is tokenized into a sequence of tokens, the model learns to predict what comes next. Given a sequence of tokens as input, the model calculates a probability distribution over all possible next tokens. This distribution tells us how likely each potential next token is, given the current context.
|
|
||||||
|
|
||||||
For instance, if the input sequence contains tokens representing a question, the model might predict that certain response tokens have higher probabilities than others. The sampling process then selects one token from this distribution—controlled by parameters like top_k, top_p, and temperature—to serve as the next token in the sequence.
|
|
||||||
|
|
||||||
Once a token is selected, it is appended to the input sequence, and the model repeats this process. The updated sequence is then fed back into the model to predict the next token. This iterative process continues until either a special end-of-sequence token is generated, or the maximum sequence length is reached. These control tokens are essential because without them, the model would continue generating tokens indefinitely, eventually exhausting available memory.
|
|
||||||
|
|
||||||
#### 2. Causal Mask
|
|
||||||
|
|
||||||
Transformers use attention mechanism. The input shape is generally [bsz, seq_len], and the output is [bsz, seq_len, n_dim]. To predict the next token, the model's input and output must be offset by one position. The target predicted by the model must be offset by one position, and during training we also use the offset-by-one method:
|
|
||||||
|
|
||||||
```
|
|
||||||
sequence : [[1, 2, 3, 4, 5, 6]]
|
|
||||||
input_ids: [[1, 2, 3, 4, 5]]
|
|
||||||
target_ids: [[2, 3, 4, 5, 6]]
|
|
||||||
```
|
|
||||||
|
|
||||||
The attention score calculation formula is:
|
|
||||||
|
|
||||||
$$ s_{ij} = softmax(\frac{q_i^Tk_j}{\sqrt{d_k}}) $$
|
|
||||||
$$ s_{ij} := s_{ij} + mask_{ij} $$
|
|
||||||
|
|
||||||
Here, the attention score represents the degree to which the model attends to the similarity between two tokens.
|
|
||||||
|
|
||||||
For decoder-only structure models, to prevent the model from "stealing" information from future positions, a mask needs to be added during attention calculation. We need to apply a mask before attention score calculation. This mask is typically a lower triangular matrix, and for a sequence of length n, its shape is [n, n]. Below is an example of how to create such a causal mask matrix for a sequence of length 5:
|
|
||||||
|
|
||||||
```
|
|
||||||
[[0, -inf, -inf, -inf, -inf],
|
|
||||||
[0, 0, -inf, -inf, -inf],
|
|
||||||
[0, 0, 0, -inf, -inf],
|
|
||||||
[0, 0, 0, 0, -inf],
|
|
||||||
[0, 0, 0, 0, 0]]
|
|
||||||
```
|
|
||||||
|
|
||||||
In this matrix, 0 represents positions that can be attended to, while -inf represents positions that should be masked (i.e., should not be attended to). Because this matrix ensures that after the softmax, the parts of the attention scores where $j > i$ change from `inf` to 0, meaning the model cannot see future information.
|
|
||||||
|
|
||||||
#### 3. Rotary Position Embedding
|
|
||||||
|
|
||||||
Rotary Position Embedding (RoPE) is a position encoding method designed to solve the problem of lacking direct modeling of sequence position information in Transformer models. Unlike traditional position encodings (such as sine and cosine function position encodings), RoPE embeds position information directly into the Query (Q) and Key (K) vectors, allowing the model to more naturally handle relative position relationships in sequences.
|
|
||||||
|
|
||||||
$$ q_i = R_i W_q x_i $$
|
|
||||||
$$ k_j = R_j W_k x_j $$
|
|
||||||
$$ q_i^T k_j = (R_i W_q x_i)^T( R_j W_k x_j) = x_i^T W_q^T R_{i-j} W_k x_j $$
|
|
||||||
|
|
||||||
The $R_{i-j}$ controls the attenuation of attention for different tokens at different relative distances. When the absolute value of $i - j$ is larger, the degree of attenuation is stronger. This approach allows the model to learn relative position relationships, enabling the model to scale and adapt to longer sequences.
|
|
||||||
|
|
||||||
## KV Cache Implementation
|
|
||||||
|
|
||||||
According to the attention calculation formula:
|
|
||||||
|
|
||||||
$$
|
|
||||||
\begin{align*}
|
|
||||||
o_i &= \sum_j s_{ij} v_{j} \newline
|
|
||||||
s_{ij} &= \text{softmax}\left( \frac{q_{i} k_{j}}{\sqrt{d_k}} \right)
|
|
||||||
\end{align*}
|
|
||||||
$$
|
|
||||||
|
|
||||||
Since the model is an autoregressive model, we only need to calculate for the last part of the sequence, meaning the index $i$ is fixed as the last element of the sequence, and we compute $o_{n}$:
|
|
||||||
|
|
||||||
$$
|
|
||||||
\begin{align*}
|
|
||||||
o_n &= \sum_j s_{j}v_{j} \newline
|
|
||||||
s_j &= \text{softmax}\left(\frac{q_n k_{j}}{\sqrt{d_k}} \right)
|
|
||||||
\end{align*}
|
|
||||||
$$
|
|
||||||
|
|
||||||
If we expand the expression:
|
|
||||||
|
|
||||||
$$
|
|
||||||
o_n = \sum_j \text{softmax}\left(\frac{q_n k_{j}}{\sqrt{d_k}}\right)v_{j}
|
|
||||||
$$
|
|
||||||
|
|
||||||
In the above expression, only k and v have length indices, while $q$ does not. Therefore, during the calculation process, the input of $q$ is fixed as the last token from the previous input, while $k$ and $v$ need to be cached for parts of different lengths. Also, when caching, note that position encoding calculation should be performed before KV cache computation, otherwise there will be position encoding calculation errors.
|
|
||||||
|
|
||||||
### 4. AutoModel Loading
|
|
||||||
|
|
||||||
The project now uses the **AutoModel** base class for flexible model loading and saving:
|
|
||||||
|
|
||||||
```python
|
|
||||||
from astrai.model import AutoModel
|
|
||||||
|
|
||||||
# Load model from checkpoint
|
|
||||||
model = AutoModel.from_pretrained("path/to/model")
|
|
||||||
|
|
||||||
# Save model to new directory
|
|
||||||
model.save_pretrained("path/to/save")
|
|
||||||
```
|
|
||||||
|
|
||||||
The Transformer model is registered via `@AutoModel.register('transformer')` decorator, allowing easy extension for new model types. The `from_pretrained` method automatically loads the `config.json` to determine the model type and uses safetensors format for weights.
|
|
||||||
|
|
||||||
### 5. Continuous Batching Inference
|
|
||||||
|
|
||||||
The inference engine supports **continuous batching** for efficient batch processing:
|
|
||||||
|
|
||||||
```python
|
|
||||||
from astrai.inference import InferenceEngine, GenerationRequest
|
|
||||||
|
|
||||||
# Create inference engine with continuous batching
|
|
||||||
engine = InferenceEngine(
|
|
||||||
model=model,
|
|
||||||
tokenizer=tokenizer,
|
|
||||||
max_batch_size=8,
|
|
||||||
max_seq_len=4096,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Use GenerationRequest with messages format
|
|
||||||
request = GenerationRequest(
|
|
||||||
messages=[
|
|
||||||
{"role": "system", "content": "You are a helpful assistant."},
|
|
||||||
{"role": "user", "content": "Hello"},
|
|
||||||
],
|
|
||||||
temperature=0.8,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50,
|
|
||||||
max_len=1024,
|
|
||||||
stream=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Generate with streaming
|
|
||||||
for token in engine.generate_with_request(request):
|
|
||||||
print(token, end="", flush=True)
|
|
||||||
```
|
|
||||||
|
|
||||||
The continuous batching feature allows dynamic batch composition where new requests can join at any time and completed requests are released immediately.
|
|
||||||
|
|
||||||
## HTTP API Usage
|
|
||||||
|
|
||||||
The inference server provides HTTP endpoints for remote inference. Start the server first:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python -m scripts.tools.server --port 8000
|
|
||||||
```
|
|
||||||
|
|
||||||
### OpenAI-Compatible Endpoint
|
|
||||||
|
|
||||||
The server provides an OpenAI-compatible chat completion endpoint at `/v1/chat/completions`:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
|
||||||
-H "Content-Type: application/json" \
|
|
||||||
-d '{
|
|
||||||
"messages": [
|
|
||||||
{"role": "system", "content": "You are a helpful assistant."},
|
|
||||||
{"role": "user", "content": "Hello, how are you?"}
|
|
||||||
],
|
|
||||||
"temperature": 0.8,
|
|
||||||
"max_tokens": 2048,
|
|
||||||
"stream": false
|
|
||||||
}'
|
|
||||||
```
|
|
||||||
|
|
||||||
**Request Parameters:**
|
|
||||||
| Parameter | Type | Default | Description |
|
|
||||||
|-----------|------|---------|-------------|
|
|
||||||
| `messages` | List[dict] | Required | Chat messages with role and content |
|
|
||||||
| `temperature` | float | 0.8 | Sampling temperature (0.0-2.0) |
|
|
||||||
| `top_p` | float | 0.95 | Nucleus sampling threshold |
|
|
||||||
| `top_k` | int | 50 | Top-k sampling parameter |
|
|
||||||
| `max_tokens` | int | 2048 | Maximum tokens to generate |
|
|
||||||
| `stream` | bool | false | Enable streaming response |
|
|
||||||
| `system_prompt` | str | None | System prompt override |
|
|
||||||
|
|
||||||
**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"
|
|
||||||
}
|
|
||||||
]
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
### 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`.
|
|
||||||
|
|
||||||
### Health Check
|
|
||||||
|
|
||||||
|
|
||||||
### 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, "engine_ready": true}
|
|
||||||
|
|
||||||
curl http://localhost:8000/stats
|
|
||||||
# {"requests_total": 10, "tokens_generated": 5000, ...}
|
|
||||||
```
|
|
||||||
|
|
||||||
> Document Update Time: 2026-04-09
|
|
||||||
+31
-88
@@ -6,18 +6,18 @@
|
|||||||
|
|
||||||
| Parameter | Description | Default |
|
| Parameter | Description | Default |
|
||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
| `--train_type` | Training type (`seq`, `sft`, `dpo`) | required |
|
| `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`) | required |
|
||||||
| `--data_root_path` | Dataset root directory | required |
|
| `--data_root_path` | Dataset root directory | required |
|
||||||
| `--param_path` | Model parameters or checkpoint path | required |
|
| `--param_path` | Model parameters or checkpoint path | required |
|
||||||
| `--n_epoch` | Total training epochs | 1 |
|
| `--n_epoch` | Total training epochs | 1 |
|
||||||
| `--batch_size` | Batch size | 1 |
|
| `--batch_per_device` | Batch size per device | 1 |
|
||||||
| `--accumulation_steps` | Gradient accumulation steps between optimizer steps | 1 |
|
| `--grad_accum_steps` | Gradient accumulation steps between optimizer steps | 1 |
|
||||||
|
|
||||||
### Learning Rate Scheduling
|
### Learning Rate Scheduling
|
||||||
|
|
||||||
| Parameter | Description | Default |
|
| Parameter | Description | Default |
|
||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
| `--warmup_steps` | Warmup steps | 1000 |
|
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
|
||||||
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
||||||
| `--max_grad_norm` | Maximum gradient norm for clipping | 1.0 |
|
| `--max_grad_norm` | Maximum gradient norm for clipping | 1.0 |
|
||||||
|
|
||||||
@@ -60,95 +60,38 @@
|
|||||||
| Parameter | Description | Default | Used by |
|
| Parameter | Description | Default | Used by |
|
||||||
|-----------|-------------|---------|---------|
|
|-----------|-------------|---------|---------|
|
||||||
| `--dpo_beta` | DPO beta value | 0.1 | `dpo` |
|
| `--dpo_beta` | DPO beta value | 0.1 | `dpo` |
|
||||||
| `--label_smoothing` | Label smoothing for cross-entropy loss | 0.1 | `seq`, `sft` |
|
| `--label_smoothing` | Label smoothing for cross-entropy loss | 0.05 | `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
|
### Usage Example
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python scripts/tools/train.py \
|
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||||
--train_type seq \
|
|
||||||
--data_root_path /path/to/dataset \
|
nohup python scripts/tools/train.py \
|
||||||
--param_path /path/to/model \
|
--nprocs=4 \
|
||||||
--n_epoch 3 \
|
--train_type=seq \
|
||||||
--batch_size 4 \
|
--data_root_path=/path/to/dataset \
|
||||||
--accumulation_steps 8 \
|
--param_path=/path/to/model \
|
||||||
--max_lr 3e-4 \
|
--batch_per_device=4 \
|
||||||
--warmup_steps 2000 \
|
--grad_accum_steps=8 \
|
||||||
--max_grad_norm 1.0 \
|
--warmup_ratio=0.05 \
|
||||||
--ckpt_interval 5000 \
|
--max_lr=1e-4 \
|
||||||
--ckpt_dir ./checkpoints \
|
--max_grad_norm=1.0 \
|
||||||
--num_workers 4 \
|
--adamw_beta1=0.9 \
|
||||||
--nprocs 1 \
|
--adamw_beta2=0.95 \
|
||||||
--device_type cuda
|
--adamw_weight_decay=0.01 \
|
||||||
|
--window_size=2048 \
|
||||||
|
--ckpt_interval=10000 \
|
||||||
|
--ckpt_dir=./checkpoint \
|
||||||
|
--random_seed=3407 \
|
||||||
|
--label_smoothing=0.05 \
|
||||||
|
> out.log 2> err.log &
|
||||||
```
|
```
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## Generation Parameters
|
> Document Update Time: 2026-05-17
|
||||||
|
|
||||||
### GenerationRequest Parameters
|
|
||||||
|
|
||||||
| Parameter | Description | Default Value |
|
|
||||||
|-----------|-------------|---------------|
|
|
||||||
| `messages` | List of message dictionaries (role, content) | required |
|
|
||||||
| `temperature` | Sampling temperature (higher = more random) | 1.0 |
|
|
||||||
| `top_p` | Nucleus sampling threshold | 1.0 |
|
|
||||||
| `top_k` | Top-k sampling count | 50 |
|
|
||||||
| `max_len` | Maximum generation length | 1024 |
|
|
||||||
| `stream` | Whether to stream output | False |
|
|
||||||
|
|
||||||
### Usage Example
|
|
||||||
|
|
||||||
```python
|
|
||||||
import torch
|
|
||||||
from astrai.model import AutoModel
|
|
||||||
from astrai.tokenize import AutoTokenizer
|
|
||||||
from astrai.inference import InferenceEngine, GenerationRequest
|
|
||||||
|
|
||||||
# Load model using AutoModel
|
|
||||||
model = AutoModel.from_pretrained("your_model_dir")
|
|
||||||
|
|
||||||
# Load tokenizer
|
|
||||||
tokenizer = AutoTokenizer.from_pretrained("your_model_dir")
|
|
||||||
|
|
||||||
# Create engine with separate model and tokenizer
|
|
||||||
engine = InferenceEngine(
|
|
||||||
model=model,
|
|
||||||
tokenizer=tokenizer,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Build request with messages format
|
|
||||||
request = GenerationRequest(
|
|
||||||
messages=[
|
|
||||||
{"role": "system", "content": "You are a helpful assistant."},
|
|
||||||
{"role": "user", "content": "Hello"},
|
|
||||||
],
|
|
||||||
temperature=0.8,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50,
|
|
||||||
max_len=1024,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Generate (streaming)
|
|
||||||
for token in engine.generate_with_request(request):
|
|
||||||
print(token, end="", flush=True)
|
|
||||||
|
|
||||||
# Or use simple generate interface
|
|
||||||
result = engine.generate(
|
|
||||||
prompt="Hello",
|
|
||||||
stream=False,
|
|
||||||
max_tokens=1024,
|
|
||||||
temperature=0.8,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50,
|
|
||||||
)
|
|
||||||
```
|
|
||||||
|
|
||||||
### Generation Modes
|
|
||||||
|
|
||||||
| Mode | Description |
|
|
||||||
|------|-------------|
|
|
||||||
| `stream=True` | Streaming output, yields token by token |
|
|
||||||
| `stream=False` | Non-streaming output, returns complete result |
|
|
||||||
|
|
||||||
> Document Update Time: 2026-04-09
|
|
||||||
@@ -0,0 +1,225 @@
|
|||||||
|
# Training
|
||||||
|
|
||||||
|
## Model Architecture
|
||||||
|
|
||||||
|
The model uses a decoder-only Transformer with **GQA** (Grouped Query Attention) and optional **MLA** (Multi-head Latent Attention). 1.0 billion parameters, Chinese–English bilingual.
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
flowchart TB
|
||||||
|
subgraph Layers["Transformer Layers"]
|
||||||
|
direction TB
|
||||||
|
A[Input Embedding] --> B[Transformer Block\nLayer 1]
|
||||||
|
B --> C[Transformer Block\nLayer ...]
|
||||||
|
C --> D[Transformer Block\nLayer ...]
|
||||||
|
D --> E[RMSNorm]
|
||||||
|
E --> F[Linear]
|
||||||
|
F --> G[SoftMax]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph TransformerBlock["Transformer Block"]
|
||||||
|
direction TB
|
||||||
|
H[x] --> I[RMSNorm]
|
||||||
|
I --> J[Linear → Q/K/V]
|
||||||
|
J --> K[Q]; J --> L[K]; J --> M[V]
|
||||||
|
K --> N[RoPE]; L --> O[RoPE]
|
||||||
|
N --> P["Q @ K^T / sqrt(d)"]; O --> P
|
||||||
|
P --> Q[Masked SoftMax]; Q --> R[S @ V]; M --> R
|
||||||
|
R --> S[Linear]; S --> T[+]; H --> T
|
||||||
|
T --> U[RMSNorm]
|
||||||
|
U --> V["Linear (gate)"]; U --> W["Linear (up)"]
|
||||||
|
V --> X[SiLU]; X --> Y[×]; W --> Y
|
||||||
|
Y --> Z["Linear (down)"]; Z --> AA[+]; T --> AA
|
||||||
|
AA --> BB[x']
|
||||||
|
end
|
||||||
|
```
|
||||||
|
|
||||||
|
### Autoregression
|
||||||
|
|
||||||
|
Given a token sequence, the model predicts the probability of the next token. Each generated token is appended to the input and fed back, repeating until an end-of-sequence token or max length.
|
||||||
|
|
||||||
|
### Causal Mask
|
||||||
|
|
||||||
|
```
|
||||||
|
sequence : [[1, 2, 3, 4, 5, 6]]
|
||||||
|
input_ids: [[1, 2, 3, 4, 5]]
|
||||||
|
target_ids: [[2, 3, 4, 5, 6]]
|
||||||
|
```
|
||||||
|
|
||||||
|
Lower-triangular mask prevents attending to future positions:
|
||||||
|
|
||||||
|
```
|
||||||
|
[[0, -inf, -inf, -inf, -inf],
|
||||||
|
[0, 0, -inf, -inf, -inf],
|
||||||
|
[0, 0, 0, -inf, -inf],
|
||||||
|
[0, 0, 0, 0, -inf],
|
||||||
|
[0, 0, 0, 0, 0]]
|
||||||
|
```
|
||||||
|
|
||||||
|
### Rotary Position Embedding (RoPE)
|
||||||
|
|
||||||
|
RoPE embeds position into Q/K vectors via complex rotation:
|
||||||
|
|
||||||
|
$$ q_i = R_i W_q x_i, \quad k_j = R_j W_k x_j, \quad q_i^T k_j = x_i^T W_q^T R_{i-j} W_k x_j $$
|
||||||
|
|
||||||
|
The complex rotation `freqs_cis` is pre-computed once (`cos, sin` pairs per position). `apply_rotary_emb` multiplies Q/K as complex numbers.
|
||||||
|
|
||||||
|
## Training Loop
|
||||||
|
|
||||||
|
Two-level loop: **epoch** → **batch**. Optimizer step fires every `grad_accum_steps` batches.
|
||||||
|
|
||||||
|
```
|
||||||
|
on_train_begin
|
||||||
|
on_epoch_begin
|
||||||
|
for batch in dataloader:
|
||||||
|
on_batch_begin
|
||||||
|
loss = strategy(batch)
|
||||||
|
(loss / grad_accum_steps).backward()
|
||||||
|
iteration += 1
|
||||||
|
on_batch_end
|
||||||
|
|
||||||
|
if iteration % grad_accum_steps == 0:
|
||||||
|
on_step_begin
|
||||||
|
optimizer.step()
|
||||||
|
optimizer.zero_grad()
|
||||||
|
on_step_end
|
||||||
|
scheduler.step()
|
||||||
|
on_epoch_end
|
||||||
|
on_train_end
|
||||||
|
```
|
||||||
|
|
||||||
|
### Callback Lifecycle
|
||||||
|
|
||||||
|
| Hook | Fires | Default callback |
|
||||||
|
|------|-------|-----------------|
|
||||||
|
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback` |
|
||||||
|
| `on_step_begin` | Every accumulation window | `GradientClippingCallback` |
|
||||||
|
| `on_batch_end` | Every batch | `CheckpointCallback`, `MetricLoggerCallback`, `ProgressBarCallback` |
|
||||||
|
| `on_step_end` | Every accumulation window | `ValidationCallback` |
|
||||||
|
| `on_train_end` | Training ends | `CheckpointCallback`, `MetricLoggerCallback` (final save) |
|
||||||
|
|
||||||
|
Default callbacks: `gradient_checkpointing` (activation checkpointing, optional), `progress_bar` (tqdm), `checkpoint` (safetensors, rank-0), `metric_logger` (JSONL, rank-0), `gradient_clipping`, `validation` (periodic validation on val_dataset).
|
||||||
|
|
||||||
|
## Strategies
|
||||||
|
|
||||||
|
### SEQ (Pre-training)
|
||||||
|
|
||||||
|
Next-token cross-entropy with optional label smoothing:
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{\text{PT}} = -\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta)
|
||||||
|
$$
|
||||||
|
|
||||||
|
Keys: `input_ids`, `target_ids`
|
||||||
|
|
||||||
|
### SFT (Supervised Fine-Tuning)
|
||||||
|
|
||||||
|
Masked cross-entropy (`ignore_index=-100`) over response tokens:
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{\text{SFT}} = -\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
|
||||||
|
$$
|
||||||
|
|
||||||
|
Keys: `input_ids`, `target_ids`, `loss_mask`
|
||||||
|
|
||||||
|
### DPO (Direct Preference Optimization)
|
||||||
|
|
||||||
|
Frozen reference model, preference margin via log-ratio:
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{\text{DPO}} = -\mathbb{E}\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]
|
||||||
|
$$
|
||||||
|
|
||||||
|
Parameters: `beta=0.1`. Keys: `chosen`, `rejected`, `chosen_mask`, `rejected_mask`.
|
||||||
|
|
||||||
|
### GRPO (Group Relative Policy Optimization)
|
||||||
|
|
||||||
|
On-policy PPO with group-normalized advantages:
|
||||||
|
|
||||||
|
$$
|
||||||
|
\text{Advantage}_i = \frac{r_i - \mu}{\sigma + \epsilon}
|
||||||
|
$$
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{\text{GRPO}} = -\mathbb{E}\left[\min\left(\frac{\pi_\theta}{\pi_{\text{ref}}}A,\; \text{clip}\left(\frac{\pi_\theta}{\pi_{\text{ref}}}, 1-\epsilon, 1+\epsilon\right)A\right)\right] + \lambda \cdot \mathbb{E}\left[(\log\pi_\theta - \log\pi_{\text{ref}})^2\right]
|
||||||
|
$$
|
||||||
|
|
||||||
|
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`, `sync_interval=200`.
|
||||||
|
|
||||||
|
Keys: `prompts`, `responses`, `masks`, `rewards`.
|
||||||
|
|
||||||
|
## LR Schedulers
|
||||||
|
|
||||||
|
| Type | Class | Description |
|
||||||
|
|------|-------|-------------|
|
||||||
|
| Cosine | `CosineScheduler` | Linear warmup → cosine decay to `min_rate` |
|
||||||
|
| SGDR | `SGDRScheduler` | Cosine annealing with warm restarts (`t_mult=2`) |
|
||||||
|
|
||||||
|
Created by `SchedulerFactory.create(optimizer, schedule_type, **kwargs)`.
|
||||||
|
|
||||||
|
## Gradient Checkpointing
|
||||||
|
|
||||||
|
Trades compute for memory by recomputing activations during backward pass. Specify module types via `gradient_checkpointing_modules`:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
|
config = TrainConfig(..., gradient_checkpointing_modules=[DecoderBlock])
|
||||||
|
```
|
||||||
|
|
||||||
|
Callback wraps each `DecoderBlock.forward` with `torch.utils.checkpoint.checkpoint(use_reentrant=False)`, compatible with `torch.compile`. Uses `nn.Module.apply()` for traversal — works through DDP wrappers without manual unwrap. Empty list (default) means no-op.
|
||||||
|
|
||||||
|
## Checkpoint
|
||||||
|
|
||||||
|
```
|
||||||
|
Checkpoint(state_dict, epoch, iteration, extra, meta)
|
||||||
|
├── save(save_dir) rank-0 only: meta.json (includes training config) + state_dict.safetensors + optional extra.pt
|
||||||
|
└── load(save_dir) broadcasts metadata from rank-0
|
||||||
|
```
|
||||||
|
|
||||||
|
Optimizer/scheduler state persisted by default via `Checkpoint.extra`.
|
||||||
|
Training config (`TrainConfig.to_dict()`) saved into `meta.json` during training via `CheckpointCallback`.
|
||||||
|
|
||||||
|
## TrainContextBuilder (Builder Pattern)
|
||||||
|
|
||||||
|
```python
|
||||||
|
context = (
|
||||||
|
TrainContextBuilder(config)
|
||||||
|
.with_checkpoint(checkpoint)
|
||||||
|
.build()
|
||||||
|
)
|
||||||
|
# Returns TrainContext with model, strategy, optimizer, scheduler, dataloader, checkpoint
|
||||||
|
```
|
||||||
|
|
||||||
|
- Loads checkpoint weights if provided
|
||||||
|
- Wraps model with `parallel_wrapper` if `nprocs > 1`
|
||||||
|
- Creates `ResumableDistributedSampler` for shuffle+resume
|
||||||
|
- Builds strategy via `StrategyFactory.create(train_type, ...)`
|
||||||
|
|
||||||
|
## Training CLI
|
||||||
|
|
||||||
|
```bash
|
||||||
|
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||||
|
|
||||||
|
nohup python scripts/tools/train.py \
|
||||||
|
--nprocs=4 \
|
||||||
|
--train_type=seq \
|
||||||
|
--data_root_path=/path/to/dataset \
|
||||||
|
--param_path=/path/to/model \
|
||||||
|
--batch_per_device=4 \
|
||||||
|
--grad_accum_steps=8 \
|
||||||
|
--warmup_ratio=0.05 \
|
||||||
|
--max_lr=1e-4 \
|
||||||
|
--max_grad_norm=1.0 \
|
||||||
|
--adamw_beta1=0.9 \
|
||||||
|
--adamw_beta2=0.95 \
|
||||||
|
--adamw_weight_decay=0.01 \
|
||||||
|
--window_size=2048 \
|
||||||
|
--ckpt_interval=10000 \
|
||||||
|
--ckpt_dir=./checkpoint \
|
||||||
|
--random_seed=3407 \
|
||||||
|
--label_smoothing=0.05 \
|
||||||
|
> out.log 2> err.log &
|
||||||
|
```
|
||||||
|
|
||||||
|
Full parameter reference at [params.md](params.md).
|
||||||
|
|
||||||
|
> Document Update Time: 2026-05-17
|
||||||
+7
-5
@@ -1,8 +1,9 @@
|
|||||||
__version__ = "1.3.3"
|
__version__ = "1.3.6"
|
||||||
__author__ = "ViperEkura"
|
__author__ = "ViperEkura"
|
||||||
|
|
||||||
from astrai.config import (
|
from astrai.config import (
|
||||||
ModelConfig,
|
AutoRegressiveLMConfig,
|
||||||
|
EncoderConfig,
|
||||||
TrainConfig,
|
TrainConfig,
|
||||||
)
|
)
|
||||||
from astrai.dataset import DatasetFactory
|
from astrai.dataset import DatasetFactory
|
||||||
@@ -11,13 +12,14 @@ from astrai.inference import (
|
|||||||
GenerationRequest,
|
GenerationRequest,
|
||||||
InferenceEngine,
|
InferenceEngine,
|
||||||
)
|
)
|
||||||
from astrai.model import AutoModel, Transformer
|
from astrai.model import AutoModel, AutoRegressiveLM
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
from astrai.trainer import CallbackFactory, SchedulerFactory, StrategyFactory, Trainer
|
from astrai.trainer import CallbackFactory, SchedulerFactory, StrategyFactory, Trainer
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"Transformer",
|
"AutoRegressiveLM",
|
||||||
"ModelConfig",
|
"AutoRegressiveLMConfig",
|
||||||
|
"EncoderConfig",
|
||||||
"TrainConfig",
|
"TrainConfig",
|
||||||
"DatasetFactory",
|
"DatasetFactory",
|
||||||
"AutoTokenizer",
|
"AutoTokenizer",
|
||||||
|
|||||||
@@ -1,8 +1,16 @@
|
|||||||
from astrai.config.model_config import ModelConfig
|
from astrai.config.model_config import (
|
||||||
|
AutoRegressiveLMConfig,
|
||||||
|
BaseModelConfig,
|
||||||
|
ConfigFactory,
|
||||||
|
EncoderConfig,
|
||||||
|
)
|
||||||
from astrai.config.train_config import TrainConfig
|
from astrai.config.train_config import TrainConfig
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
# Model configuration
|
# Model configuration
|
||||||
"ModelConfig",
|
"BaseModelConfig",
|
||||||
|
"AutoRegressiveLMConfig",
|
||||||
|
"EncoderConfig",
|
||||||
|
"ConfigFactory",
|
||||||
"TrainConfig",
|
"TrainConfig",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -0,0 +1,77 @@
|
|||||||
|
import json
|
||||||
|
from dataclasses import MISSING, dataclass, fields
|
||||||
|
from typing import Any, Dict, Optional, Self, get_type_hints
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class BaseConfig:
|
||||||
|
def to_dict(self) -> Dict[str, Any]:
|
||||||
|
d = {}
|
||||||
|
for fld in fields(self):
|
||||||
|
v = getattr(self, fld.name)
|
||||||
|
if isinstance(v, (str, int, float, bool)):
|
||||||
|
d[fld.name] = v
|
||||||
|
elif v is None:
|
||||||
|
d[fld.name] = None
|
||||||
|
elif isinstance(v, (dict, list)):
|
||||||
|
try:
|
||||||
|
json.dumps(v)
|
||||||
|
d[fld.name] = v
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
pass
|
||||||
|
return d
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_dict(cls, d: Dict[str, Any]) -> Self:
|
||||||
|
hints = get_type_hints(cls)
|
||||||
|
inst = cls.__new__(cls)
|
||||||
|
for fld in fields(cls):
|
||||||
|
if fld.name in d:
|
||||||
|
v = d[fld.name]
|
||||||
|
target = cls._unwrap_optional(hints.get(fld.name))
|
||||||
|
if target is not None:
|
||||||
|
try:
|
||||||
|
v = cls._coerce(v, target)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
pass
|
||||||
|
object.__setattr__(inst, fld.name, v)
|
||||||
|
elif fld.default is not MISSING:
|
||||||
|
object.__setattr__(inst, fld.name, fld.default)
|
||||||
|
elif fld.default_factory is not MISSING:
|
||||||
|
object.__setattr__(inst, fld.name, fld.default_factory())
|
||||||
|
else:
|
||||||
|
object.__setattr__(inst, fld.name, None)
|
||||||
|
return inst
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _unwrap_optional(tp) -> Optional[type]:
|
||||||
|
if tp is None:
|
||||||
|
return None
|
||||||
|
origin = getattr(tp, "__origin__", None)
|
||||||
|
if origin is not None:
|
||||||
|
args = getattr(tp, "__args__", ())
|
||||||
|
non_none = [a for a in args if a is not type(None)]
|
||||||
|
return non_none[0] if non_none else None
|
||||||
|
return tp
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _coerce(value: Any, target_type: type) -> Any:
|
||||||
|
if target_type is bool and isinstance(value, bool):
|
||||||
|
return value
|
||||||
|
if (
|
||||||
|
target_type is int
|
||||||
|
and isinstance(value, (int, float))
|
||||||
|
and not isinstance(value, bool)
|
||||||
|
):
|
||||||
|
return int(value)
|
||||||
|
if (
|
||||||
|
target_type is float
|
||||||
|
and isinstance(value, (int, float))
|
||||||
|
and not isinstance(value, bool)
|
||||||
|
):
|
||||||
|
return float(value)
|
||||||
|
if target_type is str and isinstance(value, str):
|
||||||
|
return value
|
||||||
|
if isinstance(value, target_type):
|
||||||
|
return value
|
||||||
|
raise TypeError
|
||||||
@@ -1,42 +1,90 @@
|
|||||||
import json
|
import json
|
||||||
from dataclasses import asdict, dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Optional, Self
|
from typing import Any, Dict, Optional, Self
|
||||||
|
|
||||||
|
from astrai.config.base import BaseConfig
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
|
||||||
|
|
||||||
|
class ConfigFactory(BaseFactory[BaseConfig]):
|
||||||
|
"""Factory that dispatches config classes by ``model_type``."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def load(cls, raw: Dict[str, Any]) -> BaseConfig:
|
||||||
|
model_type = raw.get("model_type") or "autoregressive_lm"
|
||||||
|
config_cls = cls.get_component_class(model_type)
|
||||||
|
return config_cls.from_dict(raw)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ModelConfig:
|
class BaseModelConfig(BaseConfig):
|
||||||
# basic config
|
"""Base config with ``model_type`` dispatch and file I/O."""
|
||||||
|
|
||||||
model_type: Optional[str] = None
|
model_type: Optional[str] = None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_file(cls, config_path: str) -> Self:
|
||||||
|
with open(config_path, "r") as f:
|
||||||
|
raw: Dict[str, Any] = json.load(f)
|
||||||
|
return cls.from_dict(raw)
|
||||||
|
|
||||||
|
def to_file(self, config_path: str):
|
||||||
|
d = self.to_dict()
|
||||||
|
config_dict = {k: v for k, v in d.items() if v is not None}
|
||||||
|
with open(config_path, "w") as f:
|
||||||
|
json.dump(config_dict, f, indent=4)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
@ConfigFactory.register("autoregressive_lm")
|
||||||
|
class AutoRegressiveLMConfig(BaseModelConfig):
|
||||||
|
"""Configuration for autoregressive language model."""
|
||||||
|
|
||||||
vocab_size: Optional[int] = None
|
vocab_size: Optional[int] = None
|
||||||
dim: Optional[int] = None
|
dim: Optional[int] = None
|
||||||
|
|
||||||
n_layers: Optional[int] = None
|
n_layers: Optional[int] = None
|
||||||
norm_eps: Optional[float] = None
|
norm_eps: Optional[float] = None
|
||||||
dim_ffn: Optional[int] = None
|
dim_ffn: Optional[int] = None
|
||||||
tie_weight: Optional[bool] = None
|
tie_weight: Optional[bool] = None
|
||||||
|
|
||||||
# RoPE
|
|
||||||
max_len: Optional[int] = None
|
max_len: Optional[int] = None
|
||||||
rope_theta: Optional[float] = None
|
rope_theta: Optional[float] = None
|
||||||
|
|
||||||
# GQA
|
attn_type: str = "gqa"
|
||||||
n_heads: Optional[int] = None
|
n_heads: Optional[int] = None
|
||||||
n_kv_heads: Optional[int] = None
|
n_kv_heads: Optional[int] = None
|
||||||
use_qk_norm: Optional[bool] = None
|
use_qk_norm: Optional[bool] = None
|
||||||
use_gated_attention: Optional[bool] = None
|
use_gated_attention: Optional[bool] = None
|
||||||
|
|
||||||
def load(self, config_path: str) -> Self:
|
kv_lora_rank: Optional[int] = None
|
||||||
config = {}
|
qk_nope_head_dim: Optional[int] = None
|
||||||
with open(config_path, "r") as f:
|
qk_rope_head_dim: Optional[int] = None
|
||||||
config.update(json.load(f))
|
|
||||||
|
|
||||||
for key, value in config.items():
|
ffn_type: str = "mlp"
|
||||||
if hasattr(self, key):
|
n_routed_experts: Optional[int] = None
|
||||||
setattr(self, key, value)
|
n_shared_experts: Optional[int] = None
|
||||||
|
n_activated_experts: Optional[int] = None
|
||||||
|
topk_method: Optional[str] = None
|
||||||
|
|
||||||
return self
|
|
||||||
|
|
||||||
def save(self, config_path: str):
|
@dataclass
|
||||||
config_dict = {k: v for k, v in asdict(self).items() if v is not None}
|
@ConfigFactory.register("embedding")
|
||||||
with open(config_path, "w") as f:
|
class EncoderConfig(BaseModelConfig):
|
||||||
json.dump(config_dict, f, indent=4)
|
"""Configuration for embedding encoder model."""
|
||||||
|
|
||||||
|
vocab_size: Optional[int] = None
|
||||||
|
dim: Optional[int] = None
|
||||||
|
n_layers: Optional[int] = None
|
||||||
|
norm_eps: Optional[float] = None
|
||||||
|
dim_ffn: Optional[int] = None
|
||||||
|
|
||||||
|
max_len: Optional[int] = None
|
||||||
|
rope_theta: Optional[float] = None
|
||||||
|
|
||||||
|
n_heads: Optional[int] = None
|
||||||
|
n_kv_heads: Optional[int] = None
|
||||||
|
use_qk_norm: Optional[bool] = None
|
||||||
|
use_gated_attention: Optional[bool] = None
|
||||||
|
|
||||||
|
pooling_type: Optional[str] = None
|
||||||
|
normalize_embeddings: Optional[bool] = None
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field, fields
|
||||||
from typing import Callable, List, Optional
|
from typing import Callable, List, Optional
|
||||||
|
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
@@ -6,27 +6,43 @@ from torch.optim import Optimizer
|
|||||||
from torch.optim.lr_scheduler import LRScheduler
|
from torch.optim.lr_scheduler import LRScheduler
|
||||||
from torch.utils.data import Dataset
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
|
from astrai.config.base import BaseConfig
|
||||||
|
|
||||||
|
|
||||||
|
def required(**kw):
|
||||||
|
return {"required": True, **kw}
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class TrainConfig:
|
class TrainConfig(BaseConfig):
|
||||||
# basic setting
|
# basic setting
|
||||||
model: nn.Module = field(default=None, metadata={"help": "Model for training."})
|
model: nn.Module = field(
|
||||||
strategy: str = field(default=None, metadata={"help": "Training strategy."})
|
default=None, metadata=required(help="Model for training.")
|
||||||
dataset: Dataset = field(default=None, metadata={"help": "Dataset for training."})
|
)
|
||||||
|
strategy: str = field(default=None, metadata=required(help="Training strategy."))
|
||||||
|
dataset: Dataset = field(
|
||||||
|
default=None, metadata=required(help="Dataset for training.")
|
||||||
|
)
|
||||||
optimizer_fn: Callable[[nn.Module], Optimizer] = field(
|
optimizer_fn: Callable[[nn.Module], Optimizer] = field(
|
||||||
default=None, metadata={"help": "Optimizer factory for training."}
|
default=None, metadata=required(help="Optimizer factory for training.")
|
||||||
)
|
)
|
||||||
scheduler_fn: Callable[[Optimizer], LRScheduler] = field(
|
scheduler_fn: Callable[[Optimizer], LRScheduler] = field(
|
||||||
default=None, metadata={"help": "Scheduler factory for training."}
|
default=None, metadata=required(help="Scheduler factory for training.")
|
||||||
)
|
)
|
||||||
n_epoch: int = field(default=1, metadata={"help": "Number of epochs 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."})
|
batch_per_device: int = field(
|
||||||
accumulation_steps: int = field(
|
default=4, metadata={"help": "Batch size per device."}
|
||||||
|
)
|
||||||
|
grad_accum_steps: int = field(
|
||||||
default=1, metadata={"help": "Number of iterations between steps."}
|
default=1, metadata={"help": "Number of iterations between steps."}
|
||||||
)
|
)
|
||||||
max_grad_norm: float = field(
|
max_grad_norm: float = field(
|
||||||
default=1.0, metadata={"help": "Maximum gradient norm."}
|
default=1.0, metadata={"help": "Maximum gradient norm."}
|
||||||
)
|
)
|
||||||
|
gradient_checkpointing_modules: list = field(
|
||||||
|
default_factory=list,
|
||||||
|
metadata={"help": "Module types to enable activation checkpointing for."},
|
||||||
|
)
|
||||||
|
|
||||||
# checkpoint setting
|
# checkpoint setting
|
||||||
start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."})
|
start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."})
|
||||||
@@ -40,6 +56,19 @@ class TrainConfig:
|
|||||||
default=5000, metadata={"help": "Number of iterations between checkpoints."}
|
default=5000, metadata={"help": "Number of iterations between checkpoints."}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# metric setting
|
||||||
|
log_dir: str = field(
|
||||||
|
default="./checkpoint/logs", metadata={"help": "Directory for metric logs."}
|
||||||
|
)
|
||||||
|
log_interval: int = field(
|
||||||
|
default=100,
|
||||||
|
metadata={"help": "Number of batch iterations between metric logs."},
|
||||||
|
)
|
||||||
|
metrics: List[str] = field(
|
||||||
|
default_factory=lambda: ["loss", "lr"],
|
||||||
|
metadata={"help": "Metrics to record during training."},
|
||||||
|
)
|
||||||
|
|
||||||
# dataloader setting
|
# dataloader setting
|
||||||
random_seed: int = field(default=3407, metadata={"help": "Random seed."})
|
random_seed: int = field(default=3407, metadata={"help": "Random seed."})
|
||||||
num_workers: int = field(
|
num_workers: int = field(
|
||||||
@@ -72,14 +101,23 @@ class TrainConfig:
|
|||||||
state_dict_fn: Optional[Callable] = field(
|
state_dict_fn: Optional[Callable] = field(
|
||||||
default=None, metadata={"help": "Parallel function for state dict saving."}
|
default=None, metadata={"help": "Parallel function for state dict saving."}
|
||||||
)
|
)
|
||||||
|
start_method: str = field(
|
||||||
|
default="spawn",
|
||||||
|
metadata={"help": "Multiprocessing start method (spawn/fork/forkserver)."},
|
||||||
|
)
|
||||||
|
|
||||||
# others
|
# others
|
||||||
device_ids: Optional[List[int]] = field(
|
|
||||||
default=None, metadata={"help": "Device ids for distributed training."}
|
|
||||||
)
|
|
||||||
device_type: str = field(
|
device_type: str = field(
|
||||||
default="cuda", metadata={"help": "Device type for distributed training."}
|
default="cuda", metadata={"help": "Device type for distributed training."}
|
||||||
)
|
)
|
||||||
|
val_dataset: Optional[Dataset] = field(
|
||||||
|
default=None, metadata={"help": "Dataset for validation."}
|
||||||
|
)
|
||||||
|
val_step: int = field(
|
||||||
|
default=1000,
|
||||||
|
metadata={"help": "Number of optimizer steps between validation runs."},
|
||||||
|
)
|
||||||
|
|
||||||
extra_kwargs: dict = field(
|
extra_kwargs: dict = field(
|
||||||
default_factory=dict, metadata={"help": "Other arguments."}
|
default_factory=dict, metadata={"help": "Other arguments."}
|
||||||
)
|
)
|
||||||
@@ -88,14 +126,6 @@ class TrainConfig:
|
|||||||
self.validate()
|
self.validate()
|
||||||
|
|
||||||
def validate(self):
|
def validate(self):
|
||||||
required_fields = [
|
for fld in fields(self):
|
||||||
"model",
|
if fld.metadata.get("required") and getattr(self, fld.name) is None:
|
||||||
"strategy",
|
raise ValueError(f"TrainConfig.{fld.name} is required but got None.")
|
||||||
"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,19 +1,35 @@
|
|||||||
from astrai.dataset.dataset import (
|
from astrai.dataset.dataset import (
|
||||||
BaseDataset,
|
BaseDataset,
|
||||||
BaseSegmentFetcher,
|
|
||||||
DatasetFactory,
|
DatasetFactory,
|
||||||
MultiSegmentFetcher,
|
|
||||||
)
|
)
|
||||||
from astrai.dataset.sampler import ResumableDistributedSampler
|
from astrai.dataset.sampler import ResumableDistributedSampler
|
||||||
|
from astrai.dataset.storage import (
|
||||||
|
BaseSegmentFetcher,
|
||||||
|
BaseStorage,
|
||||||
|
H5Storage,
|
||||||
|
JSONStorage,
|
||||||
|
MultiSegmentFetcher,
|
||||||
|
StorageFactory,
|
||||||
|
detect_format,
|
||||||
|
load_h5,
|
||||||
|
load_json,
|
||||||
|
save_h5,
|
||||||
|
save_json,
|
||||||
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
# Base classes
|
|
||||||
"BaseDataset",
|
"BaseDataset",
|
||||||
# Factory
|
|
||||||
"DatasetFactory",
|
"DatasetFactory",
|
||||||
# Fetchers
|
|
||||||
"BaseSegmentFetcher",
|
"BaseSegmentFetcher",
|
||||||
"MultiSegmentFetcher",
|
"MultiSegmentFetcher",
|
||||||
# Sampler
|
"BaseStorage",
|
||||||
|
"H5Storage",
|
||||||
|
"JSONStorage",
|
||||||
|
"StorageFactory",
|
||||||
|
"detect_format",
|
||||||
|
"save_h5",
|
||||||
|
"load_h5",
|
||||||
|
"save_json",
|
||||||
|
"load_json",
|
||||||
"ResumableDistributedSampler",
|
"ResumableDistributedSampler",
|
||||||
]
|
]
|
||||||
|
|||||||
+108
-127
@@ -1,140 +1,97 @@
|
|||||||
"""Dataset implementations with factory pattern for training."""
|
"""Dataset implementations with factory pattern for training."""
|
||||||
|
|
||||||
import bisect
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import Dict, List, Optional, Union
|
from typing import Dict, List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
from torch.utils.data import Dataset
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
|
from astrai.dataset.storage import (
|
||||||
|
BaseStorage,
|
||||||
|
StorageFactory,
|
||||||
|
detect_format,
|
||||||
|
)
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.serialization import load_h5
|
|
||||||
|
|
||||||
|
|
||||||
class BaseSegmentFetcher:
|
|
||||||
"""Fetches data segments across multiple tensor segments.
|
|
||||||
|
|
||||||
Maintains cumulative lengths for efficient range queries across
|
|
||||||
multiple discontinuous segments.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, segments: List[Tensor]):
|
|
||||||
self.segments = segments
|
|
||||||
self.cum_lengths = []
|
|
||||||
|
|
||||||
total = 0
|
|
||||||
for seg in segments:
|
|
||||||
total += torch.numel(seg)
|
|
||||||
self.cum_lengths.append(total)
|
|
||||||
|
|
||||||
self.total_length = total
|
|
||||||
|
|
||||||
def __len__(self) -> int:
|
|
||||||
return self.total_length
|
|
||||||
|
|
||||||
def fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
|
||||||
"""Fetch data in the range [begin_idx, end_idx).
|
|
||||||
|
|
||||||
Args:
|
|
||||||
begin_idx: Starting index (inclusive)
|
|
||||||
end_idx: Ending index (exclusive)
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Concatenated tensor of data in the specified range
|
|
||||||
"""
|
|
||||||
if not (
|
|
||||||
0 <= begin_idx < self.total_length and 0 <= end_idx <= self.total_length
|
|
||||||
):
|
|
||||||
raise ValueError("begin_idx or end_idx out of bounds")
|
|
||||||
if begin_idx >= end_idx:
|
|
||||||
return torch.tensor([], dtype=torch.long)
|
|
||||||
|
|
||||||
# Find segment boundaries for the range
|
|
||||||
seg_start_idx = bisect.bisect_right(self.cum_lengths, begin_idx)
|
|
||||||
seg_end_idx = bisect.bisect_left(self.cum_lengths, end_idx)
|
|
||||||
|
|
||||||
result_segments = []
|
|
||||||
|
|
||||||
for i in range(seg_start_idx, seg_end_idx + 1):
|
|
||||||
prev_cum = self.cum_lengths[i - 1] if i > 0 else 0
|
|
||||||
start = max(begin_idx - prev_cum, 0)
|
|
||||||
end = min(end_idx - prev_cum, len(self.segments[i]))
|
|
||||||
data = self.segments[i][start:end]
|
|
||||||
result_segments.append(data)
|
|
||||||
|
|
||||||
return torch.cat(result_segments, dim=0)
|
|
||||||
|
|
||||||
|
|
||||||
class MultiSegmentFetcher:
|
|
||||||
"""Manages multiple segment fetchers for different data keys.
|
|
||||||
|
|
||||||
Each key corresponds to a different type of data (e.g., "sequence", "mask").
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, multi_segments: Dict):
|
|
||||||
self.multi_keys = list(multi_segments.keys())
|
|
||||||
self.multi_fetchers = {
|
|
||||||
key: BaseSegmentFetcher(segments)
|
|
||||||
for key, segments in multi_segments.items()
|
|
||||||
}
|
|
||||||
|
|
||||||
def __len__(self) -> int:
|
|
||||||
"""Returns the minimum length across all fetchers."""
|
|
||||||
len_list = [len(seg) for seg in self.multi_fetchers.values()]
|
|
||||||
return min(len_list)
|
|
||||||
|
|
||||||
def key_fetch(
|
|
||||||
self, begin_idx: int, end_idx: int, keys: Union[str, List[str]]
|
|
||||||
) -> Dict:
|
|
||||||
"""Fetch data for specific keys.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
begin_idx: Starting index
|
|
||||||
end_idx: Ending index
|
|
||||||
keys: Single key or list of keys to fetch
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Dictionary of tensors if multiple keys, single tensor if one key
|
|
||||||
"""
|
|
||||||
fetch_dict = {}
|
|
||||||
keys = [keys] if isinstance(keys, str) else keys
|
|
||||||
|
|
||||||
for key in keys:
|
|
||||||
fetcher = self.multi_fetchers[key]
|
|
||||||
fetch_tensor = fetcher.fetch_data(begin_idx, end_idx)
|
|
||||||
fetch_dict[key] = fetch_tensor
|
|
||||||
|
|
||||||
return fetch_dict if len(keys) > 1 else fetch_dict[keys[0]]
|
|
||||||
|
|
||||||
def fetch_data(self, begin_idx: int, end_idx: int) -> Dict:
|
|
||||||
"""Fetch all keys."""
|
|
||||||
return self.key_fetch(begin_idx, end_idx, self.multi_keys)
|
|
||||||
|
|
||||||
|
|
||||||
class BaseDataset(Dataset, ABC):
|
class BaseDataset(Dataset, ABC):
|
||||||
"""Abstract base class for all dataset types.
|
"""Abstract base class for all dataset types.
|
||||||
|
|
||||||
Implements common functionality for window-based data fetching.
|
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):
|
def __init__(self, window_size: int, stride: int):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.segments = {}
|
|
||||||
self.window_size = window_size
|
self.window_size = window_size
|
||||||
self.stride = stride
|
self.stride = stride
|
||||||
self.total_samples = None
|
self.storage: Optional[BaseStorage] = None
|
||||||
self.fetcher: Optional[MultiSegmentFetcher] = None
|
|
||||||
|
|
||||||
def load(self, load_path: str):
|
@property
|
||||||
"""Load dataset from HDF5 file.
|
def required_keys(self) -> List[str]:
|
||||||
|
"""Return required storage keys for this dataset type.
|
||||||
|
|
||||||
|
Subclasses should override to specify expected keys.
|
||||||
|
"""
|
||||||
|
return []
|
||||||
|
|
||||||
|
def _validate_keys(self):
|
||||||
|
if not self.required_keys:
|
||||||
|
return
|
||||||
|
actual_keys = set(self.storage.keys)
|
||||||
|
missing = [k for k in self.required_keys if k not in actual_keys]
|
||||||
|
if missing:
|
||||||
|
raise KeyError(
|
||||||
|
f"Dataset {type(self).__name__} requires keys {self.required_keys}, "
|
||||||
|
f"but storage at {self._load_path} only has {sorted(actual_keys)}. "
|
||||||
|
f"Missing: {missing}"
|
||||||
|
)
|
||||||
|
|
||||||
|
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:
|
Args:
|
||||||
load_path: Path to the HDF5 data file
|
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.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
KeyError: If the loaded storage is missing required keys.
|
||||||
"""
|
"""
|
||||||
self.segments = load_h5(load_path)
|
if storage_type is None:
|
||||||
self.fetcher = MultiSegmentFetcher(self.segments)
|
storage_type = detect_format(load_path)
|
||||||
self.total_samples = len(self.fetcher)
|
self.storage = StorageFactory.create(storage_type)
|
||||||
|
self._load_path = load_path
|
||||||
|
self.storage.load(load_path, tokenizer=tokenizer)
|
||||||
|
self._validate_keys()
|
||||||
|
|
||||||
|
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:
|
def get_index(self, index: int) -> tuple:
|
||||||
"""Calculate begin and end indices for a sample.
|
"""Calculate begin and end indices for a sample.
|
||||||
@@ -145,10 +102,16 @@ class BaseDataset(Dataset, ABC):
|
|||||||
Returns:
|
Returns:
|
||||||
Tuple of (begin_idx, end_idx)
|
Tuple of (begin_idx, end_idx)
|
||||||
"""
|
"""
|
||||||
assert self.total_samples > self.window_size
|
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, self.total_samples - 1 - self.window_size)
|
begin_idx = min(index * self.stride, total - 1 - self.window_size)
|
||||||
end_idx = min(begin_idx + self.window_size, self.total_samples - 1)
|
end_idx = min(begin_idx + self.window_size, total - 1)
|
||||||
|
|
||||||
return begin_idx, end_idx
|
return begin_idx, end_idx
|
||||||
|
|
||||||
@@ -161,10 +124,12 @@ class BaseDataset(Dataset, ABC):
|
|||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
def __len__(self) -> int:
|
def __len__(self) -> int:
|
||||||
assert self.total_samples is not None
|
if self.storage is None:
|
||||||
if self.total_samples <= self.window_size:
|
|
||||||
return 0
|
return 0
|
||||||
return (self.total_samples - 1 - self.window_size) // self.stride + 1
|
total = len(self.storage)
|
||||||
|
if total <= self.window_size:
|
||||||
|
return 0
|
||||||
|
return (total - 1 - self.window_size) // self.stride + 1
|
||||||
|
|
||||||
|
|
||||||
class DatasetFactory(BaseFactory["BaseDataset"]):
|
class DatasetFactory(BaseFactory["BaseDataset"]):
|
||||||
@@ -209,6 +174,8 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
|||||||
load_path: str,
|
load_path: str,
|
||||||
window_size: int,
|
window_size: int,
|
||||||
stride: Optional[int] = None,
|
stride: Optional[int] = None,
|
||||||
|
storage_type: Optional[str] = None,
|
||||||
|
tokenizer=None,
|
||||||
) -> "BaseDataset":
|
) -> "BaseDataset":
|
||||||
"""Create and load a dataset in one step.
|
"""Create and load a dataset in one step.
|
||||||
|
|
||||||
@@ -217,6 +184,8 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
|||||||
load_path: Path to the data file
|
load_path: Path to the data file
|
||||||
window_size: Window size for data sampling
|
window_size: Window size for data sampling
|
||||||
stride: Stride between consecutive samples (default: same as window_size)
|
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:
|
Returns:
|
||||||
Loaded dataset instance
|
Loaded dataset instance
|
||||||
@@ -225,7 +194,7 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
|||||||
stride = window_size
|
stride = window_size
|
||||||
|
|
||||||
dataset = cls.create(train_type, window_size, stride)
|
dataset = cls.create(train_type, window_size, stride)
|
||||||
dataset.load(load_path)
|
dataset.load(load_path, storage_type=storage_type, tokenizer=tokenizer)
|
||||||
|
|
||||||
return dataset
|
return dataset
|
||||||
|
|
||||||
@@ -235,10 +204,6 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
|||||||
return cls.list_registered()
|
return cls.list_registered()
|
||||||
|
|
||||||
|
|
||||||
# ============== Dataset Classes ==============
|
|
||||||
# All dataset classes are registered at class definition time using the decorator
|
|
||||||
|
|
||||||
|
|
||||||
@DatasetFactory.register("seq")
|
@DatasetFactory.register("seq")
|
||||||
class SEQDataset(BaseDataset):
|
class SEQDataset(BaseDataset):
|
||||||
"""Dataset for sequential next-token prediction training."""
|
"""Dataset for sequential next-token prediction training."""
|
||||||
@@ -246,8 +211,12 @@ class SEQDataset(BaseDataset):
|
|||||||
def __init__(self, window_size: int, stride: int):
|
def __init__(self, window_size: int, stride: int):
|
||||||
super().__init__(window_size, stride)
|
super().__init__(window_size, stride)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def required_keys(self) -> List[str]:
|
||||||
|
return ["sequence"]
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
def _fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
||||||
return self.fetcher.key_fetch(begin_idx, end_idx, "sequence")
|
return self.storage.fetch(begin_idx, end_idx, "sequence")
|
||||||
|
|
||||||
def __getitem__(self, index):
|
def __getitem__(self, index):
|
||||||
begin_idx, end_idx = self.get_index(index)
|
begin_idx, end_idx = self.get_index(index)
|
||||||
@@ -265,8 +234,12 @@ class SFTDataset(BaseDataset):
|
|||||||
def __init__(self, window_size: int, stride: int):
|
def __init__(self, window_size: int, stride: int):
|
||||||
super().__init__(window_size, stride)
|
super().__init__(window_size, stride)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def required_keys(self) -> List[str]:
|
||||||
|
return ["sequence", "loss_mask"]
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||||
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
return self.storage.fetch(begin_idx, end_idx, key)
|
||||||
|
|
||||||
def __getitem__(self, index):
|
def __getitem__(self, index):
|
||||||
begin_idx, end_idx = self.get_index(index)
|
begin_idx, end_idx = self.get_index(index)
|
||||||
@@ -289,8 +262,12 @@ class DPODataset(BaseDataset):
|
|||||||
def __init__(self, window_size: int, stride: int):
|
def __init__(self, window_size: int, stride: int):
|
||||||
super().__init__(window_size, stride)
|
super().__init__(window_size, stride)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def required_keys(self) -> List[str]:
|
||||||
|
return ["chosen", "rejected", "chosen_mask", "rejected_mask"]
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||||
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
return self.storage.fetch(begin_idx, end_idx, key)
|
||||||
|
|
||||||
def __getitem__(self, index: int):
|
def __getitem__(self, index: int):
|
||||||
begin_idx, end_idx = self.get_index(index)
|
begin_idx, end_idx = self.get_index(index)
|
||||||
@@ -319,8 +296,12 @@ class GRPODataset(BaseDataset):
|
|||||||
def __init__(self, window_size: int, stride: int):
|
def __init__(self, window_size: int, stride: int):
|
||||||
super().__init__(window_size, stride)
|
super().__init__(window_size, stride)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def required_keys(self) -> List[str]:
|
||||||
|
return ["prompts", "responses", "masks", "rewards"]
|
||||||
|
|
||||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||||
return self.fetcher.key_fetch(begin_idx, end_idx, key)
|
return self.storage.fetch(begin_idx, end_idx, key)
|
||||||
|
|
||||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||||
begin_idx, end_idx = self.get_index(index)
|
begin_idx, end_idx = self.get_index(index)
|
||||||
|
|||||||
@@ -0,0 +1,301 @@
|
|||||||
|
"""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
|
||||||
|
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
|
||||||
|
|
||||||
|
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 StorageFactory(BaseFactory["BaseStorage"]):
|
||||||
|
"""Factory for creating storage backends by type name.
|
||||||
|
|
||||||
|
Example:
|
||||||
|
@StorageFactory.register("custom")
|
||||||
|
class CustomStorage(BaseStorage):
|
||||||
|
...
|
||||||
|
|
||||||
|
storage = StorageFactory.create("custom")
|
||||||
|
"""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _validate_component(cls, storage_cls: type) -> None:
|
||||||
|
if not issubclass(storage_cls, BaseStorage):
|
||||||
|
raise TypeError(f"{storage_cls.__name__} must inherit from BaseStorage")
|
||||||
|
|
||||||
|
|
||||||
|
@StorageFactory.register("h5")
|
||||||
|
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)
|
||||||
|
|
||||||
|
|
||||||
|
@StorageFactory.register("json")
|
||||||
|
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)
|
||||||
@@ -1,5 +1,6 @@
|
|||||||
"""Base factory class for extensible component registration."""
|
"""Base factory class for extensible component registration."""
|
||||||
|
|
||||||
|
import inspect
|
||||||
from abc import ABC
|
from abc import ABC
|
||||||
from typing import Callable, Dict, Generic, List, Optional, Tuple, Type, TypeVar
|
from typing import Callable, Dict, Generic, List, Optional, Tuple, Type, TypeVar
|
||||||
|
|
||||||
@@ -122,6 +123,10 @@ class BaseFactory(ABC, Generic[T]):
|
|||||||
def create(cls, name: str, *args, **kwargs) -> T:
|
def create(cls, name: str, *args, **kwargs) -> T:
|
||||||
"""Create a component instance by name.
|
"""Create a component instance by name.
|
||||||
|
|
||||||
|
Filters kwargs to match the component's __init__ signature,
|
||||||
|
so components don't need to declare **kwargs just to absorb
|
||||||
|
parameters meant for other components.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
name: Registered name of the component
|
name: Registered name of the component
|
||||||
*args: Positional arguments passed to component constructor
|
*args: Positional arguments passed to component constructor
|
||||||
@@ -139,6 +144,17 @@ class BaseFactory(ABC, Generic[T]):
|
|||||||
f"Supported types: {sorted(cls._registry.list_names())}"
|
f"Supported types: {sorted(cls._registry.list_names())}"
|
||||||
)
|
)
|
||||||
component_cls = cls._registry.get(name)
|
component_cls = cls._registry.get(name)
|
||||||
|
sig = inspect.signature(component_cls.__init__)
|
||||||
|
has_var_kwargs = any(
|
||||||
|
p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()
|
||||||
|
)
|
||||||
|
if not has_var_kwargs:
|
||||||
|
valid = {
|
||||||
|
p.name
|
||||||
|
for p in sig.parameters.values()
|
||||||
|
if p.name != "self" and p.kind != inspect.Parameter.VAR_KEYWORD
|
||||||
|
}
|
||||||
|
kwargs = {k: v for k, v in kwargs.items() if k in valid}
|
||||||
return component_cls(*args, **kwargs)
|
return component_cls(*args, **kwargs)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -155,6 +171,26 @@ class BaseFactory(ABC, Generic[T]):
|
|||||||
"""
|
"""
|
||||||
pass
|
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
|
@classmethod
|
||||||
def list_registered(cls) -> list:
|
def list_registered(cls) -> list:
|
||||||
"""List all registered component names.
|
"""List all registered component names.
|
||||||
|
|||||||
@@ -1,19 +1,46 @@
|
|||||||
"""Inference module for continuous batching.
|
"""Inference module for continuous batching.
|
||||||
|
|
||||||
Layers:
|
Layers:
|
||||||
- engine.py: Facade (InferenceEngine), Value Object (GenerationParams, GenerationRequest)
|
- core/: Core inference loop (cache, executor, scheduler, task)
|
||||||
- scheduler.py: Continuous-batching loop, Task state machine, TaskStatus enum
|
- api/: HTTP protocol handlers (OpenAI, Anthropic)
|
||||||
- cache.py: Object Pool (SlotAllocator), PrefixCacheManager
|
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
|
||||||
- sampling.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy)
|
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy)
|
||||||
- server.py: FastAPI HTTP server (OpenAI-compatible endpoints)
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
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 (
|
from astrai.inference.engine import (
|
||||||
GenerationParams,
|
|
||||||
GenerationRequest,
|
GenerationRequest,
|
||||||
InferenceEngine,
|
InferenceEngine,
|
||||||
)
|
)
|
||||||
from astrai.inference.sampling import (
|
from astrai.inference.sample import (
|
||||||
BaseSamplingStrategy,
|
BaseSamplingStrategy,
|
||||||
SamplingPipeline,
|
SamplingPipeline,
|
||||||
TemperatureStrategy,
|
TemperatureStrategy,
|
||||||
@@ -21,21 +48,27 @@ from astrai.inference.sampling import (
|
|||||||
TopPStrategy,
|
TopPStrategy,
|
||||||
sample,
|
sample,
|
||||||
)
|
)
|
||||||
from astrai.inference.scheduler import (
|
|
||||||
InferenceScheduler,
|
|
||||||
Task,
|
|
||||||
TaskStatus,
|
|
||||||
)
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
# Engine / Requests
|
# Engine / Requests
|
||||||
"InferenceEngine",
|
"InferenceEngine",
|
||||||
"GenerationRequest",
|
"GenerationRequest",
|
||||||
"GenerationParams",
|
# Core scheduler
|
||||||
# Scheduler
|
|
||||||
"InferenceScheduler",
|
"InferenceScheduler",
|
||||||
|
"Executor",
|
||||||
|
"STOP",
|
||||||
"Task",
|
"Task",
|
||||||
|
"TaskManager",
|
||||||
"TaskStatus",
|
"TaskStatus",
|
||||||
|
# Core cache
|
||||||
|
"Allocator",
|
||||||
|
"KVCache",
|
||||||
|
"KvcacheView",
|
||||||
|
"PagePool",
|
||||||
|
"PrefixCache",
|
||||||
|
"Storage",
|
||||||
|
"TaskTable",
|
||||||
|
"page_hash",
|
||||||
# Sampling (Strategy pattern)
|
# Sampling (Strategy pattern)
|
||||||
"sample",
|
"sample",
|
||||||
"BaseSamplingStrategy",
|
"BaseSamplingStrategy",
|
||||||
@@ -43,4 +76,17 @@ __all__ = [
|
|||||||
"TopKStrategy",
|
"TopKStrategy",
|
||||||
"TopPStrategy",
|
"TopPStrategy",
|
||||||
"SamplingPipeline",
|
"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,445 @@
|
|||||||
|
"""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 get_stop_sequences(self) -> List[str]:
|
||||||
|
stop = self.request.stop
|
||||||
|
if stop is None:
|
||||||
|
return []
|
||||||
|
return [stop] if isinstance(stop, str) else stop
|
||||||
|
|
||||||
|
def on_token(
|
||||||
|
self, ctx: StreamContext, token: str, stop_checker: StopChecker
|
||||||
|
) -> Optional[str]:
|
||||||
|
return stop_checker.check(ctx.accumulated)
|
||||||
|
|
||||||
|
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,167 @@
|
|||||||
|
"""
|
||||||
|
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
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
@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 _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
|
||||||
|
|
||||||
|
|
||||||
|
def _get_engine() -> InferenceEngine:
|
||||||
|
engine = app.state.engine
|
||||||
|
if engine is None:
|
||||||
|
raise HTTPException(status_code=503, detail="Engine not initialized")
|
||||||
|
return engine
|
||||||
|
|
||||||
|
|
||||||
|
@app.get("/health")
|
||||||
|
async def health():
|
||||||
|
return {
|
||||||
|
"status": "ok",
|
||||||
|
"model_loaded": app.state.engine is not None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@app.get("/stats")
|
||||||
|
async def get_stats():
|
||||||
|
return _get_engine().get_stats()
|
||||||
|
|
||||||
|
|
||||||
|
@app.post("/v1/chat/completions")
|
||||||
|
async def chat_completion(request: ChatCompletionRequest):
|
||||||
|
engine = _get_engine()
|
||||||
|
handler = OpenAIHandler(request, engine)
|
||||||
|
return await handler.handle()
|
||||||
|
|
||||||
|
|
||||||
|
@app.post("/v1/messages")
|
||||||
|
async def create_message(request: MessagesRequest):
|
||||||
|
engine = _get_engine()
|
||||||
|
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,
|
||||||
|
reload=reload,
|
||||||
|
)
|
||||||
@@ -1,135 +0,0 @@
|
|||||||
"""Page-based KV cache with page-table-indirected read/write.
|
|
||||||
|
|
||||||
Provides:
|
|
||||||
- PagedCache: paged KV cache combining page pool and tensor storage.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from typing import List, Tuple
|
|
||||||
|
|
||||||
import torch
|
|
||||||
from torch import Tensor
|
|
||||||
|
|
||||||
STOP = object()
|
|
||||||
|
|
||||||
|
|
||||||
class PagedCache:
|
|
||||||
"""Paged KV cache with page-table-indirected read/write.
|
|
||||||
|
|
||||||
Combines:
|
|
||||||
- Page pool (ref-counted alloc/free via bitmask)
|
|
||||||
- KV tensor storage (k_cache, v_cache)
|
|
||||||
|
|
||||||
Call :meth:`bind` to obtain a batch view for the attention layers.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
n_layers: int,
|
|
||||||
n_pages: int,
|
|
||||||
page_size: int,
|
|
||||||
n_kv_heads: int,
|
|
||||||
head_dim: int,
|
|
||||||
device: torch.device,
|
|
||||||
dtype: torch.dtype,
|
|
||||||
):
|
|
||||||
self.page_size = page_size
|
|
||||||
self._free_mask = (1 << n_pages) - 1
|
|
||||||
self._refs: List[int] = [0] * n_pages
|
|
||||||
self.k_cache = torch.empty(
|
|
||||||
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
|
|
||||||
device=device,
|
|
||||||
dtype=dtype,
|
|
||||||
)
|
|
||||||
self.v_cache = torch.empty(
|
|
||||||
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
|
|
||||||
device=device,
|
|
||||||
dtype=dtype,
|
|
||||||
)
|
|
||||||
|
|
||||||
def alloc(self) -> int:
|
|
||||||
lsb = self._free_mask & -self._free_mask
|
|
||||||
if lsb == 0:
|
|
||||||
return -1
|
|
||||||
idx = lsb.bit_length() - 1
|
|
||||||
self._free_mask ^= lsb
|
|
||||||
self._refs[idx] = 1
|
|
||||||
return idx
|
|
||||||
|
|
||||||
def alloc_n(self, n: int) -> List[int]:
|
|
||||||
pages = [self.alloc() for _ in range(n)]
|
|
||||||
if any(p < 0 for p in pages):
|
|
||||||
for p in pages:
|
|
||||||
if p >= 0:
|
|
||||||
self.free(p)
|
|
||||||
return []
|
|
||||||
return pages
|
|
||||||
|
|
||||||
def free(self, idx: int) -> None:
|
|
||||||
self._refs[idx] -= 1
|
|
||||||
if self._refs[idx] == 0:
|
|
||||||
self._free_mask |= 1 << idx
|
|
||||||
|
|
||||||
def bind(self, page_table: Tensor, total_len: int = 0) -> "CacheView":
|
|
||||||
return CacheView(self, page_table, total_len)
|
|
||||||
|
|
||||||
def write(
|
|
||||||
self, layer_id: int, page_table: Tensor, start_pos: int, k: Tensor, v: Tensor
|
|
||||||
) -> None:
|
|
||||||
seq_len = k.size(1)
|
|
||||||
if seq_len == 0:
|
|
||||||
return
|
|
||||||
page_size = self.page_size
|
|
||||||
written = 0
|
|
||||||
first_page = start_pos // page_size
|
|
||||||
last_page = (start_pos + seq_len - 1) // page_size
|
|
||||||
for pi in range(first_page, last_page + 1):
|
|
||||||
phys_pages = page_table[:, pi]
|
|
||||||
page_start = pi * page_size
|
|
||||||
write_start = max(page_start, start_pos)
|
|
||||||
write_end = min(page_start + page_size, start_pos + seq_len)
|
|
||||||
offset = write_start - page_start
|
|
||||||
chunk = write_end - write_start
|
|
||||||
self.k_cache[layer_id, phys_pages, offset : offset + chunk] = k[
|
|
||||||
:, written : written + chunk
|
|
||||||
]
|
|
||||||
self.v_cache[layer_id, phys_pages, offset : offset + chunk] = v[
|
|
||||||
:, written : written + chunk
|
|
||||||
]
|
|
||||||
written += chunk
|
|
||||||
|
|
||||||
def gather(self, layer_id: int, page_table: Tensor) -> Tuple[Tensor, Tensor]:
|
|
||||||
k_parts, v_parts = [], []
|
|
||||||
for pi in range(page_table.size(1)):
|
|
||||||
phys_pages = page_table[:, pi]
|
|
||||||
if not (phys_pages >= 0).any():
|
|
||||||
break
|
|
||||||
k_parts.append(self.k_cache[layer_id, phys_pages])
|
|
||||||
v_parts.append(self.v_cache[layer_id, phys_pages])
|
|
||||||
k = torch.cat(k_parts, dim=1)
|
|
||||||
v = torch.cat(v_parts, dim=1)
|
|
||||||
return k, v
|
|
||||||
|
|
||||||
|
|
||||||
class CacheView:
|
|
||||||
"""Per-batch view that bundles PagedCache + page_table + total_len.
|
|
||||||
|
|
||||||
Attention layers receive this as ``paged_cache`` and only see
|
|
||||||
``write()`` / ``gather()``, never raw page tables or length params.
|
|
||||||
"""
|
|
||||||
|
|
||||||
__slots__ = ("_cache", "_page_table", "_total_len")
|
|
||||||
|
|
||||||
def __init__(self, cache: PagedCache, page_table: Tensor, total_len: int = 0):
|
|
||||||
self._cache = cache
|
|
||||||
self._page_table = page_table
|
|
||||||
self._total_len = total_len
|
|
||||||
|
|
||||||
def write(self, layer_id: int, start_pos: int, k: Tensor, v: Tensor) -> None:
|
|
||||||
self._cache.write(layer_id, self._page_table, start_pos, k, v)
|
|
||||||
|
|
||||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
|
||||||
k, v = self._cache.gather(layer_id, self._page_table)
|
|
||||||
if self._total_len:
|
|
||||||
k = k[:, : self._total_len]
|
|
||||||
v = v[:, : self._total_len]
|
|
||||||
return k, v
|
|
||||||
@@ -0,0 +1,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,195 @@
|
|||||||
|
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 = 2048,
|
||||||
|
page_size: int = 64,
|
||||||
|
device: Optional[str] = None,
|
||||||
|
dtype: Optional[torch.dtype] = None,
|
||||||
|
):
|
||||||
|
config = model.config
|
||||||
|
|
||||||
|
if max_seq_len is not None:
|
||||||
|
self.max_seq_len = max_seq_len
|
||||||
|
elif config.max_len is not None:
|
||||||
|
self.max_seq_len = config.max_len
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
"max_seq_len must be provided either as argument "
|
||||||
|
"or in model config (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()
|
||||||
+129
-276
@@ -1,146 +1,55 @@
|
|||||||
"""Unified inference engine for continuous batching.
|
"""Unified inference engine for continuous batching."""
|
||||||
|
|
||||||
Layers:
|
|
||||||
- GenerationParams: Immutable value object for sampling parameters.
|
|
||||||
- GenerationRequest: User-facing request DTO with validation.
|
|
||||||
- _Result: Thread-safe token accumulator (Observer pattern).
|
|
||||||
- InferenceEngine: Facade over InferenceScheduler + async wrapper.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import gc
|
import gc
|
||||||
import threading
|
import threading
|
||||||
from dataclasses import dataclass
|
from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Tuple, Union
|
||||||
from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Union
|
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
from astrai.inference.cache import STOP
|
from astrai.inference.core.scheduler import InferenceScheduler
|
||||||
from astrai.inference.scheduler import InferenceScheduler
|
from astrai.inference.core.task import STOP
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
def _validate_sampling_params(
|
||||||
class GenerationParams:
|
top_k: int, top_p: float, temperature: float, max_tokens: Optional[int] = None
|
||||||
"""Immutable value object for sampling hyperparameters."""
|
):
|
||||||
|
if not (isinstance(top_k, int) and top_k >= 0):
|
||||||
top_k: int = 50
|
raise ValueError("top_k must be a non-negative integer")
|
||||||
top_p: float = 1.0
|
if not (0.0 <= top_p <= 1.0):
|
||||||
temperature: float = 1.0
|
raise ValueError("top_p must be a float between 0.0 and 1.0")
|
||||||
max_tokens: int = 1024
|
if not (isinstance(temperature, (int, float)) and temperature >= 0):
|
||||||
|
raise ValueError("temperature must be a non-negative number")
|
||||||
|
|
||||||
|
|
||||||
class GenerationRequest:
|
class GenerateResult:
|
||||||
"""Request parameters for text generation.
|
"""Thread-safe token accumulator for streaming and non-streaming modes."""
|
||||||
|
|
||||||
Encapsulates messages, sampling parameters (via GenerationParams),
|
|
||||||
and streaming preference for a single generation request.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
messages: List[Dict[str, str]],
|
|
||||||
top_k: int = 50,
|
|
||||||
top_p: float = 1.0,
|
|
||||||
temperature: float = 1.0,
|
|
||||||
max_len: int = 1024,
|
|
||||||
stream: bool = False,
|
|
||||||
):
|
|
||||||
"""Initializes a generation request.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
messages: Conversation history as list of {"role": ..., "content": ...}.
|
|
||||||
top_k: Top-k sampling count (0 disables).
|
|
||||||
top_p: Nucleus sampling probability threshold.
|
|
||||||
temperature: Sampling temperature.
|
|
||||||
max_len: Maximum tokens to generate.
|
|
||||||
stream: Whether to return output as a token stream.
|
|
||||||
"""
|
|
||||||
self.messages = messages
|
|
||||||
self.params = GenerationParams(
|
|
||||||
top_k=top_k,
|
|
||||||
top_p=top_p,
|
|
||||||
temperature=temperature,
|
|
||||||
max_tokens=max_len,
|
|
||||||
)
|
|
||||||
self.stream = stream
|
|
||||||
self._validate()
|
|
||||||
|
|
||||||
@property
|
|
||||||
def top_k(self) -> int:
|
|
||||||
return self.params.top_k
|
|
||||||
|
|
||||||
@property
|
|
||||||
def top_p(self) -> float:
|
|
||||||
return self.params.top_p
|
|
||||||
|
|
||||||
@property
|
|
||||||
def temperature(self) -> float:
|
|
||||||
return self.params.temperature
|
|
||||||
|
|
||||||
@property
|
|
||||||
def max_len(self) -> int:
|
|
||||||
return self.params.max_tokens
|
|
||||||
|
|
||||||
def _validate(self):
|
|
||||||
"""Validates sampling parameter ranges."""
|
|
||||||
if not (isinstance(self.top_k, int) and self.top_k >= 0):
|
|
||||||
raise ValueError("top_k must be a non-negative integer")
|
|
||||||
if not (0.0 <= self.top_p <= 1.0):
|
|
||||||
raise ValueError("top_p must be a float between 0.0 and 1.0")
|
|
||||||
if not (isinstance(self.temperature, (int, float)) and self.temperature >= 0):
|
|
||||||
raise ValueError("temperature must be a non-negative number")
|
|
||||||
|
|
||||||
|
|
||||||
class _Result:
|
|
||||||
"""Thread-safe token accumulator for streaming and non-streaming modes.
|
|
||||||
|
|
||||||
Supports multiple concurrent generation tasks with per-index result tracking.
|
|
||||||
Uses a threading.Event for efficient waiting on completion.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, count: int = 1):
|
def __init__(self, count: int = 1):
|
||||||
"""Initializes the accumulator.
|
self._cond = threading.Condition()
|
||||||
|
|
||||||
Args:
|
|
||||||
count: Number of concurrent generation tasks to track.
|
|
||||||
"""
|
|
||||||
self._lock = threading.Lock()
|
|
||||||
self._event = threading.Event()
|
self._event = threading.Event()
|
||||||
self.tokens: List[str] = []
|
self.tokens: List[Tuple[int, str]] = []
|
||||||
self.results: List[str] = [""] * count
|
self.results: List[str] = [""] * count
|
||||||
self._done: List[bool] = [False] * count
|
self._done: List[bool] = [False] * count
|
||||||
self._completed = 0
|
self._completed = 0
|
||||||
self._total = count
|
self._total = count
|
||||||
|
|
||||||
def append(self, token: str, idx: int = 0):
|
def append(self, token: str, idx: int = 0):
|
||||||
"""Appends a token to the result buffer.
|
with self._cond:
|
||||||
|
self.tokens.append((idx, token))
|
||||||
In non-streaming mode, tokens are concatenated into results[idx].
|
|
||||||
The sentinel STOP marks a task as complete.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
token: The decoded token string, or STOP sentinel.
|
|
||||||
idx: Index of the generation task this token belongs to.
|
|
||||||
"""
|
|
||||||
with self._lock:
|
|
||||||
self.tokens.append(token)
|
|
||||||
if token is not STOP:
|
if token is not STOP:
|
||||||
self.results[idx] += token
|
self.results[idx] += token
|
||||||
else:
|
else:
|
||||||
if not self._done[idx]:
|
if not self._done[idx]:
|
||||||
self._done[idx] = True
|
self._done[idx] = True
|
||||||
self._completed += 1
|
self._completed += 1
|
||||||
self._event.set()
|
self._cond.notify_all()
|
||||||
|
self._event.set()
|
||||||
|
|
||||||
def pop_all(self) -> List[str]:
|
def pop_all(self) -> List[Tuple[int, str]]:
|
||||||
"""Returns and clears all accumulated tokens.
|
with self._cond:
|
||||||
|
|
||||||
Returns:
|
|
||||||
List of token strings since the last call.
|
|
||||||
"""
|
|
||||||
with self._lock:
|
|
||||||
out = self.tokens.copy()
|
out = self.tokens.copy()
|
||||||
self.tokens.clear()
|
self.tokens.clear()
|
||||||
if not out:
|
if not out:
|
||||||
@@ -148,36 +57,47 @@ class _Result:
|
|||||||
return out
|
return out
|
||||||
|
|
||||||
def wait(self, timeout: Optional[float] = None) -> bool:
|
def wait(self, timeout: Optional[float] = None) -> bool:
|
||||||
"""Blocks until new tokens arrive or the timeout expires.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
timeout: Maximum wait time in seconds (None = infinite).
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True if the event was set (new data available), False on timeout.
|
|
||||||
"""
|
|
||||||
return self._event.wait(timeout=timeout)
|
return self._event.wait(timeout=timeout)
|
||||||
|
|
||||||
def get_results(self) -> List[str]:
|
def wait_completion(self, timeout: float = 300.0) -> None:
|
||||||
"""Returns all accumulated results for non-streaming mode.
|
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)"
|
||||||
|
)
|
||||||
|
|
||||||
Returns:
|
def get_results(self) -> List[str]:
|
||||||
List of complete generated strings, one per task index.
|
with self._cond:
|
||||||
"""
|
|
||||||
with self._lock:
|
|
||||||
return self.results.copy()
|
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:
|
class InferenceEngine:
|
||||||
"""Unified inference engine backed by continuous-batching scheduler.
|
"""Unified inference engine backed by continuous-batching scheduler."""
|
||||||
|
|
||||||
Usage:
|
|
||||||
with InferenceEngine(model, tokenizer) as engine:
|
|
||||||
for token in engine.generate("hello", stream=True):
|
|
||||||
print(token, end="")
|
|
||||||
|
|
||||||
text = engine.generate("hello")
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -188,17 +108,6 @@ class InferenceEngine:
|
|||||||
max_prompt_len: int = 2048,
|
max_prompt_len: int = 2048,
|
||||||
page_size: int = 128,
|
page_size: int = 128,
|
||||||
):
|
):
|
||||||
"""Initializes the inference engine.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
model: The model instance.
|
|
||||||
tokenizer: The tokenizer instance.
|
|
||||||
max_batch_size: Maximum number of concurrent tasks.
|
|
||||||
max_seq_len: Maximum sequence length.
|
|
||||||
max_prompt_len: Maximum prompt tokens.
|
|
||||||
compile: Whether to compile the model with torch.compile.
|
|
||||||
page_size: Number of tokens per KV cache page.
|
|
||||||
"""
|
|
||||||
self.model = model
|
self.model = model
|
||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
self.scheduler = InferenceScheduler(
|
self.scheduler = InferenceScheduler(
|
||||||
@@ -223,25 +132,12 @@ class InferenceEngine:
|
|||||||
self,
|
self,
|
||||||
prompt: Union[str, List[str]],
|
prompt: Union[str, List[str]],
|
||||||
stream: bool = False,
|
stream: bool = False,
|
||||||
max_tokens: int = 1024,
|
max_tokens: Optional[int] = None,
|
||||||
temperature: float = 1.0,
|
temperature: float = 1.0,
|
||||||
top_p: float = 1.0,
|
top_p: float = 1.0,
|
||||||
top_k: int = 50,
|
top_k: int = 50,
|
||||||
) -> Union[Generator[str, None, None], str, List[str]]:
|
) -> Union[Generator, str, List[str]]:
|
||||||
"""Generates text from a prompt.
|
_validate_sampling_params(top_k, top_p, temperature, max_tokens)
|
||||||
|
|
||||||
Args:
|
|
||||||
prompt: Single string or list of strings for batch generation.
|
|
||||||
stream: If True, returns a generator yielding tokens one by one.
|
|
||||||
max_tokens: Maximum number of tokens to generate.
|
|
||||||
temperature: Sampling temperature.
|
|
||||||
top_p: Nucleus sampling probability threshold.
|
|
||||||
top_k: Top-k sampling count (0 disables).
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Generator (stream=True), single string (non-stream, single prompt),
|
|
||||||
or list of strings (non-stream, batch prompts).
|
|
||||||
"""
|
|
||||||
is_batch = isinstance(prompt, list)
|
is_batch = isinstance(prompt, list)
|
||||||
prompts = prompt if is_batch else [prompt]
|
prompts = prompt if is_batch else [prompt]
|
||||||
|
|
||||||
@@ -257,26 +153,12 @@ class InferenceEngine:
|
|||||||
def generate_async(
|
def generate_async(
|
||||||
self,
|
self,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
max_tokens: int = 1024,
|
max_tokens: Optional[int] = None,
|
||||||
temperature: float = 1.0,
|
temperature: float = 1.0,
|
||||||
top_p: float = 1.0,
|
top_p: float = 1.0,
|
||||||
top_k: int = 50,
|
top_k: int = 50,
|
||||||
) -> AsyncGenerator[str, None]:
|
) -> AsyncGenerator[str, None]:
|
||||||
"""Async streaming generator that does not block the event loop.
|
_validate_sampling_params(top_k, top_p, temperature, max_tokens)
|
||||||
|
|
||||||
Runs the synchronous generator in a background thread pool executor,
|
|
||||||
yielding tokens to the async consumer as they arrive.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
prompt: Input text to generate from.
|
|
||||||
max_tokens: Maximum tokens to generate.
|
|
||||||
temperature: Sampling temperature.
|
|
||||||
top_p: Nucleus sampling threshold.
|
|
||||||
top_k: Top-k sampling count.
|
|
||||||
|
|
||||||
Yields:
|
|
||||||
Decoded token strings as they are generated.
|
|
||||||
"""
|
|
||||||
sync_gen = self._generate_streaming(
|
sync_gen = self._generate_streaming(
|
||||||
[prompt], False, max_tokens, temperature, top_p, top_k
|
[prompt], False, max_tokens, temperature, top_p, top_k
|
||||||
)
|
)
|
||||||
@@ -293,14 +175,6 @@ class InferenceEngine:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _next_token(gen: Generator) -> Optional[str]:
|
def _next_token(gen: Generator) -> Optional[str]:
|
||||||
"""Retrieves the next token from a synchronous generator.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
gen: A synchronous generator yielding token strings.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
The next token, or None if the generator is exhausted.
|
|
||||||
"""
|
|
||||||
try:
|
try:
|
||||||
return next(gen)
|
return next(gen)
|
||||||
except StopIteration:
|
except StopIteration:
|
||||||
@@ -309,77 +183,80 @@ class InferenceEngine:
|
|||||||
def generate_with_request(
|
def generate_with_request(
|
||||||
self, request: GenerationRequest
|
self, request: GenerationRequest
|
||||||
) -> Union[Generator[str, None, None], str, List[str]]:
|
) -> Union[Generator[str, None, None], str, List[str]]:
|
||||||
"""Generates text from a structured GenerationRequest.
|
|
||||||
|
|
||||||
Applies the chat template to the request's messages before generation.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
request: A GenerationRequest with messages and parameters.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Generator, string, or list of strings (see generate()).
|
|
||||||
"""
|
|
||||||
prompt = self.tokenizer.apply_chat_template(request.messages, tokenize=False)
|
prompt = self.tokenizer.apply_chat_template(request.messages, tokenize=False)
|
||||||
return self.generate(
|
return self.generate(
|
||||||
prompt=prompt,
|
prompt=prompt,
|
||||||
stream=request.stream,
|
stream=request.stream,
|
||||||
max_tokens=request.params.max_tokens,
|
max_tokens=request.max_tokens,
|
||||||
temperature=request.params.temperature,
|
temperature=request.temperature,
|
||||||
top_p=request.params.top_p,
|
top_p=request.top_p,
|
||||||
top_k=request.params.top_k,
|
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(
|
def _generate_streaming(
|
||||||
self,
|
self,
|
||||||
prompts: List[str],
|
prompts: List[str],
|
||||||
is_batch: bool,
|
is_batch: bool,
|
||||||
max_tokens: int,
|
max_tokens: Optional[int],
|
||||||
temperature: float,
|
temperature: float,
|
||||||
top_p: float,
|
top_p: float,
|
||||||
top_k: int,
|
top_k: int,
|
||||||
) -> Generator[str, None, None]:
|
) -> Generator:
|
||||||
"""Internal streaming generator.
|
result, task_ids = self._submit_tasks(
|
||||||
|
prompts, max_tokens, temperature, top_p, top_k
|
||||||
Polls the _Result accumulator in a loop, yielding tokens as they arrive.
|
|
||||||
Cleans up the scheduler task on GeneratorExit.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
prompts: List of prompts (only first is used; batch not yet supported).
|
|
||||||
is_batch: If True, raises NotImplementedError.
|
|
||||||
max_tokens: Maximum tokens to generate.
|
|
||||||
temperature: Sampling temperature.
|
|
||||||
top_p: Nucleus sampling threshold.
|
|
||||||
top_k: Top-k sampling count.
|
|
||||||
|
|
||||||
Yields:
|
|
||||||
Decoded token strings.
|
|
||||||
"""
|
|
||||||
if is_batch:
|
|
||||||
raise NotImplementedError("Batch streaming not yet supported")
|
|
||||||
|
|
||||||
result = _Result()
|
|
||||||
|
|
||||||
task_id = self.scheduler.add_task(
|
|
||||||
prompt=prompts[0],
|
|
||||||
max_tokens=max_tokens,
|
|
||||||
temperature=temperature,
|
|
||||||
top_p=top_p,
|
|
||||||
top_k=top_k,
|
|
||||||
stream_callback=lambda tok: result.append(tok, 0),
|
|
||||||
)
|
)
|
||||||
|
n = len(prompts)
|
||||||
|
remaining = n
|
||||||
|
finished = [False] * n
|
||||||
|
|
||||||
def gen():
|
def gen():
|
||||||
|
nonlocal remaining
|
||||||
try:
|
try:
|
||||||
while True:
|
while remaining > 0:
|
||||||
tokens = result.pop_all()
|
items = result.pop_all()
|
||||||
for token in tokens:
|
for idx, token in items:
|
||||||
if token is STOP:
|
if token is STOP:
|
||||||
return
|
if not finished[idx]:
|
||||||
yield token
|
finished[idx] = True
|
||||||
if not result.wait(timeout=0.05):
|
remaining -= 1
|
||||||
pass
|
else:
|
||||||
|
yield (idx, token) if is_batch else token
|
||||||
|
if remaining > 0:
|
||||||
|
result.wait(timeout=0.05)
|
||||||
finally:
|
finally:
|
||||||
self.scheduler.remove_task(task_id)
|
for tid in task_ids:
|
||||||
|
self.scheduler.remove_task(tid)
|
||||||
|
|
||||||
return gen()
|
return gen()
|
||||||
|
|
||||||
@@ -387,56 +264,32 @@ class InferenceEngine:
|
|||||||
self,
|
self,
|
||||||
prompts: List[str],
|
prompts: List[str],
|
||||||
is_batch: bool,
|
is_batch: bool,
|
||||||
max_tokens: int,
|
max_tokens: Optional[int],
|
||||||
temperature: float,
|
temperature: float,
|
||||||
top_p: float,
|
top_p: float,
|
||||||
top_k: int,
|
top_k: int,
|
||||||
) -> Union[str, List[str]]:
|
) -> Union[str, List[str]]:
|
||||||
"""Internal non-streaming generator.
|
result, task_ids = self._submit_tasks(
|
||||||
|
prompts, max_tokens, temperature, top_p, top_k
|
||||||
|
)
|
||||||
|
|
||||||
Submits all prompts to the scheduler and waits for all to complete.
|
try:
|
||||||
|
result.wait_completion()
|
||||||
|
except TimeoutError:
|
||||||
|
for tid in task_ids:
|
||||||
|
self.scheduler.remove_task(tid)
|
||||||
|
raise
|
||||||
|
|
||||||
Args:
|
for tid in task_ids:
|
||||||
prompts: List of prompt strings.
|
self.scheduler.remove_task(tid)
|
||||||
is_batch: Whether multiple prompts were provided.
|
|
||||||
max_tokens: Maximum tokens to generate.
|
|
||||||
temperature: Sampling temperature.
|
|
||||||
top_p: Nucleus sampling threshold.
|
|
||||||
top_k: Top-k sampling count.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Single string for one prompt, list of strings for batch.
|
|
||||||
"""
|
|
||||||
result = _Result(count=len(prompts))
|
|
||||||
|
|
||||||
for i, p in enumerate(prompts):
|
|
||||||
|
|
||||||
def make_cb(idx):
|
|
||||||
return lambda tok: result.append(tok, idx)
|
|
||||||
|
|
||||||
self.scheduler.add_task(
|
|
||||||
prompt=p,
|
|
||||||
max_tokens=max_tokens,
|
|
||||||
temperature=temperature,
|
|
||||||
top_p=top_p,
|
|
||||||
top_k=top_k,
|
|
||||||
stream_callback=make_cb(i),
|
|
||||||
)
|
|
||||||
|
|
||||||
result.wait()
|
|
||||||
res = result.get_results()
|
res = result.get_results()
|
||||||
return res if is_batch else res[0]
|
return res if is_batch else res[0]
|
||||||
|
|
||||||
def get_stats(self) -> Dict[str, Any]:
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
"""Returns current engine statistics.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Dict with total_tasks, total_tokens, active_tasks, waiting_queue.
|
|
||||||
"""
|
|
||||||
return self.scheduler.get_stats()
|
return self.scheduler.get_stats()
|
||||||
|
|
||||||
def shutdown(self) -> None:
|
def shutdown(self) -> None:
|
||||||
"""Shuts down the engine, stops the scheduler, and frees GPU memory."""
|
|
||||||
self.scheduler.stop()
|
self.scheduler.stop()
|
||||||
if torch.cuda.is_available():
|
if torch.cuda.is_available():
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
|
|||||||
@@ -64,16 +64,26 @@ class TopKStrategy(BaseSamplingStrategy):
|
|||||||
def apply(self, logits, filter_value=-float("inf")):
|
def apply(self, logits, filter_value=-float("inf")):
|
||||||
tk = self.top_k
|
tk = self.top_k
|
||||||
if isinstance(tk, Tensor):
|
if isinstance(tk, Tensor):
|
||||||
|
tk = tk.to(logits.device, non_blocking=True).long().clamp(min=0)
|
||||||
max_k = int(tk.max().item())
|
max_k = int(tk.max().item())
|
||||||
if max_k <= 0:
|
if max_k <= 0:
|
||||||
return logits
|
return logits
|
||||||
k = min(max_k, logits.size(-1))
|
max_k = min(max_k, logits.size(-1))
|
||||||
elif tk > 0:
|
values, _ = torch.topk(logits, max_k, dim=-1)
|
||||||
k = min(tk, logits.size(-1))
|
per_row_k = tk.clamp(max=max_k)
|
||||||
else:
|
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
|
return logits
|
||||||
thresholds = torch.topk(logits, k, dim=-1)[0][..., -1:]
|
if tk > 0:
|
||||||
logits[logits < thresholds] = filter_value
|
k = min(tk, logits.size(-1))
|
||||||
|
thresholds = torch.topk(logits, k, dim=-1)[0][..., -1:]
|
||||||
|
logits[logits < thresholds] = filter_value
|
||||||
return logits
|
return logits
|
||||||
|
|
||||||
|
|
||||||
@@ -1,384 +0,0 @@
|
|||||||
"""Inference scheduler for single-GPU continuous batching with paged KV cache."""
|
|
||||||
|
|
||||||
import logging
|
|
||||||
import threading
|
|
||||||
import time
|
|
||||||
import uuid
|
|
||||||
from enum import Enum
|
|
||||||
from typing import Any, Callable, Dict, List, Optional
|
|
||||||
|
|
||||||
import torch
|
|
||||||
from torch import Tensor
|
|
||||||
|
|
||||||
from astrai.inference.cache import STOP, PagedCache
|
|
||||||
from astrai.inference.sampling import sample
|
|
||||||
from astrai.model.automodel import AutoModel
|
|
||||||
from astrai.tokenize.tokenizer import AutoTokenizer
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
|
|
||||||
class TaskStatus(Enum):
|
|
||||||
"""Task states in the continuous batching lifecycle."""
|
|
||||||
|
|
||||||
PENDING = "pending"
|
|
||||||
RUNNING = "running"
|
|
||||||
FINISHED = "finished"
|
|
||||||
ABORTED = "aborted"
|
|
||||||
|
|
||||||
|
|
||||||
class Task:
|
|
||||||
"""Represents a single generation request with paged KV cache tracking."""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
task_id: str,
|
|
||||||
prompt_ids: List[int],
|
|
||||||
max_tokens: int = 1024,
|
|
||||||
temperature: float = 1.0,
|
|
||||||
top_p: float = 1.0,
|
|
||||||
top_k: int = 50,
|
|
||||||
stream_callback: Optional[Callable[[str], None]] = None,
|
|
||||||
):
|
|
||||||
self.task_id = task_id
|
|
||||||
self.prompt_ids = prompt_ids
|
|
||||||
self.max_tokens = max_tokens
|
|
||||||
self.temperature = temperature
|
|
||||||
self.top_p = top_p
|
|
||||||
self.top_k = top_k
|
|
||||||
|
|
||||||
self.status = TaskStatus.PENDING
|
|
||||||
self.output_ids: List[int] = []
|
|
||||||
self.input_tokens: int = 0
|
|
||||||
self.output_tokens: int = 0
|
|
||||||
self.page_table: List[int] = []
|
|
||||||
self.n_pages: int = 0
|
|
||||||
self.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.output_tokens >= self.max_tokens:
|
|
||||||
return True
|
|
||||||
if self.output_ids and self.output_ids[-1] in stop_ids:
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
|
|
||||||
|
|
||||||
class InferenceScheduler:
|
|
||||||
"""Continuous batching scheduler with paged KV cache.
|
|
||||||
|
|
||||||
Runs a background generation loop with four phases per iteration:
|
|
||||||
1. Cleanup finished tasks and release resources.
|
|
||||||
2. Refill active batch from the waiting queue.
|
|
||||||
3. Prefill newly activated tasks.
|
|
||||||
4. Decode the largest same-position group of active tasks.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
model: AutoModel,
|
|
||||||
tokenizer: AutoTokenizer,
|
|
||||||
max_batch_size: int = 16,
|
|
||||||
max_seq_len: Optional[int] = None,
|
|
||||||
max_prompt_len: int = 512,
|
|
||||||
page_size: int = 64,
|
|
||||||
device: str = "cuda",
|
|
||||||
dtype: torch.dtype = torch.bfloat16,
|
|
||||||
):
|
|
||||||
config = model.config
|
|
||||||
|
|
||||||
self.model = model
|
|
||||||
self.tokenizer = tokenizer
|
|
||||||
self.max_batch_size = max_batch_size
|
|
||||||
self.max_seq_len = max_seq_len or config.max_len
|
|
||||||
self.max_prompt_len = max_prompt_len
|
|
||||||
self.page_size = page_size
|
|
||||||
self.device = device or next(model.parameters()).device
|
|
||||||
self.dtype = dtype or next(model.parameters()).dtype
|
|
||||||
|
|
||||||
n_kv_heads = config.n_kv_heads
|
|
||||||
head_dim = config.dim // config.n_heads
|
|
||||||
n_layers = config.n_layers
|
|
||||||
n_pages = (max_batch_size * self.max_seq_len + page_size - 1) // page_size
|
|
||||||
|
|
||||||
self.page_cache = PagedCache(
|
|
||||||
n_layers,
|
|
||||||
n_pages,
|
|
||||||
page_size,
|
|
||||||
n_kv_heads,
|
|
||||||
head_dim,
|
|
||||||
self.device,
|
|
||||||
self.dtype,
|
|
||||||
)
|
|
||||||
|
|
||||||
self.waiting_queue: List[Task] = []
|
|
||||||
self.active_tasks: List[Task] = []
|
|
||||||
|
|
||||||
self._running = False
|
|
||||||
self._task_event = threading.Event()
|
|
||||||
self._lock = threading.Lock()
|
|
||||||
|
|
||||||
self._total_tasks = 0
|
|
||||||
self._total_tokens = 0
|
|
||||||
|
|
||||||
def _n_pages_for(self, n_tokens: int) -> int:
|
|
||||||
return (n_tokens + self.page_size - 1) // self.page_size
|
|
||||||
|
|
||||||
def add_task(
|
|
||||||
self,
|
|
||||||
prompt: str,
|
|
||||||
max_tokens: int = 1024,
|
|
||||||
temperature: float = 1.0,
|
|
||||||
top_p: float = 1.0,
|
|
||||||
top_k: int = 50,
|
|
||||||
stream_callback: Optional[Callable[[str], None]] = None,
|
|
||||||
) -> str:
|
|
||||||
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
|
|
||||||
prompt_ids = self.tokenizer.encode(prompt)
|
|
||||||
if len(prompt_ids) > self.max_prompt_len:
|
|
||||||
prompt_ids = prompt_ids[-self.max_prompt_len :]
|
|
||||||
|
|
||||||
task = Task(
|
|
||||||
task_id=task_id,
|
|
||||||
prompt_ids=prompt_ids,
|
|
||||||
max_tokens=max_tokens,
|
|
||||||
temperature=temperature,
|
|
||||||
top_p=top_p,
|
|
||||||
top_k=top_k,
|
|
||||||
stream_callback=stream_callback,
|
|
||||||
)
|
|
||||||
|
|
||||||
with self._lock:
|
|
||||||
self.waiting_queue.append(task)
|
|
||||||
self._total_tasks += 1
|
|
||||||
|
|
||||||
self._task_event.set()
|
|
||||||
return task_id
|
|
||||||
|
|
||||||
def remove_task(self, task_id: str) -> None:
|
|
||||||
with self._lock:
|
|
||||||
removed_active = [t for t in self.active_tasks if t.task_id == task_id]
|
|
||||||
self.waiting_queue = [t for t in self.waiting_queue if t.task_id != task_id]
|
|
||||||
self.active_tasks = [t for t in self.active_tasks if t.task_id != task_id]
|
|
||||||
|
|
||||||
for task in removed_active:
|
|
||||||
self._free_pages(task.page_table)
|
|
||||||
task.page_table.clear()
|
|
||||||
task.n_pages = 0
|
|
||||||
|
|
||||||
def _free_pages(self, indices: List[int]) -> None:
|
|
||||||
for idx in indices:
|
|
||||||
self.page_cache.free(idx)
|
|
||||||
|
|
||||||
def _remove_finished_tasks(self) -> None:
|
|
||||||
finished = []
|
|
||||||
for task in self.active_tasks:
|
|
||||||
if task.is_finished(self.tokenizer.stop_ids):
|
|
||||||
task.status = TaskStatus.FINISHED
|
|
||||||
task.finish_time = time.time()
|
|
||||||
finished.append(task)
|
|
||||||
self._total_tokens += task.output_tokens
|
|
||||||
|
|
||||||
for task in finished:
|
|
||||||
self._free_pages(task.page_table)
|
|
||||||
task.page_table.clear()
|
|
||||||
task.n_pages = 0
|
|
||||||
|
|
||||||
self.active_tasks = [
|
|
||||||
t for t in self.active_tasks if t.status != TaskStatus.FINISHED
|
|
||||||
]
|
|
||||||
|
|
||||||
def _refill_active_batch(self) -> None:
|
|
||||||
available = self.max_batch_size - len(self.active_tasks)
|
|
||||||
if available <= 0:
|
|
||||||
return
|
|
||||||
|
|
||||||
to_add: List[Task] = []
|
|
||||||
with self._lock:
|
|
||||||
n = min(available, len(self.waiting_queue))
|
|
||||||
for _ in range(n):
|
|
||||||
to_add.append(self.waiting_queue.pop(0))
|
|
||||||
|
|
||||||
failed: List[Task] = []
|
|
||||||
for task in to_add:
|
|
||||||
prompt_len = len(task.prompt_ids)
|
|
||||||
n_pages = self._n_pages_for(prompt_len)
|
|
||||||
task.page_table = self.page_cache.alloc_n(n_pages)
|
|
||||||
if not task.page_table:
|
|
||||||
failed.append(task)
|
|
||||||
continue
|
|
||||||
task.n_pages = len(task.page_table)
|
|
||||||
task.status = TaskStatus.RUNNING
|
|
||||||
self.active_tasks.append(task)
|
|
||||||
|
|
||||||
if failed:
|
|
||||||
with self._lock:
|
|
||||||
self.waiting_queue[:0] = failed
|
|
||||||
|
|
||||||
def _execute_prefill(self) -> None:
|
|
||||||
to_prefill = [t for t in self.active_tasks if t.output_tokens == 0]
|
|
||||||
if not to_prefill:
|
|
||||||
return
|
|
||||||
|
|
||||||
for t in to_prefill:
|
|
||||||
prompt_len = len(t.prompt_ids)
|
|
||||||
t.input_tokens = prompt_len
|
|
||||||
t.output_tokens = 0
|
|
||||||
|
|
||||||
groups: Dict[int, List[Task]] = {}
|
|
||||||
for t in to_prefill:
|
|
||||||
groups.setdefault(len(t.prompt_ids), []).append(t)
|
|
||||||
|
|
||||||
for prompt_len, group in groups.items():
|
|
||||||
self._execute_prefill_batch(group, prompt_len)
|
|
||||||
|
|
||||||
def _execute_prefill_batch(self, tasks: List[Task], prompt_len: int) -> None:
|
|
||||||
tasks = sorted(tasks, key=lambda t: t.task_id)
|
|
||||||
batch_sz = len(tasks)
|
|
||||||
|
|
||||||
input_ids = torch.zeros(
|
|
||||||
batch_sz,
|
|
||||||
prompt_len,
|
|
||||||
dtype=torch.long,
|
|
||||||
device=self.device,
|
|
||||||
)
|
|
||||||
input_mask = torch.ones(
|
|
||||||
batch_sz,
|
|
||||||
prompt_len,
|
|
||||||
dtype=torch.bool,
|
|
||||||
device=self.device,
|
|
||||||
)
|
|
||||||
|
|
||||||
for i, t in enumerate(tasks):
|
|
||||||
input_ids[i] = torch.tensor(t.prompt_ids, device=self.device)
|
|
||||||
|
|
||||||
page_tables = self._make_page_table_tensor(tasks)
|
|
||||||
|
|
||||||
with torch.inference_mode():
|
|
||||||
self.model(
|
|
||||||
input_ids,
|
|
||||||
input_mask=input_mask,
|
|
||||||
start_pos=0,
|
|
||||||
paged_cache=self.page_cache.bind(page_tables, total_len=prompt_len),
|
|
||||||
)
|
|
||||||
|
|
||||||
def _execute_decode(self, tasks: List[Task], start_pos: int) -> None:
|
|
||||||
if not tasks:
|
|
||||||
return
|
|
||||||
|
|
||||||
tasks = sorted(tasks, key=lambda t: t.task_id)
|
|
||||||
batch_sz = len(tasks)
|
|
||||||
|
|
||||||
input_ids = torch.zeros(batch_sz, dtype=torch.long, device=self.device)
|
|
||||||
for i, t in enumerate(tasks):
|
|
||||||
input_ids[i] = t.output_ids[-1] if t.output_ids else t.prompt_ids[-1]
|
|
||||||
|
|
||||||
active_mask = torch.ones((batch_sz, 1), dtype=torch.bool, device=self.device)
|
|
||||||
|
|
||||||
page_tables = self._make_page_table_tensor(tasks)
|
|
||||||
total_len = start_pos + 1
|
|
||||||
|
|
||||||
with torch.inference_mode():
|
|
||||||
outputs = self.model(
|
|
||||||
input_ids.unsqueeze(1),
|
|
||||||
input_mask=active_mask,
|
|
||||||
paged_cache=self.page_cache.bind(page_tables, total_len=total_len),
|
|
||||||
start_pos=start_pos,
|
|
||||||
)
|
|
||||||
logits = outputs["logits"][:, -1, :]
|
|
||||||
|
|
||||||
next_tokens = sample(
|
|
||||||
logits,
|
|
||||||
temperature=torch.tensor(
|
|
||||||
[t.temperature for t in tasks], device=logits.device
|
|
||||||
),
|
|
||||||
top_k=torch.tensor([t.top_k for t in tasks], device=logits.device),
|
|
||||||
top_p=torch.tensor([t.top_p for t in tasks], device=logits.device),
|
|
||||||
).tolist()
|
|
||||||
|
|
||||||
for t, ntok in zip(tasks, next_tokens):
|
|
||||||
t.output_ids.append(ntok)
|
|
||||||
t.output_tokens += 1
|
|
||||||
pos = t.input_tokens + t.output_tokens
|
|
||||||
self._maybe_alloc_page(t, pos)
|
|
||||||
if t.stream_callback:
|
|
||||||
t.stream_callback(self.tokenizer.decode([ntok]))
|
|
||||||
|
|
||||||
for t in tasks:
|
|
||||||
if t.is_finished(self.tokenizer.stop_ids):
|
|
||||||
if t.stream_callback:
|
|
||||||
t.stream_callback(STOP)
|
|
||||||
|
|
||||||
def _make_page_table_tensor(self, tasks: List[Task]) -> Tensor:
|
|
||||||
max_pages = max(t.n_pages for t in tasks)
|
|
||||||
rows = [t.page_table + [-1] * (max_pages - t.n_pages) for t in tasks]
|
|
||||||
return torch.tensor(rows, dtype=torch.long, device=self.device)
|
|
||||||
|
|
||||||
def _maybe_alloc_page(self, task: Task, pos: int) -> None:
|
|
||||||
needed = self._n_pages_for(pos + 1)
|
|
||||||
while task.n_pages < needed:
|
|
||||||
p = self.page_cache.alloc()
|
|
||||||
if p < 0:
|
|
||||||
break
|
|
||||||
task.page_table.append(p)
|
|
||||||
task.n_pages += 1
|
|
||||||
|
|
||||||
def _run_generation_loop(self) -> None:
|
|
||||||
try:
|
|
||||||
while self._running:
|
|
||||||
self._remove_finished_tasks()
|
|
||||||
self._refill_active_batch()
|
|
||||||
|
|
||||||
if not self.active_tasks and not self.waiting_queue:
|
|
||||||
self._task_event.clear()
|
|
||||||
self._task_event.wait(timeout=1.0)
|
|
||||||
continue
|
|
||||||
|
|
||||||
self._execute_prefill()
|
|
||||||
|
|
||||||
pos_groups: Dict[int, List[Task]] = {}
|
|
||||||
for t in self.active_tasks:
|
|
||||||
pos_groups.setdefault(t.next_pos, []).append(t)
|
|
||||||
|
|
||||||
if pos_groups:
|
|
||||||
best_pos = max(pos_groups, key=lambda p: len(pos_groups[p]))
|
|
||||||
self._execute_decode(pos_groups[best_pos], best_pos)
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
|
||||||
for task in self.active_tasks:
|
|
||||||
if task.stream_callback:
|
|
||||||
task.stream_callback(STOP)
|
|
||||||
for task in self.waiting_queue:
|
|
||||||
if task.stream_callback:
|
|
||||||
task.stream_callback(STOP)
|
|
||||||
raise
|
|
||||||
|
|
||||||
def start(self) -> None:
|
|
||||||
if not self._running:
|
|
||||||
self._running = True
|
|
||||||
t = threading.Thread(target=self._run_generation_loop, daemon=True)
|
|
||||||
t.start()
|
|
||||||
self._loop_thread = t
|
|
||||||
|
|
||||||
def stop(self) -> None:
|
|
||||||
self._running = False
|
|
||||||
self._task_event.set()
|
|
||||||
if hasattr(self, "_loop_thread"):
|
|
||||||
self._loop_thread.join(timeout=2.0)
|
|
||||||
self.waiting_queue.clear()
|
|
||||||
self.active_tasks.clear()
|
|
||||||
if torch.cuda.is_available():
|
|
||||||
torch.cuda.empty_cache()
|
|
||||||
|
|
||||||
def get_stats(self) -> Dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"total_tasks": self._total_tasks,
|
|
||||||
"total_tokens": self._total_tokens,
|
|
||||||
"active_tasks": len(self.active_tasks),
|
|
||||||
"waiting_queue": len(self.waiting_queue),
|
|
||||||
}
|
|
||||||
@@ -1,486 +0,0 @@
|
|||||||
"""
|
|
||||||
OpenAI / Anthropic-compatible chat completion server backed by continuous-batching inference.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import json
|
|
||||||
import logging
|
|
||||||
import time
|
|
||||||
import uuid
|
|
||||||
from contextlib import asynccontextmanager
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any, Dict, List, Optional, Union
|
|
||||||
|
|
||||||
import torch
|
|
||||||
import uvicorn
|
|
||||||
from fastapi import FastAPI, HTTPException
|
|
||||||
from fastapi.responses import StreamingResponse
|
|
||||||
from pydantic import BaseModel, Field
|
|
||||||
|
|
||||||
from astrai.inference.engine import InferenceEngine
|
|
||||||
from astrai.model import AutoModel
|
|
||||||
from astrai.tokenize import AutoTokenizer
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
_project_root = Path(__file__).parent.parent.parent
|
|
||||||
|
|
||||||
|
|
||||||
class ServerState:
|
|
||||||
def __init__(self):
|
|
||||||
self.engine: Optional[InferenceEngine] = None
|
|
||||||
self.config: Dict[str, Any] = {
|
|
||||||
"device": "cuda",
|
|
||||||
"dtype": torch.bfloat16,
|
|
||||||
"param_path": None,
|
|
||||||
"max_batch_size": 16,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
_state = ServerState()
|
|
||||||
|
|
||||||
|
|
||||||
class ChatMessage(BaseModel):
|
|
||||||
role: str
|
|
||||||
content: str
|
|
||||||
|
|
||||||
|
|
||||||
class ChatCompletionRequest(BaseModel):
|
|
||||||
"""OpenAI Chat Completion API request body."""
|
|
||||||
|
|
||||||
model: str = "astrai"
|
|
||||||
messages: List[ChatMessage]
|
|
||||||
temperature: Optional[float] = Field(default=1.0, ge=0.0, le=2.0)
|
|
||||||
top_p: Optional[float] = Field(default=1.0, ge=0.0, le=1.0)
|
|
||||||
top_k: Optional[int] = Field(default=50, ge=1)
|
|
||||||
stream: Optional[bool] = False
|
|
||||||
stop: Optional[Union[str, List[str]]] = None
|
|
||||||
max_tokens: Optional[int] = Field(default=2048, ge=1)
|
|
||||||
n: Optional[int] = Field(default=1, ge=1)
|
|
||||||
presence_penalty: Optional[float] = Field(default=0.0, ge=-2.0, le=2.0)
|
|
||||||
frequency_penalty: Optional[float] = Field(default=0.0, ge=-2.0, le=2.0)
|
|
||||||
logit_bias: Optional[Dict[int, float]] = None
|
|
||||||
user: Optional[str] = None
|
|
||||||
|
|
||||||
|
|
||||||
class AnthropicMessage(BaseModel):
|
|
||||||
role: str
|
|
||||||
content: Union[str, List[Dict[str, Any]]]
|
|
||||||
|
|
||||||
|
|
||||||
class MessagesRequest(BaseModel):
|
|
||||||
"""Anthropic Messages API request body."""
|
|
||||||
|
|
||||||
model: str = "astrai"
|
|
||||||
max_tokens: int = Field(default=1024, ge=1)
|
|
||||||
messages: List[AnthropicMessage]
|
|
||||||
system: Optional[str] = None
|
|
||||||
temperature: Optional[float] = Field(default=1.0, ge=0.0, le=2.0)
|
|
||||||
top_p: Optional[float] = Field(default=1.0, ge=0.0, le=1.0)
|
|
||||||
top_k: Optional[int] = Field(default=50, ge=1)
|
|
||||||
stream: Optional[bool] = False
|
|
||||||
stop_sequences: Optional[List[str]] = None
|
|
||||||
|
|
||||||
|
|
||||||
def configure_server(
|
|
||||||
device: str = "cuda",
|
|
||||||
dtype: torch.dtype = torch.bfloat16,
|
|
||||||
param_path: Optional[Path] = None,
|
|
||||||
max_batch_size: int = 16,
|
|
||||||
):
|
|
||||||
_state.config.update(
|
|
||||||
device=device,
|
|
||||||
dtype=dtype,
|
|
||||||
param_path=param_path,
|
|
||||||
max_batch_size=max_batch_size,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@asynccontextmanager
|
|
||||||
async def lifespan(app: FastAPI):
|
|
||||||
try:
|
|
||||||
load_model(
|
|
||||||
param_path=_state.config["param_path"],
|
|
||||||
device=_state.config["device"],
|
|
||||||
dtype=_state.config["dtype"],
|
|
||||||
max_batch_size=_state.config["max_batch_size"],
|
|
||||||
)
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Failed to load model: {e}")
|
|
||||||
raise
|
|
||||||
yield
|
|
||||||
if _state.engine:
|
|
||||||
_state.engine.shutdown()
|
|
||||||
logger.info("Inference engine shutdown complete")
|
|
||||||
|
|
||||||
|
|
||||||
app = FastAPI(title="AstrAI Inference Server", version="0.2.0", lifespan=lifespan)
|
|
||||||
|
|
||||||
|
|
||||||
def load_model(
|
|
||||||
param_path: Optional[Path] = None,
|
|
||||||
device: str = "cuda",
|
|
||||||
dtype: torch.dtype = torch.bfloat16,
|
|
||||||
max_batch_size: int = 16,
|
|
||||||
):
|
|
||||||
if param_path is None:
|
|
||||||
param_path = _project_root / "params"
|
|
||||||
if not param_path.exists():
|
|
||||||
raise FileNotFoundError(f"Parameter directory not found: {param_path}")
|
|
||||||
|
|
||||||
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
|
||||||
model = AutoModel.from_pretrained(param_path)
|
|
||||||
model.to(device=device, dtype=dtype)
|
|
||||||
logger.info(f"Model loaded on {device} with dtype {dtype}")
|
|
||||||
|
|
||||||
_state.engine = InferenceEngine(
|
|
||||||
model=model,
|
|
||||||
tokenizer=tokenizer,
|
|
||||||
max_batch_size=max_batch_size,
|
|
||||||
)
|
|
||||||
logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}")
|
|
||||||
|
|
||||||
|
|
||||||
def _get_engine() -> InferenceEngine:
|
|
||||||
if _state.engine is None:
|
|
||||||
raise HTTPException(status_code=503, detail="Engine not initialized")
|
|
||||||
return _state.engine
|
|
||||||
|
|
||||||
|
|
||||||
def _make_chunk(
|
|
||||||
delta: Dict[str, str],
|
|
||||||
finish_reason: Optional[str] = None,
|
|
||||||
*,
|
|
||||||
resp_id: str,
|
|
||||||
created: int,
|
|
||||||
model: str,
|
|
||||||
index: int = 0,
|
|
||||||
) -> str:
|
|
||||||
"""Build a single SSE ``data:`` chunk matching OpenAI streaming format."""
|
|
||||||
data = {
|
|
||||||
"id": resp_id,
|
|
||||||
"object": "chat.completion.chunk",
|
|
||||||
"created": created,
|
|
||||||
"model": model,
|
|
||||||
"choices": [
|
|
||||||
{
|
|
||||||
"index": index,
|
|
||||||
"delta": delta,
|
|
||||||
"finish_reason": finish_reason,
|
|
||||||
}
|
|
||||||
],
|
|
||||||
}
|
|
||||||
return f"data: {json.dumps(data, ensure_ascii=False)}\n\n"
|
|
||||||
|
|
||||||
|
|
||||||
@app.get("/health")
|
|
||||||
async def health():
|
|
||||||
return {
|
|
||||||
"status": "ok",
|
|
||||||
"model_loaded": _state.engine is not None,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
@app.get("/stats")
|
|
||||||
async def get_stats():
|
|
||||||
return _get_engine().get_stats()
|
|
||||||
|
|
||||||
|
|
||||||
@app.post("/v1/chat/completions")
|
|
||||||
async def chat_completion(request: ChatCompletionRequest):
|
|
||||||
"""OpenAI-compatible chat completion endpoint (streaming + non-streaming)."""
|
|
||||||
engine = _get_engine()
|
|
||||||
resp_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
|
||||||
created = int(time.time())
|
|
||||||
model = request.model
|
|
||||||
|
|
||||||
prompt = engine.tokenizer.apply_chat_template(
|
|
||||||
[{"role": m.role, "content": m.content} for m in request.messages],
|
|
||||||
tokenize=False,
|
|
||||||
)
|
|
||||||
prompt_tokens = len(engine.tokenizer.encode(prompt))
|
|
||||||
|
|
||||||
if request.stream:
|
|
||||||
agen = engine.generate_async(
|
|
||||||
prompt=prompt,
|
|
||||||
max_tokens=request.max_tokens,
|
|
||||||
temperature=request.temperature,
|
|
||||||
top_p=request.top_p,
|
|
||||||
top_k=request.top_k,
|
|
||||||
)
|
|
||||||
|
|
||||||
async def event_stream():
|
|
||||||
yield _make_chunk(
|
|
||||||
{"role": "assistant"},
|
|
||||||
finish_reason=None,
|
|
||||||
resp_id=resp_id,
|
|
||||||
created=created,
|
|
||||||
model=model,
|
|
||||||
)
|
|
||||||
|
|
||||||
completion_tokens = 0
|
|
||||||
async for token in agen:
|
|
||||||
yield _make_chunk(
|
|
||||||
{"content": token},
|
|
||||||
finish_reason=None,
|
|
||||||
resp_id=resp_id,
|
|
||||||
created=created,
|
|
||||||
model=model,
|
|
||||||
)
|
|
||||||
completion_tokens += 1
|
|
||||||
|
|
||||||
yield _make_chunk(
|
|
||||||
{},
|
|
||||||
finish_reason="stop",
|
|
||||||
resp_id=resp_id,
|
|
||||||
created=created,
|
|
||||||
model=model,
|
|
||||||
)
|
|
||||||
|
|
||||||
usage = {
|
|
||||||
"prompt_tokens": prompt_tokens,
|
|
||||||
"completion_tokens": completion_tokens,
|
|
||||||
"total_tokens": prompt_tokens + completion_tokens,
|
|
||||||
}
|
|
||||||
yield f"data: {json.dumps(usage, ensure_ascii=False)}\n\n"
|
|
||||||
yield "data: [DONE]\n\n"
|
|
||||||
|
|
||||||
return StreamingResponse(
|
|
||||||
event_stream(),
|
|
||||||
media_type="text/event-stream",
|
|
||||||
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
|
||||||
)
|
|
||||||
|
|
||||||
completion_tokens = 0
|
|
||||||
chunks: List[str] = []
|
|
||||||
agen = engine.generate_async(
|
|
||||||
prompt=prompt,
|
|
||||||
max_tokens=request.max_tokens,
|
|
||||||
temperature=request.temperature,
|
|
||||||
top_p=request.top_p,
|
|
||||||
top_k=request.top_k,
|
|
||||||
)
|
|
||||||
async for token in agen:
|
|
||||||
chunks.append(token)
|
|
||||||
completion_tokens += 1
|
|
||||||
content = "".join(chunks)
|
|
||||||
|
|
||||||
return {
|
|
||||||
"id": resp_id,
|
|
||||||
"object": "chat.completion",
|
|
||||||
"created": created,
|
|
||||||
"model": model,
|
|
||||||
"choices": [
|
|
||||||
{
|
|
||||||
"index": 0,
|
|
||||||
"message": {"role": "assistant", "content": content},
|
|
||||||
"finish_reason": "stop",
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"usage": {
|
|
||||||
"prompt_tokens": prompt_tokens,
|
|
||||||
"completion_tokens": completion_tokens,
|
|
||||||
"total_tokens": prompt_tokens + completion_tokens,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def _make_anthropic_sse(event: str, data: Dict[str, Any]) -> str:
|
|
||||||
return f"event: {event}\ndata: {json.dumps(data, ensure_ascii=False)}\n\n"
|
|
||||||
|
|
||||||
|
|
||||||
def _check_stop_sequence(text: str, stop_sequences: List[str]) -> Optional[str]:
|
|
||||||
for seq in stop_sequences:
|
|
||||||
if seq and seq in text:
|
|
||||||
return seq
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
def _extract_text_content(content: Union[str, List[Dict[str, Any]]]) -> str:
|
|
||||||
if isinstance(content, str):
|
|
||||||
return content
|
|
||||||
if isinstance(content, list):
|
|
||||||
for block in content:
|
|
||||||
if isinstance(block, dict) and block.get("type") == "text":
|
|
||||||
return block.get("text", "")
|
|
||||||
return ""
|
|
||||||
|
|
||||||
|
|
||||||
def _build_anthropic_messages(
|
|
||||||
messages: List[AnthropicMessage], system: Optional[str]
|
|
||||||
) -> List[Dict[str, str]]:
|
|
||||||
result: List[Dict[str, str]] = []
|
|
||||||
if system:
|
|
||||||
result.append({"role": "system", "content": system})
|
|
||||||
for m in messages:
|
|
||||||
content = _extract_text_content(m.content)
|
|
||||||
if content:
|
|
||||||
result.append({"role": m.role, "content": content})
|
|
||||||
return result
|
|
||||||
|
|
||||||
|
|
||||||
@app.post("/v1/messages")
|
|
||||||
async def create_message(request: MessagesRequest):
|
|
||||||
"""Anthropic-compatible Messages API endpoint (streaming + non-streaming)."""
|
|
||||||
engine = _get_engine()
|
|
||||||
resp_id = f"msg_{uuid.uuid4().hex[:24]}"
|
|
||||||
model = request.model
|
|
||||||
|
|
||||||
chat_messages = _build_anthropic_messages(request.messages, request.system)
|
|
||||||
prompt = engine.tokenizer.apply_chat_template(chat_messages, tokenize=False)
|
|
||||||
prompt_tokens = len(engine.tokenizer.encode(prompt))
|
|
||||||
|
|
||||||
stop_sequences = request.stop_sequences or []
|
|
||||||
|
|
||||||
if request.stream:
|
|
||||||
agen = engine.generate_async(
|
|
||||||
prompt=prompt,
|
|
||||||
max_tokens=request.max_tokens,
|
|
||||||
temperature=request.temperature,
|
|
||||||
top_p=request.top_p,
|
|
||||||
top_k=request.top_k,
|
|
||||||
)
|
|
||||||
|
|
||||||
async def event_stream():
|
|
||||||
yield _make_anthropic_sse(
|
|
||||||
"message_start",
|
|
||||||
{
|
|
||||||
"type": "message_start",
|
|
||||||
"message": {
|
|
||||||
"id": resp_id,
|
|
||||||
"type": "message",
|
|
||||||
"role": "assistant",
|
|
||||||
"model": model,
|
|
||||||
"content": [],
|
|
||||||
"usage": {"input_tokens": prompt_tokens},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
yield _make_anthropic_sse(
|
|
||||||
"content_block_start",
|
|
||||||
{
|
|
||||||
"type": "content_block_start",
|
|
||||||
"index": 0,
|
|
||||||
"content_block": {"type": "text", "text": ""},
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
completion_tokens = 0
|
|
||||||
accumulated = ""
|
|
||||||
stopped_seq: Optional[str] = None
|
|
||||||
async for token in agen:
|
|
||||||
accumulated += token
|
|
||||||
completion_tokens += 1
|
|
||||||
|
|
||||||
matched = _check_stop_sequence(accumulated, stop_sequences)
|
|
||||||
if matched:
|
|
||||||
text = accumulated[: accumulated.rfind(matched)]
|
|
||||||
stopped_seq = matched
|
|
||||||
if text:
|
|
||||||
yield _make_anthropic_sse(
|
|
||||||
"content_block_delta",
|
|
||||||
{
|
|
||||||
"type": "content_block_delta",
|
|
||||||
"index": 0,
|
|
||||||
"delta": {"type": "text_delta", "text": text},
|
|
||||||
},
|
|
||||||
)
|
|
||||||
break
|
|
||||||
|
|
||||||
yield _make_anthropic_sse(
|
|
||||||
"content_block_delta",
|
|
||||||
{
|
|
||||||
"type": "content_block_delta",
|
|
||||||
"index": 0,
|
|
||||||
"delta": {"type": "text_delta", "text": token},
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
yield _make_anthropic_sse(
|
|
||||||
"content_block_stop",
|
|
||||||
{"type": "content_block_stop", "index": 0},
|
|
||||||
)
|
|
||||||
|
|
||||||
stop_reason = "stop_sequence" if stopped_seq else "end_turn"
|
|
||||||
yield _make_anthropic_sse(
|
|
||||||
"message_delta",
|
|
||||||
{
|
|
||||||
"type": "message_delta",
|
|
||||||
"delta": {"stop_reason": stop_reason, "stop_sequence": stopped_seq},
|
|
||||||
"usage": {"output_tokens": completion_tokens},
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
yield _make_anthropic_sse(
|
|
||||||
"message_stop",
|
|
||||||
{"type": "message_stop"},
|
|
||||||
)
|
|
||||||
|
|
||||||
return StreamingResponse(
|
|
||||||
event_stream(),
|
|
||||||
media_type="text/event-stream",
|
|
||||||
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
|
||||||
)
|
|
||||||
|
|
||||||
completion_tokens = 0
|
|
||||||
chunks: List[str] = []
|
|
||||||
agen = engine.generate_async(
|
|
||||||
prompt=prompt,
|
|
||||||
max_tokens=request.max_tokens,
|
|
||||||
temperature=request.temperature,
|
|
||||||
top_p=request.top_p,
|
|
||||||
top_k=request.top_k,
|
|
||||||
)
|
|
||||||
stopped_seq: Optional[str] = None
|
|
||||||
accumulated = ""
|
|
||||||
async for token in agen:
|
|
||||||
chunks.append(token)
|
|
||||||
completion_tokens += 1
|
|
||||||
accumulated += token
|
|
||||||
matched = _check_stop_sequence(accumulated, stop_sequences)
|
|
||||||
if matched:
|
|
||||||
stopped_seq = matched
|
|
||||||
break
|
|
||||||
|
|
||||||
content = "".join(chunks)
|
|
||||||
if stopped_seq:
|
|
||||||
idx = content.rfind(stopped_seq)
|
|
||||||
if idx != -1:
|
|
||||||
content = content[:idx]
|
|
||||||
|
|
||||||
return {
|
|
||||||
"id": resp_id,
|
|
||||||
"type": "message",
|
|
||||||
"role": "assistant",
|
|
||||||
"model": model,
|
|
||||||
"content": [{"type": "text", "text": content}],
|
|
||||||
"stop_reason": "stop_sequence" if stopped_seq else "end_turn",
|
|
||||||
"stop_sequence": stopped_seq,
|
|
||||||
"usage": {
|
|
||||||
"input_tokens": prompt_tokens,
|
|
||||||
"output_tokens": completion_tokens,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def run_server(
|
|
||||||
host: str = "0.0.0.0",
|
|
||||||
port: int = 8000,
|
|
||||||
reload: bool = False,
|
|
||||||
device: str = "cuda",
|
|
||||||
dtype: torch.dtype = torch.bfloat16,
|
|
||||||
param_path: Optional[Path] = None,
|
|
||||||
max_batch_size: int = 16,
|
|
||||||
):
|
|
||||||
configure_server(
|
|
||||||
device=device,
|
|
||||||
dtype=dtype,
|
|
||||||
param_path=param_path,
|
|
||||||
max_batch_size=max_batch_size,
|
|
||||||
)
|
|
||||||
uvicorn.run(
|
|
||||||
"astrai.inference.server:app",
|
|
||||||
host=host,
|
|
||||||
port=port,
|
|
||||||
reload=reload,
|
|
||||||
)
|
|
||||||
@@ -1,12 +1,11 @@
|
|||||||
from astrai.model.automodel import AutoModel
|
from astrai.model.automodel import AutoModel
|
||||||
from astrai.model.module import (
|
from astrai.model.components.attention import GQA
|
||||||
GQA,
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
MLP,
|
from astrai.model.components.linear import Linear
|
||||||
DecoderBlock,
|
from astrai.model.components.mlp import MLP
|
||||||
Linear,
|
from astrai.model.components.norm import RMSNorm
|
||||||
RMSNorm,
|
from astrai.model.encoder import EmbeddingEncoder
|
||||||
)
|
from astrai.model.transformer import AutoRegressiveLM
|
||||||
from astrai.model.transformer import Transformer
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
# Modules
|
# Modules
|
||||||
@@ -16,6 +15,7 @@ __all__ = [
|
|||||||
"GQA",
|
"GQA",
|
||||||
"DecoderBlock",
|
"DecoderBlock",
|
||||||
# Models
|
# Models
|
||||||
"Transformer",
|
"AutoRegressiveLM",
|
||||||
|
"EmbeddingEncoder",
|
||||||
"AutoModel",
|
"AutoModel",
|
||||||
]
|
]
|
||||||
|
|||||||
+15
-42
@@ -2,15 +2,16 @@
|
|||||||
AutoModel base class for model loading and saving.
|
AutoModel base class for model loading and saving.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Self, Type, Union
|
from typing import Self, Union
|
||||||
|
|
||||||
import safetensors.torch as st
|
import safetensors.torch as st
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
from astrai.config import ModelConfig
|
from astrai.config.model_config import BaseModelConfig, ConfigFactory
|
||||||
from astrai.factory import Registry
|
from astrai.factory import BaseFactory
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
@@ -39,65 +40,37 @@ def _disable_random_init(enable: bool = True):
|
|||||||
setattr(nn.init, name, orig_func)
|
setattr(nn.init, name, orig_func)
|
||||||
|
|
||||||
|
|
||||||
class AutoModel(nn.Module):
|
class AutoModel(BaseFactory["AutoModel"], nn.Module):
|
||||||
"""
|
"""
|
||||||
Autoregressive language model base class.
|
Autoregressive language model base class.
|
||||||
Provides model loading/saving and generation capabilities.
|
Provides model loading/saving, registration, and generation.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
_registry = Registry()
|
def __init__(self, config: BaseModelConfig):
|
||||||
|
|
||||||
def __init__(self, config: ModelConfig):
|
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.config = config
|
self.config = config
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def register(cls, model_type: str):
|
|
||||||
"""
|
|
||||||
Class method decorator to register model type.
|
|
||||||
|
|
||||||
Usage:
|
|
||||||
@AutoModel.register('transformer')
|
|
||||||
class Transformer(AutoModel):
|
|
||||||
...
|
|
||||||
"""
|
|
||||||
|
|
||||||
def decorator(sub_cls: Type["AutoModel"]) -> Type["AutoModel"]:
|
|
||||||
cls._registry.register(model_type.lower(), sub_cls)
|
|
||||||
return sub_cls
|
|
||||||
|
|
||||||
return decorator
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def get_model_class(cls, model_type: str) -> Type["AutoModel"]:
|
|
||||||
"""Get model class by model_type string."""
|
|
||||||
model_type = model_type.lower()
|
|
||||||
if not cls._registry.contains(model_type):
|
|
||||||
available = cls._registry.list_names()
|
|
||||||
raise ValueError(
|
|
||||||
f"Unknown model_type: {model_type}. Available: {available}"
|
|
||||||
)
|
|
||||||
return cls._registry.get(model_type)
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_pretrained(
|
def from_pretrained(
|
||||||
cls,
|
cls,
|
||||||
path: Union[str, Path],
|
path: Union[str, Path],
|
||||||
disable_random_init: bool = True,
|
disable_random_init: bool = True,
|
||||||
|
strict: bool = True,
|
||||||
) -> nn.Module:
|
) -> nn.Module:
|
||||||
|
|
||||||
model_path = Path(path)
|
model_path = Path(path)
|
||||||
|
|
||||||
# Load config
|
# Load config
|
||||||
config = ModelConfig()
|
|
||||||
config_path = model_path / "config.json"
|
config_path = model_path / "config.json"
|
||||||
if config_path.exists():
|
if config_path.exists():
|
||||||
config.load(str(config_path))
|
with open(config_path, "r") as f:
|
||||||
|
raw = json.load(f)
|
||||||
|
config = ConfigFactory.load(raw)
|
||||||
|
model_type = config.model_type or "autoregressive_lm"
|
||||||
else:
|
else:
|
||||||
raise FileNotFoundError(f"Config file not found: {config_path}")
|
raise FileNotFoundError(f"Config file not found: {config_path}")
|
||||||
|
|
||||||
model_type = config.model_type or "transformer"
|
actual_cls = AutoModel.get_component_class(model_type)
|
||||||
actual_cls = cls.get_model_class(model_type)
|
|
||||||
|
|
||||||
with _disable_random_init(enable=disable_random_init):
|
with _disable_random_init(enable=disable_random_init):
|
||||||
model = actual_cls(config)
|
model = actual_cls(config)
|
||||||
@@ -106,7 +79,7 @@ class AutoModel(nn.Module):
|
|||||||
weights_path = model_path / "model.safetensors"
|
weights_path = model_path / "model.safetensors"
|
||||||
if weights_path.exists():
|
if weights_path.exists():
|
||||||
state_dict = st.load_file(str(weights_path))
|
state_dict = st.load_file(str(weights_path))
|
||||||
model.load_state_dict(state_dict, strict=False)
|
model.load_state_dict(state_dict, strict=strict)
|
||||||
|
|
||||||
return model
|
return model
|
||||||
|
|
||||||
@@ -118,7 +91,7 @@ class AutoModel(nn.Module):
|
|||||||
save_path.mkdir(parents=True, exist_ok=True)
|
save_path.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
# Save config
|
# Save config
|
||||||
self.config.save(str(save_path / "config.json"))
|
self.config.to_file(str(save_path / "config.json"))
|
||||||
|
|
||||||
# Save weights
|
# Save weights
|
||||||
st.save_file(self.state_dict(), str(save_path / "model.safetensors"))
|
st.save_file(self.state_dict(), str(save_path / "model.safetensors"))
|
||||||
|
|||||||
@@ -0,0 +1,25 @@
|
|||||||
|
from astrai.model.components.attention import GQA, MLA, repeat_kv
|
||||||
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
|
from astrai.model.components.embedding import Embedding
|
||||||
|
from astrai.model.components.linear import Linear
|
||||||
|
from astrai.model.components.mlp import MLP
|
||||||
|
from astrai.model.components.norm import RMSNorm
|
||||||
|
from astrai.model.components.rope import (
|
||||||
|
RotaryEmbedding,
|
||||||
|
apply_rotary_emb,
|
||||||
|
get_rotary_emb,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"Linear",
|
||||||
|
"RMSNorm",
|
||||||
|
"MLP",
|
||||||
|
"Embedding",
|
||||||
|
"GQA",
|
||||||
|
"MLA",
|
||||||
|
"DecoderBlock",
|
||||||
|
"RotaryEmbedding",
|
||||||
|
"apply_rotary_emb",
|
||||||
|
"get_rotary_emb",
|
||||||
|
"repeat_kv",
|
||||||
|
]
|
||||||
@@ -0,0 +1,212 @@
|
|||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.inference.core.cache import KvcacheView
|
||||||
|
from astrai.model.components.linear import Linear
|
||||||
|
from astrai.model.components.norm import RMSNorm
|
||||||
|
from astrai.model.components.rope import apply_rotary_emb
|
||||||
|
|
||||||
|
|
||||||
|
def repeat_kv(x: Tensor, n_rep: int) -> 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)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class AttnFactory(BaseFactory[nn.Module]):
|
||||||
|
@classmethod
|
||||||
|
def create(cls, attn_type: str, **kwargs) -> nn.Module:
|
||||||
|
return super().create(attn_type, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
@AttnFactory.register("gqa")
|
||||||
|
class GQA(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim: int,
|
||||||
|
n_heads: int,
|
||||||
|
n_kv_heads: int,
|
||||||
|
use_qk_norm: bool,
|
||||||
|
norm_eps: float,
|
||||||
|
use_gated_attention: bool,
|
||||||
|
layer_id: int,
|
||||||
|
):
|
||||||
|
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
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
@AttnFactory.register("mla")
|
||||||
|
class MLA(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim: int,
|
||||||
|
n_heads: int,
|
||||||
|
n_kv_heads: int,
|
||||||
|
kv_lora_rank: int,
|
||||||
|
qk_nope_head_dim: int,
|
||||||
|
qk_rope_head_dim: int,
|
||||||
|
norm_eps: float,
|
||||||
|
use_qk_norm: bool,
|
||||||
|
use_gated_attention: bool,
|
||||||
|
layer_id: int,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
self.n_heads = n_heads
|
||||||
|
self.n_kv_heads = n_kv_heads
|
||||||
|
self.kv_lora_rank = kv_lora_rank
|
||||||
|
self.qk_nope_head_dim = qk_nope_head_dim
|
||||||
|
self.qk_rope_head_dim = qk_rope_head_dim
|
||||||
|
self.head_dim = qk_nope_head_dim + qk_rope_head_dim
|
||||||
|
self.layer_id = layer_id
|
||||||
|
self.n_rep = n_heads // n_kv_heads
|
||||||
|
self.use_qk_norm = use_qk_norm
|
||||||
|
self.use_gated_attention = use_gated_attention
|
||||||
|
|
||||||
|
self.q_proj = Linear(dim, n_heads * self.head_dim, bias=False)
|
||||||
|
|
||||||
|
if self.use_qk_norm:
|
||||||
|
self.q_norm = RMSNorm(self.head_dim, norm_eps)
|
||||||
|
self.k_norm = RMSNorm(self.head_dim, norm_eps)
|
||||||
|
self.kv_a_proj = Linear(dim, kv_lora_rank, bias=False)
|
||||||
|
self.kv_norm = RMSNorm(kv_lora_rank, norm_eps)
|
||||||
|
|
||||||
|
self.kv_b_proj = Linear(
|
||||||
|
kv_lora_rank,
|
||||||
|
n_kv_heads * (2 * self.head_dim),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.o_proj = Linear(dim, dim, bias=False)
|
||||||
|
|
||||||
|
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_nope_head_dim :],
|
||||||
|
)
|
||||||
|
q_rope = apply_rotary_emb(q_rope, rotary_emb)
|
||||||
|
k_rope = apply_rotary_emb(k_rope, rotary_emb)
|
||||||
|
|
||||||
|
q = torch.cat([q_nope, q_rope], dim=-1)
|
||||||
|
k = torch.cat([k_nope, k_rope], dim=-1)
|
||||||
|
|
||||||
|
if self.use_qk_norm:
|
||||||
|
q = self.q_norm(q)
|
||||||
|
k = self.k_norm(k)
|
||||||
|
|
||||||
|
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
|
||||||
@@ -0,0 +1,59 @@
|
|||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch.nn as nn
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.inference.core.cache import KvcacheView
|
||||||
|
from astrai.model.components.attention import AttnFactory
|
||||||
|
from astrai.model.components.mlp import FFNFactory
|
||||||
|
from astrai.model.components.norm import RMSNorm
|
||||||
|
|
||||||
|
|
||||||
|
class DecoderBlock(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim: int,
|
||||||
|
n_heads: int,
|
||||||
|
dim_ffn: int,
|
||||||
|
n_kv_heads: int,
|
||||||
|
norm_eps: float,
|
||||||
|
use_qk_norm: bool,
|
||||||
|
use_gated_attention: bool,
|
||||||
|
layer_id: int,
|
||||||
|
attn_type: str = "gqa",
|
||||||
|
ffn_type: str = "mlp",
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.attention = AttnFactory.create(
|
||||||
|
attn_type,
|
||||||
|
dim=dim,
|
||||||
|
n_heads=n_heads,
|
||||||
|
n_kv_heads=n_kv_heads,
|
||||||
|
use_qk_norm=use_qk_norm,
|
||||||
|
norm_eps=norm_eps,
|
||||||
|
use_gated_attention=use_gated_attention,
|
||||||
|
layer_id=layer_id,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
self.input_norm = RMSNorm(dim, norm_eps)
|
||||||
|
self.post_attention_norm = RMSNorm(dim, norm_eps)
|
||||||
|
self.mlp = FFNFactory.create(ffn_type, dim, dim_ffn, **kwargs)
|
||||||
|
|
||||||
|
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
|
||||||
@@ -0,0 +1,16 @@
|
|||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
class Embedding(nn.Module):
|
||||||
|
def __init__(self, vocab_size: int, embedding_dim: int):
|
||||||
|
super().__init__()
|
||||||
|
self.weight = nn.Parameter(torch.empty((vocab_size, embedding_dim)))
|
||||||
|
|
||||||
|
def reset_parameters(self):
|
||||||
|
nn.init.normal_(self.weight, mean=0.0, std=0.02)
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
return F.embedding(x, self.weight)
|
||||||
@@ -0,0 +1,21 @@
|
|||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
class Linear(nn.Module):
|
||||||
|
def __init__(self, in_dim: int, out_dim: int, bias: bool = False):
|
||||||
|
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 reset_parameters(self):
|
||||||
|
nn.init.kaiming_uniform_(self.weight, a=5**0.5)
|
||||||
|
if self.bias is not None:
|
||||||
|
fan_in, _ = nn.init._calculate_fan_in_and_fan_out(self.weight)
|
||||||
|
bound = 1 / (fan_in**0.5)
|
||||||
|
nn.init.uniform_(self.bias, -bound, bound)
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
return F.linear(x, self.weight, self.bias)
|
||||||
@@ -0,0 +1,93 @@
|
|||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.model.components.linear import Linear
|
||||||
|
|
||||||
|
|
||||||
|
class FFNFactory(BaseFactory[nn.Module]):
|
||||||
|
@classmethod
|
||||||
|
def create(cls, ffn_type: str, dim: int, dim_ffn: int, **kwargs) -> nn.Module:
|
||||||
|
return super().create(ffn_type, dim, dim_ffn, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
@FFNFactory.register("mlp")
|
||||||
|
class MLP(nn.Module):
|
||||||
|
def __init__(self, dim: int, dim_ffn: int):
|
||||||
|
super().__init__()
|
||||||
|
self.up = Linear(dim, dim_ffn)
|
||||||
|
self.gate = Linear(dim, dim_ffn)
|
||||||
|
self.down = Linear(dim_ffn, dim)
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
gated = self.up(x) * F.silu(self.gate(x))
|
||||||
|
out = self.down(gated)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
@FFNFactory.register("moe")
|
||||||
|
class DeepSeekMoE(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim: int,
|
||||||
|
dim_ffn: int,
|
||||||
|
n_routed_experts: int,
|
||||||
|
n_shared_experts: int = 1,
|
||||||
|
n_activated_experts: int = 2,
|
||||||
|
topk_method: str = "greedy",
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
self.n_routed_experts = n_routed_experts
|
||||||
|
self.n_shared_experts = n_shared_experts
|
||||||
|
self.n_activated_experts = n_activated_experts
|
||||||
|
self.topk_method = topk_method
|
||||||
|
|
||||||
|
self.router = Linear(dim, n_routed_experts, bias=False)
|
||||||
|
|
||||||
|
self.shared_experts = nn.ModuleList(
|
||||||
|
[MLP(dim, dim_ffn) for _ in range(n_shared_experts)]
|
||||||
|
)
|
||||||
|
self.routed_experts = nn.ModuleList(
|
||||||
|
[MLP(dim, dim_ffn) for _ in range(n_routed_experts)]
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
bsz, seq_len, dim = x.shape
|
||||||
|
x_flat = x.view(-1, dim)
|
||||||
|
|
||||||
|
shared_out = self._shared_forward(x_flat)
|
||||||
|
routed_out = self._routed_forward(x_flat)
|
||||||
|
|
||||||
|
out = (shared_out + routed_out).view(bsz, seq_len, dim)
|
||||||
|
return out
|
||||||
|
|
||||||
|
def _shared_forward(self, x: Tensor) -> Tensor:
|
||||||
|
if self.n_shared_experts == 0:
|
||||||
|
return torch.zeros_like(x)
|
||||||
|
return sum(e(x) for e in self.shared_experts) / self.n_shared_experts
|
||||||
|
|
||||||
|
def _routed_forward(self, x: Tensor) -> Tensor:
|
||||||
|
N, D = x.shape
|
||||||
|
K = self.n_activated_experts
|
||||||
|
|
||||||
|
router_logits = self.router(x)
|
||||||
|
router_probs = torch.softmax(router_logits.float(), dim=-1).to(x.dtype)
|
||||||
|
|
||||||
|
topk_weights, topk_indices = torch.topk(router_probs, K, dim=-1)
|
||||||
|
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
||||||
|
|
||||||
|
output = torch.zeros(N, D, device=x.device, dtype=x.dtype)
|
||||||
|
for expert_idx in range(self.n_routed_experts):
|
||||||
|
expert_mask = topk_indices == expert_idx
|
||||||
|
token_idx, k_idx = expert_mask.nonzero(as_tuple=True)
|
||||||
|
if token_idx.numel() == 0:
|
||||||
|
continue
|
||||||
|
expert_input = x[token_idx]
|
||||||
|
expert_output = self.routed_experts[expert_idx](expert_input)
|
||||||
|
weights = topk_weights[token_idx, k_idx].unsqueeze(-1)
|
||||||
|
output.index_add_(0, token_idx, expert_output * weights)
|
||||||
|
|
||||||
|
return output
|
||||||
@@ -0,0 +1,15 @@
|
|||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
class RMSNorm(nn.Module):
|
||||||
|
def __init__(self, dim, norm_eps):
|
||||||
|
super().__init__()
|
||||||
|
self.weight = nn.Parameter(torch.ones(dim))
|
||||||
|
self.normalized_shape = (dim,)
|
||||||
|
self.norm_eps = norm_eps
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
return F.rms_norm(x, self.normalized_shape, self.weight, self.norm_eps)
|
||||||
@@ -0,0 +1,53 @@
|
|||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
def get_rotary_emb(
|
||||||
|
dim: int,
|
||||||
|
max_len: int,
|
||||||
|
base: float = 10000,
|
||||||
|
device: Optional[torch.device] = None,
|
||||||
|
) -> Tensor:
|
||||||
|
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: float = 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)
|
||||||
@@ -0,0 +1,100 @@
|
|||||||
|
from typing import Any, Mapping, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.config.model_config import EncoderConfig
|
||||||
|
from astrai.model.automodel import AutoModel
|
||||||
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
|
from astrai.model.components.embedding import Embedding
|
||||||
|
from astrai.model.components.norm import RMSNorm
|
||||||
|
from astrai.model.components.rope import RotaryEmbedding
|
||||||
|
from astrai.model.transformer import process_attention_mask
|
||||||
|
|
||||||
|
|
||||||
|
@AutoModel.register("embedding")
|
||||||
|
class EmbeddingEncoder(AutoModel):
|
||||||
|
def __init__(self, config: EncoderConfig):
|
||||||
|
super().__init__(config)
|
||||||
|
self.config = config
|
||||||
|
rope_dim = config.dim // config.n_heads
|
||||||
|
rope_base = config.rope_theta if config.rope_theta is not None else 10000
|
||||||
|
self.rotary_embedding = RotaryEmbedding(rope_dim, config.max_len, rope_base)
|
||||||
|
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.pooling_type = config.pooling_type or "mean"
|
||||||
|
self.normalize_embeddings = config.normalize_embeddings or False
|
||||||
|
|
||||||
|
self.apply(self._init_weights)
|
||||||
|
|
||||||
|
def _init_weights(self, module):
|
||||||
|
if hasattr(module, "reset_parameters"):
|
||||||
|
module.reset_parameters()
|
||||||
|
|
||||||
|
def load_state_dict(self, state_dict: Mapping[str, Any], strict=True, assign=False):
|
||||||
|
state_dict = dict(state_dict)
|
||||||
|
state_dict.pop("lm_head.weight", None)
|
||||||
|
return super().load_state_dict(state_dict, strict=strict, assign=assign)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
input_ids: Tensor,
|
||||||
|
input_mask: Optional[Tensor] = None,
|
||||||
|
position_ids: Optional[Tensor] = None,
|
||||||
|
) -> Tensor:
|
||||||
|
assert input_ids.ndim == 2
|
||||||
|
B, S = input_ids.shape
|
||||||
|
|
||||||
|
x = self.embed_tokens(input_ids)
|
||||||
|
|
||||||
|
if position_ids is None:
|
||||||
|
position_ids = torch.arange(S, device=x.device).unsqueeze(0).expand(B, -1)
|
||||||
|
|
||||||
|
rotary_emb = self.rotary_embedding(x, position_ids)
|
||||||
|
attn_mask = process_attention_mask(x, position_ids, input_mask, is_causal=False)
|
||||||
|
|
||||||
|
for layer in self.layers:
|
||||||
|
x = layer(x, rotary_emb, attn_mask, paged_cache=None)
|
||||||
|
|
||||||
|
hidden_states = self.norm(x)
|
||||||
|
|
||||||
|
if self.pooling_type == "cls":
|
||||||
|
pooled = hidden_states[:, 0]
|
||||||
|
elif self.pooling_type == "last":
|
||||||
|
if input_mask is not None:
|
||||||
|
lengths = input_mask.sum(dim=1) - 1
|
||||||
|
pooled = hidden_states[torch.arange(B, device=x.device), lengths]
|
||||||
|
else:
|
||||||
|
pooled = hidden_states[:, -1]
|
||||||
|
else:
|
||||||
|
if input_mask is not None:
|
||||||
|
mask = input_mask.unsqueeze(-1).to(dtype=hidden_states.dtype)
|
||||||
|
pooled = (hidden_states * mask).sum(dim=1) / mask.sum(dim=1).clamp(
|
||||||
|
min=1.0
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
pooled = hidden_states.mean(dim=1)
|
||||||
|
|
||||||
|
if self.normalize_embeddings:
|
||||||
|
pooled = torch.nn.functional.normalize(pooled, p=2, dim=-1)
|
||||||
|
|
||||||
|
return pooled
|
||||||
@@ -1,337 +0,0 @@
|
|||||||
from typing import Optional, Tuple
|
|
||||||
|
|
||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
import torch.nn.functional as F
|
|
||||||
from torch import Tensor
|
|
||||||
|
|
||||||
from astrai.inference.cache import CacheView
|
|
||||||
|
|
||||||
|
|
||||||
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
|
||||||
"""Repeat KV heads n_rep times for GQA."""
|
|
||||||
bs, slen, n_heads, head_dim = x.shape
|
|
||||||
if n_rep == 1:
|
|
||||||
return x
|
|
||||||
return (
|
|
||||||
x[:, :, :, None, :]
|
|
||||||
.expand(bs, slen, n_heads, n_rep, head_dim)
|
|
||||||
.reshape(bs, slen, n_heads * n_rep, head_dim)
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def get_rotary_emb(
|
|
||||||
dim: int,
|
|
||||||
max_len: int,
|
|
||||||
base: float = 10000,
|
|
||||||
device: Optional[torch.device] = None,
|
|
||||||
) -> Tuple[Tensor, Tensor]:
|
|
||||||
"""Precompute cos/sin for RoPE."""
|
|
||||||
theta = base ** (-torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim)
|
|
||||||
t = torch.arange(0, max_len, dtype=torch.float64, device=device)
|
|
||||||
freqs = torch.outer(t, theta)
|
|
||||||
return torch.cos(freqs).float(), torch.sin(freqs).float()
|
|
||||||
|
|
||||||
|
|
||||||
def apply_rotary_emb(x: torch.Tensor, rotary_emb: Tuple[Tensor, Tensor]) -> Tensor:
|
|
||||||
"""Apply rotary embedding via cos/sin (shape-preserving)."""
|
|
||||||
dtype = x.dtype
|
|
||||||
cos, sin = rotary_emb
|
|
||||||
cos = cos.unsqueeze(0).unsqueeze(2)
|
|
||||||
sin = sin.unsqueeze(0).unsqueeze(2)
|
|
||||||
x_real = x[..., 0::2]
|
|
||||||
x_imag = x[..., 1::2]
|
|
||||||
x_real_rot = x_real * cos - x_imag * sin
|
|
||||||
x_imag_rot = x_real * sin + x_imag * cos
|
|
||||||
x_out = torch.stack([x_real_rot, x_imag_rot], dim=-1)
|
|
||||||
x_out = x_out.view(*x_out.shape[:-2], -1)
|
|
||||||
return x_out.to(dtype)
|
|
||||||
|
|
||||||
|
|
||||||
class RotaryEmbedding(nn.Module):
|
|
||||||
def __init__(self, dim: int, max_len: int, base: int = 10000):
|
|
||||||
super().__init__()
|
|
||||||
self.dim = dim
|
|
||||||
self.max_len = max_len
|
|
||||||
self.base = base
|
|
||||||
self.max_len_cached = None
|
|
||||||
self._set_rotary_buffer(self.max_len, None)
|
|
||||||
|
|
||||||
def _set_rotary_buffer(self, max_len: int, device: Optional[torch.device] = None):
|
|
||||||
cos_cached, sin_cached = get_rotary_emb(self.dim, max_len, self.base, device)
|
|
||||||
self.register_buffer("cos_cached", cos_cached, persistent=False)
|
|
||||||
self.register_buffer("sin_cached", sin_cached, persistent=False)
|
|
||||||
self.max_len_cached = max_len
|
|
||||||
|
|
||||||
def forward(self, x: Tensor, start_pos: int = 0) -> Tuple[Tensor, Tensor]:
|
|
||||||
seq_len = x.size(1)
|
|
||||||
if self.max_len_cached < seq_len + start_pos:
|
|
||||||
self._set_rotary_buffer(self.max_len_cached * 2, x.device)
|
|
||||||
cos = self.cos_cached[start_pos : start_pos + seq_len]
|
|
||||||
sin = self.sin_cached[start_pos : start_pos + seq_len]
|
|
||||||
return (cos, sin)
|
|
||||||
|
|
||||||
|
|
||||||
class Linear(nn.Module):
|
|
||||||
def __init__(self, in_dim: int, out_dim: int, bias: bool = False):
|
|
||||||
super().__init__()
|
|
||||||
self.weight = nn.Parameter(torch.empty((out_dim, in_dim)))
|
|
||||||
self.bias = nn.Parameter(torch.zeros(out_dim)) if bias else None
|
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
|
||||||
return F.linear(x, self.weight, self.bias)
|
|
||||||
|
|
||||||
|
|
||||||
class RMSNorm(nn.Module):
|
|
||||||
def __init__(self, dim, norm_eps):
|
|
||||||
super().__init__()
|
|
||||||
self.weight = nn.Parameter(torch.ones(dim))
|
|
||||||
self.normalized_shape = (dim,)
|
|
||||||
self.norm_eps = norm_eps
|
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
|
||||||
return F.rms_norm(x, self.normalized_shape, self.weight, self.norm_eps)
|
|
||||||
|
|
||||||
|
|
||||||
class MLP(nn.Module):
|
|
||||||
def __init__(self, dim: int, dim_feed_forward: int):
|
|
||||||
super().__init__()
|
|
||||||
self.up = Linear(dim, dim_feed_forward)
|
|
||||||
self.gate = Linear(dim, dim_feed_forward)
|
|
||||||
self.down = Linear(dim_feed_forward, dim)
|
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
|
||||||
gated = self.up(x) * F.silu(self.gate(x))
|
|
||||||
out = self.down(gated)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
class GQA(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
dim: int,
|
|
||||||
n_heads: int,
|
|
||||||
n_kv_heads: int,
|
|
||||||
use_qk_norm: bool,
|
|
||||||
norm_eps: float,
|
|
||||||
use_gated_attention: bool,
|
|
||||||
layer_id: int,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
assert dim % n_heads == 0
|
|
||||||
assert n_heads % n_kv_heads == 0
|
|
||||||
|
|
||||||
self.head_dim = dim // n_heads
|
|
||||||
self.layer_id = layer_id
|
|
||||||
self.dim = dim
|
|
||||||
self.n_heads = n_heads
|
|
||||||
self.n_kv_heads = n_kv_heads
|
|
||||||
self.n_rep = n_heads // n_kv_heads
|
|
||||||
self.use_qk_norm = use_qk_norm
|
|
||||||
self.use_gated_attention = use_gated_attention
|
|
||||||
|
|
||||||
self.q_proj = Linear(dim, n_heads * self.head_dim)
|
|
||||||
self.k_proj = Linear(dim, n_kv_heads * self.head_dim)
|
|
||||||
self.v_proj = Linear(dim, n_kv_heads * self.head_dim)
|
|
||||||
self.o_proj = Linear(dim, dim)
|
|
||||||
|
|
||||||
if self.use_qk_norm:
|
|
||||||
self.q_norm = RMSNorm(self.head_dim, norm_eps)
|
|
||||||
self.k_norm = RMSNorm(self.head_dim, norm_eps)
|
|
||||||
|
|
||||||
if self.use_gated_attention:
|
|
||||||
self.gate = Linear(dim, dim)
|
|
||||||
|
|
||||||
def _split_heads(self, x: Tensor, n_heads) -> Tensor:
|
|
||||||
batch_size, seq_len, _ = x.shape
|
|
||||||
x = x.reshape(batch_size, seq_len, n_heads, self.head_dim)
|
|
||||||
return x
|
|
||||||
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
x: Tensor,
|
|
||||||
rotary_emb: Tuple[Tensor, Tensor],
|
|
||||||
mask: Tensor = None,
|
|
||||||
paged_cache: Optional[CacheView] = None,
|
|
||||||
start_pos: int = 0,
|
|
||||||
) -> Tensor:
|
|
||||||
bsz, seq_len, _ = x.size()
|
|
||||||
is_causal = mask is None
|
|
||||||
|
|
||||||
# (bsz, seq_len, dim) -> (bsz, seq_len, n_heads, head_dim)
|
|
||||||
q = self._split_heads(self.q_proj(x), self.n_heads)
|
|
||||||
k = self._split_heads(self.k_proj(x), self.n_kv_heads)
|
|
||||||
v = self._split_heads(self.v_proj(x), self.n_kv_heads)
|
|
||||||
q, k = apply_rotary_emb(q, rotary_emb), apply_rotary_emb(k, rotary_emb)
|
|
||||||
|
|
||||||
if self.use_qk_norm:
|
|
||||||
q, k = self.q_norm(q), self.k_norm(k)
|
|
||||||
|
|
||||||
if paged_cache is not None:
|
|
||||||
paged_cache.write(self.layer_id, start_pos, k, v)
|
|
||||||
k, v = paged_cache.gather(self.layer_id)
|
|
||||||
|
|
||||||
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)
|
|
||||||
|
|
||||||
# (bsz, seq_len, n_heads, head_dim) -> (bsz, n_heads, seq_len, head_dim)
|
|
||||||
q, k, v = q.permute(0, 2, 1, 3), k.permute(0, 2, 1, 3), v.permute(0, 2, 1, 3)
|
|
||||||
sdqa_out = (
|
|
||||||
F.scaled_dot_product_attention(q, k, v, mask, is_causal=is_causal)
|
|
||||||
.permute(0, 2, 1, 3)
|
|
||||||
.contiguous()
|
|
||||||
.flatten(2)
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.use_gated_attention:
|
|
||||||
sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
|
|
||||||
|
|
||||||
out = self.o_proj(sdqa_out)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
class MLA(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
dim: int,
|
|
||||||
n_heads: int,
|
|
||||||
n_kv_heads: int,
|
|
||||||
kv_lora_rank: int,
|
|
||||||
qk_nope_head_dim: int,
|
|
||||||
qk_rope_head_dim: int,
|
|
||||||
norm_eps: float,
|
|
||||||
use_gated_attention: bool,
|
|
||||||
layer_id: int,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.dim = dim
|
|
||||||
self.n_heads = n_heads
|
|
||||||
self.n_kv_heads = n_kv_heads
|
|
||||||
self.kv_lora_rank = kv_lora_rank
|
|
||||||
self.qk_nope_head_dim = qk_nope_head_dim
|
|
||||||
self.qk_rope_head_dim = qk_rope_head_dim
|
|
||||||
self.head_dim = qk_nope_head_dim + qk_rope_head_dim
|
|
||||||
self.layer_id = layer_id
|
|
||||||
self.n_rep = n_heads // n_kv_heads
|
|
||||||
self.use_gated_attention = use_gated_attention
|
|
||||||
|
|
||||||
self.q_proj = Linear(dim, n_heads * self.head_dim, bias=False)
|
|
||||||
self.kv_a_proj = Linear(dim, kv_lora_rank, bias=False)
|
|
||||||
self.kv_norm = RMSNorm(kv_lora_rank, norm_eps)
|
|
||||||
|
|
||||||
# fused KV: (k_nope, k_rope, v)
|
|
||||||
self.kv_b_proj = Linear(
|
|
||||||
kv_lora_rank,
|
|
||||||
n_kv_heads * (self.head_dim + qk_rope_head_dim + self.head_dim),
|
|
||||||
)
|
|
||||||
|
|
||||||
self.o_proj = Linear(dim, dim, bias=False)
|
|
||||||
|
|
||||||
if use_gated_attention:
|
|
||||||
self.gate = Linear(dim, dim, bias=False)
|
|
||||||
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
x: Tensor,
|
|
||||||
rotary_emb: Tuple[Tensor, Tensor],
|
|
||||||
mask: Tensor = None,
|
|
||||||
paged_cache: Optional[CacheView] = None,
|
|
||||||
start_pos: int = 0,
|
|
||||||
) -> Tensor:
|
|
||||||
bsz, seq_len, _ = x.size()
|
|
||||||
is_causal = mask is None
|
|
||||||
|
|
||||||
q = self.q_proj(x)
|
|
||||||
q = q.view(bsz, seq_len, self.n_heads, self.head_dim)
|
|
||||||
|
|
||||||
kv_compressed = self.kv_a_proj(x)
|
|
||||||
kv_compressed = self.kv_norm(kv_compressed)
|
|
||||||
|
|
||||||
kv = self.kv_b_proj(kv_compressed)
|
|
||||||
kv = kv.view(bsz, seq_len, self.n_kv_heads, -1)
|
|
||||||
|
|
||||||
k_nope, k_rope, v = torch.split(
|
|
||||||
kv, [self.qk_nope_head_dim, self.qk_rope_head_dim, self.head_dim], dim=-1
|
|
||||||
)
|
|
||||||
|
|
||||||
q_nope, q_rope = (
|
|
||||||
q[..., : self.qk_nope_head_dim],
|
|
||||||
q[..., self.qk_rope_head_dim :],
|
|
||||||
)
|
|
||||||
q_rope = apply_rotary_emb(q_rope, rotary_emb)
|
|
||||||
k_rope = apply_rotary_emb(k_rope, rotary_emb)
|
|
||||||
|
|
||||||
q = torch.cat([q_nope, q_rope], dim=-1)
|
|
||||||
k = torch.cat([k_nope, k_rope], dim=-1)
|
|
||||||
|
|
||||||
if paged_cache is not None:
|
|
||||||
paged_cache.write(self.layer_id, start_pos, k, v)
|
|
||||||
k, v = paged_cache.gather(self.layer_id)
|
|
||||||
|
|
||||||
q = q.permute(0, 2, 1, 3)
|
|
||||||
k = k.permute(0, 2, 1, 3)
|
|
||||||
v = v.permute(0, 2, 1, 3)
|
|
||||||
|
|
||||||
attn_out = F.scaled_dot_product_attention(q, k, v, mask, is_causal=is_causal)
|
|
||||||
attn_out = attn_out.permute(0, 2, 1, 3).contiguous().flatten(2)
|
|
||||||
|
|
||||||
if self.use_gated_attention:
|
|
||||||
attn_out = attn_out * F.sigmoid(self.gate(x))
|
|
||||||
|
|
||||||
out = self.o_proj(attn_out)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
class DecoderBlock(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
dim: int,
|
|
||||||
n_heads: int,
|
|
||||||
dim_ffn: int,
|
|
||||||
n_kv_heads: int,
|
|
||||||
norm_eps: int,
|
|
||||||
use_qk_norm: bool,
|
|
||||||
use_gated_attention: bool,
|
|
||||||
layer_id: int,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.attention = GQA(
|
|
||||||
dim,
|
|
||||||
n_heads,
|
|
||||||
n_kv_heads,
|
|
||||||
use_qk_norm,
|
|
||||||
norm_eps,
|
|
||||||
use_gated_attention,
|
|
||||||
layer_id,
|
|
||||||
)
|
|
||||||
self.input_norm = RMSNorm(dim, norm_eps)
|
|
||||||
self.mlp = MLP(dim, dim_ffn)
|
|
||||||
self.post_attention_norm = RMSNorm(dim, norm_eps)
|
|
||||||
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
x: Tensor,
|
|
||||||
rotary_emb: Tuple[Tensor, Tensor],
|
|
||||||
attention_mask: Optional[Tensor] = None,
|
|
||||||
paged_cache: Optional[CacheView] = None,
|
|
||||||
start_pos: int = 0,
|
|
||||||
) -> Tensor:
|
|
||||||
attn_output = self.attention(
|
|
||||||
self.input_norm(x),
|
|
||||||
rotary_emb,
|
|
||||||
attention_mask,
|
|
||||||
paged_cache,
|
|
||||||
start_pos,
|
|
||||||
)
|
|
||||||
x = attn_output + x
|
|
||||||
|
|
||||||
x = self.mlp(self.post_attention_norm(x)) + x
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
class Embedding(nn.Module):
|
|
||||||
def __init__(self, vocab_size: int, embedding_dim: int):
|
|
||||||
super().__init__()
|
|
||||||
self.weight = nn.Parameter(torch.empty((vocab_size, embedding_dim)))
|
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
|
||||||
return F.embedding(x, self.weight)
|
|
||||||
+58
-54
@@ -4,65 +4,62 @@ import torch
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.config.model_config import ModelConfig
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
from astrai.inference.cache import CacheView
|
from astrai.inference.core.cache import KvcacheView
|
||||||
from astrai.model.automodel import AutoModel
|
from astrai.model.automodel import AutoModel
|
||||||
from astrai.model.module import (
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
DecoderBlock,
|
from astrai.model.components.embedding import Embedding
|
||||||
Embedding,
|
from astrai.model.components.linear import Linear
|
||||||
Linear,
|
from astrai.model.components.norm import RMSNorm
|
||||||
RMSNorm,
|
from astrai.model.components.rope import RotaryEmbedding
|
||||||
RotaryEmbedding,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def process_attention_mask(
|
def process_attention_mask(
|
||||||
seq_mask: Tensor,
|
|
||||||
input_tensor: Tensor,
|
input_tensor: Tensor,
|
||||||
start_pos: int = 0,
|
position_ids: Optional[Tensor],
|
||||||
|
input_mask: Optional[Tensor] = None,
|
||||||
is_causal: bool = False,
|
is_causal: bool = False,
|
||||||
) -> Tensor:
|
) -> Optional[Tensor]:
|
||||||
"""Build 4D attention mask from 2D seq_mask, with optional causal masking."""
|
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
|
device = input_tensor.device
|
||||||
dtype = input_tensor.dtype
|
dtype = input_tensor.dtype
|
||||||
seq_len = input_tensor.size(1)
|
B, S = input_tensor.size()[:2]
|
||||||
|
T = position_ids.max().item() + 1
|
||||||
|
|
||||||
if seq_mask is None:
|
if input_mask is None:
|
||||||
if start_pos != 0:
|
if position_ids.min().item() == 0 and is_causal:
|
||||||
seq_mask = torch.ones((1, seq_len), dtype=torch.bool, device=device)
|
|
||||||
else:
|
|
||||||
return None
|
return None
|
||||||
|
pad = torch.ones(B, T, dtype=torch.bool, device=device)
|
||||||
|
else:
|
||||||
|
pad = input_mask[:, :T].to(device=device, dtype=torch.bool)
|
||||||
|
|
||||||
if seq_mask.dim() > 2:
|
attend = pad.view(B, 1, T).expand(B, S, T).clone()
|
||||||
return seq_mask
|
|
||||||
|
|
||||||
batch_size = seq_mask.size(0)
|
|
||||||
seq_mask = seq_mask[:, : start_pos + seq_len].to(device=device, dtype=torch.bool)
|
|
||||||
expanded_mask = seq_mask.unsqueeze(1).expand(
|
|
||||||
batch_size, seq_len, start_pos + seq_len
|
|
||||||
)
|
|
||||||
|
|
||||||
if is_causal:
|
if is_causal:
|
||||||
expanded_mask = torch.tril(expanded_mask, diagonal=start_pos)
|
attend &= position_ids.unsqueeze(-1) >= torch.arange(T, device=device)
|
||||||
|
|
||||||
attention_mask = torch.zeros_like(expanded_mask, dtype=dtype, device=device)
|
return torch.full(
|
||||||
attention_mask = attention_mask.masked_fill_(
|
(B, 1, S, T), -torch.finfo(dtype).max / 2, dtype=dtype, device=device
|
||||||
~expanded_mask, -torch.finfo(dtype).max / 2
|
).masked_fill_(attend.unsqueeze(1), 0.0)
|
||||||
).unsqueeze(1)
|
|
||||||
|
|
||||||
return attention_mask
|
|
||||||
|
|
||||||
|
|
||||||
@AutoModel.register("transformer")
|
@AutoModel.register("autoregressive_lm")
|
||||||
class Transformer(AutoModel):
|
class AutoRegressiveLM(AutoModel):
|
||||||
"""Transformer language model with paged KV cache."""
|
"""Autoregressive language model with paged KV cache."""
|
||||||
|
|
||||||
def __init__(self, config: ModelConfig):
|
def __init__(self, config: AutoRegressiveLMConfig):
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
self.config = config
|
self.config = config
|
||||||
self.rotary_embedding = RotaryEmbedding(
|
rope_dim = (
|
||||||
config.dim // config.n_heads, config.max_len
|
config.qk_rope_head_dim
|
||||||
|
if config.attn_type == "mla"
|
||||||
|
else config.dim // config.n_heads
|
||||||
)
|
)
|
||||||
|
rope_base = config.rope_theta if config.rope_theta is not None else 10000
|
||||||
|
self.rotary_embedding = RotaryEmbedding(rope_dim, config.max_len, rope_base)
|
||||||
self.embed_tokens = Embedding(config.vocab_size, config.dim)
|
self.embed_tokens = Embedding(config.vocab_size, config.dim)
|
||||||
|
|
||||||
self.layers = nn.ModuleList(
|
self.layers = nn.ModuleList(
|
||||||
@@ -76,6 +73,15 @@ class Transformer(AutoModel):
|
|||||||
config.use_qk_norm,
|
config.use_qk_norm,
|
||||||
config.use_gated_attention,
|
config.use_gated_attention,
|
||||||
layer_id,
|
layer_id,
|
||||||
|
attn_type=config.attn_type,
|
||||||
|
ffn_type=config.ffn_type,
|
||||||
|
n_routed_experts=config.n_routed_experts,
|
||||||
|
n_shared_experts=config.n_shared_experts,
|
||||||
|
n_activated_experts=config.n_activated_experts,
|
||||||
|
topk_method=config.topk_method,
|
||||||
|
kv_lora_rank=config.kv_lora_rank,
|
||||||
|
qk_nope_head_dim=config.qk_nope_head_dim,
|
||||||
|
qk_rope_head_dim=config.qk_rope_head_dim,
|
||||||
)
|
)
|
||||||
for layer_id in range(config.n_layers)
|
for layer_id in range(config.n_layers)
|
||||||
]
|
]
|
||||||
@@ -84,15 +90,14 @@ class Transformer(AutoModel):
|
|||||||
self.norm = RMSNorm(config.dim, config.norm_eps)
|
self.norm = RMSNorm(config.dim, config.norm_eps)
|
||||||
self.lm_head = Linear(config.dim, config.vocab_size)
|
self.lm_head = Linear(config.dim, config.vocab_size)
|
||||||
|
|
||||||
if self.config.tie_weight:
|
if self.config.tie_weight is True:
|
||||||
self.lm_head.weight = self.embed_tokens.weight
|
self.lm_head.weight = self.embed_tokens.weight
|
||||||
|
|
||||||
self._init_weights()
|
self.apply(self._init_weights)
|
||||||
|
|
||||||
def _init_weights(self):
|
def _init_weights(self, module):
|
||||||
for param in self.parameters():
|
if hasattr(module, "reset_parameters"):
|
||||||
if param.dim() > 1:
|
module.reset_parameters()
|
||||||
nn.init.normal_(param, mean=0.0, std=0.006)
|
|
||||||
|
|
||||||
def load_state_dict(self, state_dict: Mapping[str, Any], strict=True, assign=False):
|
def load_state_dict(self, state_dict: Mapping[str, Any], strict=True, assign=False):
|
||||||
lm_head_key = "lm_head.weight"
|
lm_head_key = "lm_head.weight"
|
||||||
@@ -100,7 +105,7 @@ class Transformer(AutoModel):
|
|||||||
|
|
||||||
state_dict = dict(state_dict)
|
state_dict = dict(state_dict)
|
||||||
|
|
||||||
if self.config.tie_weight:
|
if self.config.tie_weight is True:
|
||||||
# same tensor for embed and lm_head
|
# same tensor for embed and lm_head
|
||||||
if embed_key in state_dict:
|
if embed_key in state_dict:
|
||||||
state_dict[lm_head_key] = state_dict[embed_key]
|
state_dict[lm_head_key] = state_dict[embed_key]
|
||||||
@@ -116,7 +121,7 @@ class Transformer(AutoModel):
|
|||||||
destination=destination, prefix=prefix, keep_vars=keep_vars
|
destination=destination, prefix=prefix, keep_vars=keep_vars
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.config.tie_weight:
|
if self.config.tie_weight is True:
|
||||||
lm_head_key = prefix + "lm_head.weight"
|
lm_head_key = prefix + "lm_head.weight"
|
||||||
if lm_head_key in state_dict:
|
if lm_head_key in state_dict:
|
||||||
del state_dict[lm_head_key]
|
del state_dict[lm_head_key]
|
||||||
@@ -127,18 +132,17 @@ class Transformer(AutoModel):
|
|||||||
self,
|
self,
|
||||||
input_ids: Tensor,
|
input_ids: Tensor,
|
||||||
input_mask: Optional[Tensor] = None,
|
input_mask: Optional[Tensor] = None,
|
||||||
paged_cache: Optional[CacheView] = None,
|
paged_cache: Optional[KvcacheView] = None,
|
||||||
start_pos: int = 0,
|
position_ids: Optional[Tensor] = None,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
assert input_ids.ndim == 2
|
assert input_ids.ndim == 2
|
||||||
|
|
||||||
x = self.embed_tokens(input_ids)
|
x = self.embed_tokens(input_ids)
|
||||||
rotary_emb = self.rotary_embedding(x, start_pos)
|
rotary_emb = self.rotary_embedding(x, position_ids)
|
||||||
|
attn_mask = process_attention_mask(x, position_ids, input_mask, is_causal=True)
|
||||||
attn_mask = process_attention_mask(input_mask, x, start_pos, is_causal=True)
|
|
||||||
|
|
||||||
for layer in self.layers:
|
for layer in self.layers:
|
||||||
x = layer(x, rotary_emb, attn_mask, paged_cache, start_pos)
|
x = layer(x, rotary_emb, attn_mask, paged_cache)
|
||||||
|
|
||||||
hidden_states = self.norm(x)
|
hidden_states = self.norm(x)
|
||||||
logits = self.lm_head(hidden_states)
|
logits = self.lm_head(hidden_states)
|
||||||
|
|||||||
+12
-16
@@ -1,7 +1,7 @@
|
|||||||
import os
|
import os
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from functools import wraps
|
from functools import wraps
|
||||||
from typing import Callable, List, Optional
|
from typing import Callable
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
@@ -34,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():
|
||||||
@@ -45,15 +44,10 @@ 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)]
|
|
||||||
device_id = torch.device(device_type, device_ids[rank])
|
|
||||||
|
|
||||||
os.environ["MASTER_ADDR"] = master_addr
|
os.environ["MASTER_ADDR"] = master_addr
|
||||||
os.environ["MASTER_PORT"] = master_port
|
os.environ["MASTER_PORT"] = master_port
|
||||||
|
|
||||||
os.environ["LOCAL_RANK"] = str(rank)
|
os.environ["LOCAL_RANK"] = str(rank)
|
||||||
os.environ["WORLD_SIZE"] = str(world_size)
|
os.environ["WORLD_SIZE"] = str(world_size)
|
||||||
os.environ["LOCAL_DEVICE"] = str(device_id)
|
os.environ["LOCAL_DEVICE"] = str(device_id)
|
||||||
@@ -103,7 +97,6 @@ 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,
|
||||||
):
|
):
|
||||||
@@ -115,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)
|
||||||
|
|
||||||
@@ -131,7 +123,7 @@ 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,
|
start_method: str = "spawn",
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
# clear environment variables
|
# clear environment variables
|
||||||
@@ -147,8 +139,9 @@ def spawn_parallel_fn(
|
|||||||
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)
|
||||||
device_id = torch.device(device_type, device_ids[0])
|
os.environ["LOCAL_RANK"] = "0"
|
||||||
|
os.environ["WORLD_SIZE"] = "1"
|
||||||
os.environ["LOCAL_DEVICE"] = str(device_id)
|
os.environ["LOCAL_DEVICE"] = str(device_id)
|
||||||
|
|
||||||
func(**kwargs)
|
func(**kwargs)
|
||||||
@@ -160,11 +153,14 @@ def spawn_parallel_fn(
|
|||||||
master_addr,
|
master_addr,
|
||||||
master_port,
|
master_port,
|
||||||
device_type,
|
device_type,
|
||||||
device_ids,
|
|
||||||
func,
|
func,
|
||||||
kwargs,
|
kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
mp.spawn(
|
mp.start_processes(
|
||||||
wrapper_spawn_func, nprocs=world_size, args=wrapper_spawn_func_args, join=True
|
wrapper_spawn_func,
|
||||||
|
args=wrapper_spawn_func_args,
|
||||||
|
nprocs=world_size,
|
||||||
|
start_method=start_method,
|
||||||
|
join=True,
|
||||||
)
|
)
|
||||||
|
|||||||
+17
-40
@@ -1,63 +1,29 @@
|
|||||||
import json
|
import json
|
||||||
import os
|
import time
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Dict, List
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
import h5py
|
|
||||||
import safetensors.torch as st
|
import safetensors.torch as st
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
from torch import Tensor
|
|
||||||
|
|
||||||
from astrai.parallel.setup import get_rank
|
from astrai.parallel.setup import get_rank
|
||||||
|
|
||||||
|
|
||||||
def save_h5(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
|
|
||||||
os.makedirs(file_path, exist_ok=True)
|
|
||||||
full_file_path = os.path.join(file_path, f"{file_name}.h5")
|
|
||||||
with h5py.File(full_file_path, "w") as f:
|
|
||||||
for key, tensors in tensor_group.items():
|
|
||||||
grp = f.create_group(key)
|
|
||||||
for idx, tensor in enumerate(tensors):
|
|
||||||
arr = tensor.cpu().numpy()
|
|
||||||
grp.create_dataset(f"data_{idx}", data=arr)
|
|
||||||
|
|
||||||
|
|
||||||
def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
|
|
||||||
tensor_group: Dict[str, List[Tensor]] = {}
|
|
||||||
|
|
||||||
root_path = Path(file_path)
|
|
||||||
h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5"))
|
|
||||||
|
|
||||||
for h5_file in h5_files:
|
|
||||||
with h5py.File(h5_file, "r") as f:
|
|
||||||
for key in f.keys():
|
|
||||||
grp = f[key]
|
|
||||||
dsets = []
|
|
||||||
for dset_name in grp.keys():
|
|
||||||
dset = grp[dset_name]
|
|
||||||
tensor = torch.from_numpy(dset[:])
|
|
||||||
if share_memory:
|
|
||||||
tensor = tensor.share_memory_()
|
|
||||||
dsets.append(tensor)
|
|
||||||
|
|
||||||
if tensor_group.get(key) is None:
|
|
||||||
tensor_group[key] = []
|
|
||||||
tensor_group[key].extend(dsets)
|
|
||||||
|
|
||||||
return tensor_group
|
|
||||||
|
|
||||||
|
|
||||||
class Checkpoint:
|
class Checkpoint:
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
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,
|
||||||
|
meta: 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 {}
|
||||||
|
self.meta = meta or {}
|
||||||
|
|
||||||
def save(
|
def save(
|
||||||
self,
|
self,
|
||||||
@@ -72,11 +38,16 @@ class Checkpoint:
|
|||||||
meta = {
|
meta = {
|
||||||
"epoch": self.epoch,
|
"epoch": self.epoch,
|
||||||
"iteration": self.iteration,
|
"iteration": self.iteration,
|
||||||
|
"timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"),
|
||||||
}
|
}
|
||||||
|
meta.update(self.meta)
|
||||||
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)
|
||||||
|
|
||||||
st.save_file(self.state_dict, save_path / "state_dict.safetensors")
|
st.save_file(self.state_dict, save_path / "state_dict.safetensors")
|
||||||
|
if self.extra:
|
||||||
|
for key, value in self.extra.items():
|
||||||
|
torch.save(value, save_path / f"{key}.pt")
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def load(
|
def load(
|
||||||
@@ -99,8 +70,14 @@ class Checkpoint:
|
|||||||
|
|
||||||
state_dict = st.load_file(save_path / "state_dict.safetensors")
|
state_dict = st.load_file(save_path / "state_dict.safetensors")
|
||||||
|
|
||||||
|
extra = {}
|
||||||
|
for f in save_path.iterdir():
|
||||||
|
if f.suffix == ".pt" and f.stem not in ("meta",):
|
||||||
|
extra[f.stem] = torch.load(f, 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 or None,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -51,9 +51,26 @@ class AutoTokenizer:
|
|||||||
self.set_chat_template(config["chat_template"])
|
self.set_chat_template(config["chat_template"])
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_pretrained(cls, path: Union[str, Path], **kwargs) -> "AutoTokenizer":
|
def from_pretrained(cls, path: Union[str, Path]) -> "AutoTokenizer":
|
||||||
"""Load tokenizer from pretrained directory."""
|
"""Load tokenizer from pretrained directory.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
FileNotFoundError: If tokenizer.json is missing.
|
||||||
|
RuntimeError: If tokenizer failed to initialize.
|
||||||
|
"""
|
||||||
|
path = Path(path)
|
||||||
|
tokenizer_file = path / "tokenizer.json"
|
||||||
|
if not tokenizer_file.exists():
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"Tokenizer file not found: {tokenizer_file}. "
|
||||||
|
"A valid tokenizer.json is required."
|
||||||
|
)
|
||||||
instance = cls(path)
|
instance = cls(path)
|
||||||
|
if instance._tokenizer is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Failed to load tokenizer from {path}. "
|
||||||
|
"The tokenizer.json may be corrupted or incompatible."
|
||||||
|
)
|
||||||
return instance
|
return instance
|
||||||
|
|
||||||
def save_pretrained(self, save_path: str):
|
def save_pretrained(self, save_path: str):
|
||||||
@@ -64,6 +81,11 @@ class AutoTokenizer:
|
|||||||
save_path: Path to save the tokenizer
|
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 = Path(save_path)
|
||||||
save_path.mkdir(parents=True, exist_ok=True)
|
save_path.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
from astrai.trainer.optim import Muon
|
||||||
from astrai.trainer.schedule import BaseScheduler, SchedulerFactory
|
from astrai.trainer.schedule import BaseScheduler, SchedulerFactory
|
||||||
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
||||||
from astrai.trainer.train_callback import (
|
from astrai.trainer.train_callback import (
|
||||||
@@ -9,6 +10,8 @@ from astrai.trainer.trainer import Trainer
|
|||||||
__all__ = [
|
__all__ = [
|
||||||
# Main trainer
|
# Main trainer
|
||||||
"Trainer",
|
"Trainer",
|
||||||
|
# Optimizer
|
||||||
|
"Muon",
|
||||||
# Strategy factory
|
# Strategy factory
|
||||||
"StrategyFactory",
|
"StrategyFactory",
|
||||||
"BaseStrategy",
|
"BaseStrategy",
|
||||||
|
|||||||
@@ -1,75 +1,42 @@
|
|||||||
from typing import Dict
|
from typing import Any, Callable, Dict
|
||||||
|
|
||||||
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
|
|
||||||
def grad_norm(model: nn.Module, norm_type: int = 2) -> Dict[str, float]:
|
def _grad_stat(
|
||||||
"""Compute gradient norm for each parameter in the model."""
|
model: nn.Module, fn: Callable[[torch.Tensor], Any], default: Any
|
||||||
norms = {}
|
) -> dict:
|
||||||
|
results = {}
|
||||||
for name, param in model.named_parameters():
|
for name, param in model.named_parameters():
|
||||||
norms[name] = 0.0
|
results[name] = default
|
||||||
if param.grad:
|
if param.grad is not None:
|
||||||
norm = param.grad.data.norm(norm_type).item()
|
results[name] = fn(param.grad.data)
|
||||||
norms[name] = norm
|
return results
|
||||||
return norms
|
|
||||||
|
|
||||||
|
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]:
|
def grad_std(model: nn.Module) -> Dict[str, float]:
|
||||||
"""Compute standard deviation of gradients for each parameter."""
|
return _grad_stat(model, lambda g: g.std().item(), 0.0)
|
||||||
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]:
|
def grad_max(model: nn.Module) -> Dict[str, float]:
|
||||||
"""Find the maximum absolute gradient value for each parameter."""
|
return _grad_stat(model, lambda g: g.max().item(), -float("inf"))
|
||||||
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]:
|
def grad_min(model: nn.Module) -> Dict[str, float]:
|
||||||
"""Find the minimum absolute gradient value for each parameter."""
|
return _grad_stat(model, lambda g: g.min().item(), float("inf"))
|
||||||
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]:
|
def grad_mean(model: nn.Module) -> Dict[str, float]:
|
||||||
"""Compute mean of gradients for each parameter."""
|
return _grad_stat(model, lambda g: g.mean().item(), 0.0)
|
||||||
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]:
|
def grad_nan_num(model: nn.Module) -> Dict[str, int]:
|
||||||
"""Count the number of NaNs in gradients for each parameter."""
|
return _grad_stat(model, lambda g: g.isnan().sum().item(), 0)
|
||||||
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):
|
def ctx_get_loss(ctx):
|
||||||
@@ -80,6 +47,10 @@ def ctx_get_lr(ctx):
|
|||||||
return ctx.optimizer.param_groups[-1]["lr"]
|
return ctx.optimizer.param_groups[-1]["lr"]
|
||||||
|
|
||||||
|
|
||||||
|
def ctx_get_val_loss(ctx):
|
||||||
|
return ctx.val_loss
|
||||||
|
|
||||||
|
|
||||||
def ctx_get_grad_norm(ctx):
|
def ctx_get_grad_norm(ctx):
|
||||||
return grad_norm(ctx.model)
|
return grad_norm(ctx.model)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,113 @@
|
|||||||
|
import torch
|
||||||
|
from torch.optim import Optimizer
|
||||||
|
|
||||||
|
|
||||||
|
def _zeropower_via_newtonschulz(G: torch.Tensor, steps: int = 5):
|
||||||
|
assert G.ndim == 2
|
||||||
|
X = G.bfloat16()
|
||||||
|
scale = max(1, G.size(0) / G.size(1)) ** 0.5
|
||||||
|
X = X / (X.norm() + 1e-7) * scale
|
||||||
|
if steps == 0:
|
||||||
|
return X.type_as(G)
|
||||||
|
a, b, c = (3.4445, -4.7750, 2.0315)
|
||||||
|
for _ in range(steps):
|
||||||
|
A = X @ X.T
|
||||||
|
B = A @ X
|
||||||
|
X = a * X + b * B + c * (A @ B)
|
||||||
|
return X.type_as(G)
|
||||||
|
|
||||||
|
|
||||||
|
class Muon(Optimizer):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
params,
|
||||||
|
lr: float = 2e-3,
|
||||||
|
momentum: float = 0.95,
|
||||||
|
weight_decay: float = 0.0,
|
||||||
|
nesterov: bool = True,
|
||||||
|
ns_steps: int = 5,
|
||||||
|
adamw_lr: float = None,
|
||||||
|
adamw_betas: tuple = (0.9, 0.95),
|
||||||
|
adamw_eps: float = 1e-8,
|
||||||
|
adamw_wd: float = 0.0,
|
||||||
|
):
|
||||||
|
defaults = dict(
|
||||||
|
lr=lr,
|
||||||
|
momentum=momentum,
|
||||||
|
weight_decay=weight_decay,
|
||||||
|
nesterov=nesterov,
|
||||||
|
ns_steps=ns_steps,
|
||||||
|
adamw_lr=adamw_lr if adamw_lr is not None else lr * 0.1,
|
||||||
|
adamw_betas=adamw_betas,
|
||||||
|
adamw_eps=adamw_eps,
|
||||||
|
adamw_wd=adamw_wd,
|
||||||
|
)
|
||||||
|
super().__init__(params, defaults)
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def step(self, closure=None):
|
||||||
|
loss = None
|
||||||
|
if closure is not None:
|
||||||
|
with torch.enable_grad():
|
||||||
|
loss = closure()
|
||||||
|
for group in self.param_groups:
|
||||||
|
for p in group["params"]:
|
||||||
|
if p.grad is None:
|
||||||
|
continue
|
||||||
|
grad = p.grad
|
||||||
|
if grad.is_sparse:
|
||||||
|
raise RuntimeError("Muon does not support sparse gradients")
|
||||||
|
if p.ndim >= 2:
|
||||||
|
self._muon_update(p, grad, group)
|
||||||
|
else:
|
||||||
|
self._adamw_update(p, grad, group)
|
||||||
|
return loss
|
||||||
|
|
||||||
|
def _muon_update(self, p, grad, group):
|
||||||
|
lr = group["lr"]
|
||||||
|
momentum = group["momentum"]
|
||||||
|
wd = group["weight_decay"]
|
||||||
|
nesterov = group["nesterov"]
|
||||||
|
ns_steps = group["ns_steps"]
|
||||||
|
state = self.state[p]
|
||||||
|
|
||||||
|
p.mul_(1 - lr * wd)
|
||||||
|
|
||||||
|
if nesterov:
|
||||||
|
grad = grad.add(p, alpha=wd)
|
||||||
|
|
||||||
|
if "momentum_buffer" not in state:
|
||||||
|
state["momentum_buffer"] = torch.zeros_like(grad)
|
||||||
|
buf = state["momentum_buffer"]
|
||||||
|
buf.lerp_(grad, 1 - momentum)
|
||||||
|
|
||||||
|
update = _zeropower_via_newtonschulz(buf, steps=ns_steps)
|
||||||
|
scale = max(1, p.size(0) / p.size(1)) ** 0.5
|
||||||
|
p.add_(update, alpha=-lr * scale)
|
||||||
|
|
||||||
|
def _adamw_update(self, p, grad, group):
|
||||||
|
lr = group["adamw_lr"]
|
||||||
|
betas = group["adamw_betas"]
|
||||||
|
eps = group["adamw_eps"]
|
||||||
|
wd = group["adamw_wd"]
|
||||||
|
state = self.state[p]
|
||||||
|
|
||||||
|
if not state:
|
||||||
|
state["step"] = 0
|
||||||
|
state["exp_avg"] = torch.zeros_like(p)
|
||||||
|
state["exp_avg_sq"] = torch.zeros_like(p)
|
||||||
|
|
||||||
|
state["step"] += 1
|
||||||
|
exp_avg, exp_avg_sq = state["exp_avg"], state["exp_avg_sq"]
|
||||||
|
beta1, beta2 = betas
|
||||||
|
|
||||||
|
exp_avg.lerp_(grad, 1 - beta1)
|
||||||
|
exp_avg_sq.lerp_(grad.square(), 1 - beta2)
|
||||||
|
|
||||||
|
step = state["step"]
|
||||||
|
bias1 = 1 - beta1**step
|
||||||
|
bias2 = 1 - beta2**step
|
||||||
|
|
||||||
|
p.mul_(1 - lr * wd)
|
||||||
|
denom = exp_avg_sq.sqrt().div_(bias2**0.5).add_(eps)
|
||||||
|
p.addcdiv_(exp_avg / bias1, denom, value=-lr)
|
||||||
@@ -1,15 +1,21 @@
|
|||||||
import json
|
import json
|
||||||
|
import logging
|
||||||
import os
|
import os
|
||||||
|
import sys
|
||||||
import time
|
import time
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Callable, List, Optional, Protocol, runtime_checkable
|
from typing import IO, Callable, List, Optional, Protocol, runtime_checkable
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch.nn.utils import clip_grad_norm_
|
from torch.nn.utils import clip_grad_norm_
|
||||||
|
from torch.utils.checkpoint import checkpoint as torch_checkpoint
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.parallel import only_on_rank
|
from astrai.parallel import only_on_rank
|
||||||
|
from astrai.parallel.setup import get_current_device
|
||||||
from astrai.serialization import Checkpoint
|
from astrai.serialization import Checkpoint
|
||||||
from astrai.trainer.metric_util import (
|
from astrai.trainer.metric_util import (
|
||||||
ctx_get_grad_max,
|
ctx_get_grad_max,
|
||||||
@@ -20,9 +26,12 @@ from astrai.trainer.metric_util import (
|
|||||||
ctx_get_grad_std,
|
ctx_get_grad_std,
|
||||||
ctx_get_loss,
|
ctx_get_loss,
|
||||||
ctx_get_lr,
|
ctx_get_lr,
|
||||||
|
ctx_get_val_loss,
|
||||||
)
|
)
|
||||||
from astrai.trainer.train_context import TrainContext
|
from astrai.trainer.train_context import TrainContext
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@runtime_checkable
|
@runtime_checkable
|
||||||
class TrainCallback(Protocol):
|
class TrainCallback(Protocol):
|
||||||
@@ -69,12 +78,6 @@ class CallbackFactory(BaseFactory[TrainCallback]):
|
|||||||
callback = CallbackFactory.create("my_callback", **kwargs)
|
callback = CallbackFactory.create("my_callback", **kwargs)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def _validate_component(cls, callback_cls: type) -> None:
|
|
||||||
"""Validate that the callback class inherits from TrainCallback."""
|
|
||||||
if not issubclass(callback_cls, TrainCallback):
|
|
||||||
raise TypeError(f"{callback_cls.__name__} must inherit from TrainCallback")
|
|
||||||
|
|
||||||
|
|
||||||
@CallbackFactory.register("gradient_clipping")
|
@CallbackFactory.register("gradient_clipping")
|
||||||
class GradientClippingCallback(TrainCallback):
|
class GradientClippingCallback(TrainCallback):
|
||||||
@@ -86,27 +89,42 @@ class GradientClippingCallback(TrainCallback):
|
|||||||
self.max_grad_norm = max_grad_norm
|
self.max_grad_norm = max_grad_norm
|
||||||
|
|
||||||
def on_step_begin(self, context: TrainContext):
|
def on_step_begin(self, context: TrainContext):
|
||||||
_ = context
|
|
||||||
clip_grad_norm_(context.model.parameters(), self.max_grad_norm)
|
clip_grad_norm_(context.model.parameters(), self.max_grad_norm)
|
||||||
|
|
||||||
|
|
||||||
@CallbackFactory.register("scheduler")
|
@CallbackFactory.register("gradient_checkpointing")
|
||||||
class SchedulerCallback(TrainCallback):
|
class GradientCheckpointingCallback(TrainCallback):
|
||||||
"""
|
"""
|
||||||
Scheduler callback for trainer.
|
Activation checkpointing callback — trades compute for memory
|
||||||
|
by recomputing specified module activations during the backward pass.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
modules: Module types to apply checkpointing to.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self, modules: Optional[List[type]] = None):
|
||||||
pass
|
self.modules = tuple(modules) if modules else ()
|
||||||
|
|
||||||
|
def _enable(self, module: nn.Module):
|
||||||
|
if self.modules and isinstance(module, self.modules):
|
||||||
|
fn = module.forward
|
||||||
|
module._original_forward = fn
|
||||||
|
module.forward = lambda *a, **kw: torch_checkpoint(
|
||||||
|
fn, *a, use_reentrant=False, **kw
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _disable(module: nn.Module):
|
||||||
|
if hasattr(module, "_original_forward"):
|
||||||
|
module.forward = module._original_forward
|
||||||
|
del module._original_forward
|
||||||
|
|
||||||
def on_train_begin(self, context: TrainContext):
|
def on_train_begin(self, context: TrainContext):
|
||||||
for group in context.optimizer.param_groups:
|
context.model.apply(self._enable)
|
||||||
if "initial_lr" not in group:
|
logger.info("Gradient checkpointing enabled")
|
||||||
group["initial_lr"] = group["lr"]
|
|
||||||
|
|
||||||
def on_batch_end(self, context: TrainContext):
|
def on_train_end(self, context: TrainContext):
|
||||||
if context.scheduler:
|
context.model.apply(self._disable)
|
||||||
context.scheduler.step()
|
|
||||||
|
|
||||||
|
|
||||||
@CallbackFactory.register("checkpoint")
|
@CallbackFactory.register("checkpoint")
|
||||||
@@ -115,17 +133,23 @@ class CheckpointCallback(TrainCallback):
|
|||||||
Checkpoint callback for trainer.
|
Checkpoint callback for trainer.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
extra_keys = ("optimizer", "scheduler")
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
save_dir: str,
|
save_dir: str,
|
||||||
interval: int,
|
interval: int,
|
||||||
weight_only: bool = False,
|
weight_only: bool = False,
|
||||||
state_dict_fn: Optional[Callable[[nn.Module], dict]] = None,
|
state_dict_fn: Optional[Callable[[nn.Module], dict]] = None,
|
||||||
|
save_extra_fn: Optional[Callable[["TrainContext"], dict]] = None,
|
||||||
|
load_extra_fn: Optional[Callable[[dict, "TrainContext"], None]] = 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 or CheckpointCallback.save_extra
|
||||||
|
self.load_extra_fn = load_extra_fn or CheckpointCallback.load_extra
|
||||||
self.last_ckpt_iter = 0
|
self.last_ckpt_iter = 0
|
||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
@@ -139,13 +163,22 @@ class CheckpointCallback(TrainCallback):
|
|||||||
else context.model.state_dict()
|
else context.model.state_dict()
|
||||||
)
|
)
|
||||||
|
|
||||||
|
extra = self.save_extra_fn(context)
|
||||||
context.checkpoint = Checkpoint(
|
context.checkpoint = Checkpoint(
|
||||||
state_dict=state_dict, epoch=context.epoch, iteration=context.iteration
|
state_dict=state_dict,
|
||||||
|
epoch=context.epoch,
|
||||||
|
iteration=context.iteration,
|
||||||
|
extra=extra,
|
||||||
|
meta=context.config.to_dict(),
|
||||||
)
|
)
|
||||||
|
|
||||||
context.checkpoint.save(save_path)
|
context.checkpoint.save(save_path)
|
||||||
self.last_ckpt_iter = context.iteration
|
self.last_ckpt_iter = context.iteration
|
||||||
|
|
||||||
|
def on_train_begin(self, context: TrainContext):
|
||||||
|
if context.checkpoint and context.checkpoint.extra:
|
||||||
|
self.load_extra_fn(context.checkpoint.extra, context)
|
||||||
|
|
||||||
def on_batch_end(self, context: TrainContext):
|
def on_batch_end(self, context: TrainContext):
|
||||||
if context.iteration - self.last_ckpt_iter >= self.interval:
|
if context.iteration - self.last_ckpt_iter >= self.interval:
|
||||||
self._save_checkpoint(context)
|
self._save_checkpoint(context)
|
||||||
@@ -157,6 +190,21 @@ class CheckpointCallback(TrainCallback):
|
|||||||
def on_error(self, context: TrainContext):
|
def on_error(self, context: TrainContext):
|
||||||
self._save_checkpoint(context)
|
self._save_checkpoint(context)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def save_extra(context: TrainContext) -> dict:
|
||||||
|
extra = {}
|
||||||
|
for name in CheckpointCallback.extra_keys:
|
||||||
|
obj = getattr(context, name, None)
|
||||||
|
if obj:
|
||||||
|
extra[name] = obj.state_dict()
|
||||||
|
return extra
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def load_extra(extra: dict, context: TrainContext):
|
||||||
|
for name in CheckpointCallback.extra_keys:
|
||||||
|
if name in extra:
|
||||||
|
getattr(context, name).load_state_dict(extra[name])
|
||||||
|
|
||||||
|
|
||||||
@CallbackFactory.register("progress_bar")
|
@CallbackFactory.register("progress_bar")
|
||||||
class ProgressBarCallback(TrainCallback):
|
class ProgressBarCallback(TrainCallback):
|
||||||
@@ -164,8 +212,12 @@ class ProgressBarCallback(TrainCallback):
|
|||||||
Progress bar callback for trainer.
|
Progress bar callback for trainer.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, num_epoch: int):
|
def __init__(
|
||||||
|
self, num_epoch: int, log_interval: int = 100, file: IO[str] = sys.stdout
|
||||||
|
):
|
||||||
self.num_epoch = num_epoch
|
self.num_epoch = num_epoch
|
||||||
|
self.log_interval = log_interval
|
||||||
|
self.file = file
|
||||||
self.progress_bar: tqdm = None
|
self.progress_bar: tqdm = None
|
||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
@@ -174,16 +226,18 @@ class ProgressBarCallback(TrainCallback):
|
|||||||
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,
|
||||||
|
file=self.file,
|
||||||
)
|
)
|
||||||
|
|
||||||
@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(
|
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}",
|
}
|
||||||
}
|
if context.val_loss > 0:
|
||||||
)
|
postfix["val_loss"] = f"{context.val_loss:.4f}"
|
||||||
|
self.progress_bar.set_postfix(postfix)
|
||||||
self.progress_bar.update(1)
|
self.progress_bar.update(1)
|
||||||
|
|
||||||
@only_on_rank(0)
|
@only_on_rank(0)
|
||||||
@@ -215,6 +269,7 @@ class MetricLoggerCallback(TrainCallback):
|
|||||||
self._metric_funcs = {
|
self._metric_funcs = {
|
||||||
"loss": ctx_get_loss,
|
"loss": ctx_get_loss,
|
||||||
"lr": ctx_get_lr,
|
"lr": ctx_get_lr,
|
||||||
|
"val_loss": ctx_get_val_loss,
|
||||||
"grad_norm": ctx_get_grad_norm,
|
"grad_norm": ctx_get_grad_norm,
|
||||||
"grad_std": ctx_get_grad_std,
|
"grad_std": ctx_get_grad_std,
|
||||||
"grad_max": ctx_get_grad_max,
|
"grad_max": ctx_get_grad_max,
|
||||||
@@ -225,7 +280,7 @@ class MetricLoggerCallback(TrainCallback):
|
|||||||
|
|
||||||
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-%dT%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},
|
||||||
@@ -258,3 +313,43 @@ class MetricLoggerCallback(TrainCallback):
|
|||||||
|
|
||||||
def on_error(self, context):
|
def on_error(self, context):
|
||||||
self._save_log(context.epoch, context.iteration)
|
self._save_log(context.epoch, context.iteration)
|
||||||
|
|
||||||
|
|
||||||
|
@CallbackFactory.register("validation")
|
||||||
|
class ValidationCallback(TrainCallback):
|
||||||
|
def _run_validation(self, context: TrainContext):
|
||||||
|
context.model.eval()
|
||||||
|
|
||||||
|
total_loss = 0.0
|
||||||
|
num_batches = 0
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
for batch in context.val_dataloader:
|
||||||
|
loss = context.strategy(batch)
|
||||||
|
total_loss += loss.item()
|
||||||
|
num_batches += 1
|
||||||
|
|
||||||
|
avg_loss = total_loss / max(num_batches, 1)
|
||||||
|
|
||||||
|
if context.world_size > 1 and dist.is_initialized():
|
||||||
|
loss_tensor = torch.tensor([avg_loss], device=get_current_device())
|
||||||
|
dist.all_reduce(loss_tensor, op=dist.ReduceOp.AVG)
|
||||||
|
avg_loss = loss_tensor.item()
|
||||||
|
|
||||||
|
context.val_loss = avg_loss
|
||||||
|
context.model.train()
|
||||||
|
|
||||||
|
step_count = context.iteration // context.config.grad_accum_steps
|
||||||
|
logger.info(
|
||||||
|
f"Epoch {context.epoch + 1}, Step {step_count}, Val Loss: {avg_loss:.4f}"
|
||||||
|
)
|
||||||
|
|
||||||
|
def on_step_end(self, context: TrainContext):
|
||||||
|
if context.val_dataloader is None:
|
||||||
|
return
|
||||||
|
cfg = context.config
|
||||||
|
if cfg.val_step <= 0:
|
||||||
|
return
|
||||||
|
step_count = context.iteration // cfg.grad_accum_steps
|
||||||
|
if step_count % cfg.val_step == 0:
|
||||||
|
self._run_validation(context)
|
||||||
|
|||||||
@@ -21,10 +21,13 @@ class TrainContext:
|
|||||||
optimizer: Optimizer = field(default=None)
|
optimizer: Optimizer = field(default=None)
|
||||||
scheduler: LRScheduler = field(default=None)
|
scheduler: LRScheduler = field(default=None)
|
||||||
checkpoint: Checkpoint = field(default=None)
|
checkpoint: Checkpoint = field(default=None)
|
||||||
|
config: TrainConfig = field(default=None)
|
||||||
|
|
||||||
epoch: int = field(default=0)
|
epoch: int = field(default=0)
|
||||||
iteration: int = field(default=0)
|
iteration: int = field(default=0)
|
||||||
loss: float = field(default=0.0)
|
loss: float = field(default=0.0)
|
||||||
|
val_dataloader: DataLoader = field(default=None)
|
||||||
|
val_loss: float = field(default=0.0)
|
||||||
|
|
||||||
world_size: int = field(default=1)
|
world_size: int = field(default=1)
|
||||||
rank: int = field(default=0)
|
rank: int = field(default=0)
|
||||||
@@ -32,7 +35,10 @@ class TrainContext:
|
|||||||
|
|
||||||
|
|
||||||
class TrainContextBuilder:
|
class TrainContextBuilder:
|
||||||
def __init__(self, config: TrainConfig):
|
def __init__(
|
||||||
|
self,
|
||||||
|
config: TrainConfig,
|
||||||
|
):
|
||||||
self.config = config
|
self.config = config
|
||||||
self._checkpoint: Optional[Checkpoint] = None
|
self._checkpoint: Optional[Checkpoint] = None
|
||||||
|
|
||||||
@@ -45,6 +51,7 @@ class TrainContextBuilder:
|
|||||||
model=self.config.model,
|
model=self.config.model,
|
||||||
world_size=get_world_size(),
|
world_size=get_world_size(),
|
||||||
rank=get_rank(),
|
rank=get_rank(),
|
||||||
|
config=self.config,
|
||||||
)
|
)
|
||||||
|
|
||||||
device = get_current_device()
|
device = get_current_device()
|
||||||
@@ -67,7 +74,7 @@ class TrainContextBuilder:
|
|||||||
context.scheduler = self.config.scheduler_fn(context.optimizer)
|
context.scheduler = self.config.scheduler_fn(context.optimizer)
|
||||||
|
|
||||||
cfg = self.config
|
cfg = self.config
|
||||||
sampler_offset = context.iteration * cfg.batch_size
|
sampler_offset = context.iteration * cfg.batch_per_device
|
||||||
sampler = ResumableDistributedSampler(
|
sampler = ResumableDistributedSampler(
|
||||||
data_source=cfg.dataset,
|
data_source=cfg.dataset,
|
||||||
start_epoch=context.epoch,
|
start_epoch=context.epoch,
|
||||||
@@ -76,13 +83,30 @@ class TrainContextBuilder:
|
|||||||
)
|
)
|
||||||
context.dataloader = DataLoader(
|
context.dataloader = DataLoader(
|
||||||
cfg.dataset,
|
cfg.dataset,
|
||||||
batch_size=cfg.batch_size,
|
batch_size=cfg.batch_per_device,
|
||||||
sampler=sampler,
|
sampler=sampler,
|
||||||
num_workers=cfg.num_workers,
|
num_workers=cfg.num_workers,
|
||||||
pin_memory=cfg.pin_memory,
|
pin_memory=cfg.pin_memory,
|
||||||
prefetch_factor=cfg.prefetch_factor,
|
prefetch_factor=cfg.prefetch_factor,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if cfg.val_dataset is not None:
|
||||||
|
val_sampler = ResumableDistributedSampler(
|
||||||
|
data_source=cfg.val_dataset,
|
||||||
|
start_epoch=0,
|
||||||
|
start_iter=0,
|
||||||
|
seed=cfg.random_seed,
|
||||||
|
shuffle=False,
|
||||||
|
)
|
||||||
|
context.val_dataloader = DataLoader(
|
||||||
|
cfg.val_dataset,
|
||||||
|
batch_size=cfg.batch_per_device,
|
||||||
|
sampler=val_sampler,
|
||||||
|
num_workers=cfg.num_workers,
|
||||||
|
pin_memory=cfg.pin_memory,
|
||||||
|
prefetch_factor=cfg.prefetch_factor,
|
||||||
|
)
|
||||||
|
|
||||||
context.strategy = StrategyFactory.create(
|
context.strategy = StrategyFactory.create(
|
||||||
model=context.model,
|
model=context.model,
|
||||||
train_type=self.config.strategy,
|
train_type=self.config.strategy,
|
||||||
|
|||||||
+50
-39
@@ -25,18 +25,29 @@ class Trainer:
|
|||||||
|
|
||||||
def _get_default_callbacks(self) -> List[TrainCallback]:
|
def _get_default_callbacks(self) -> List[TrainCallback]:
|
||||||
cfg = self.train_config
|
cfg = self.train_config
|
||||||
return [
|
callbacks = [
|
||||||
|
CallbackFactory.create(
|
||||||
|
"gradient_checkpointing",
|
||||||
|
modules=cfg.gradient_checkpointing_modules,
|
||||||
|
),
|
||||||
|
CallbackFactory.create(
|
||||||
|
"checkpoint",
|
||||||
|
cfg.ckpt_dir,
|
||||||
|
cfg.ckpt_interval,
|
||||||
|
state_dict_fn=cfg.state_dict_fn,
|
||||||
|
),
|
||||||
|
CallbackFactory.create(
|
||||||
|
"metric_logger",
|
||||||
|
log_dir=cfg.log_dir,
|
||||||
|
save_interval=cfg.ckpt_interval,
|
||||||
|
log_interval=cfg.log_interval,
|
||||||
|
metrics=cfg.metrics,
|
||||||
|
),
|
||||||
CallbackFactory.create("progress_bar", cfg.n_epoch),
|
CallbackFactory.create("progress_bar", cfg.n_epoch),
|
||||||
CallbackFactory.create("checkpoint", cfg.ckpt_dir, cfg.ckpt_interval),
|
|
||||||
CallbackFactory.create("metric_logger", cfg.ckpt_dir, cfg.ckpt_interval),
|
|
||||||
CallbackFactory.create("gradient_clipping", cfg.max_grad_norm),
|
CallbackFactory.create("gradient_clipping", cfg.max_grad_norm),
|
||||||
CallbackFactory.create("scheduler"),
|
CallbackFactory.create("validation"),
|
||||||
]
|
]
|
||||||
|
return callbacks
|
||||||
def _build_context(self, checkpoint: Optional[Checkpoint]) -> TrainContext:
|
|
||||||
return (
|
|
||||||
TrainContextBuilder(self.train_config).with_checkpoint(checkpoint).build()
|
|
||||||
)
|
|
||||||
|
|
||||||
def _call_callbacks(self, method_name: str, context: TrainContext):
|
def _call_callbacks(self, method_name: str, context: TrainContext):
|
||||||
for callback in self.callbacks:
|
for callback in self.callbacks:
|
||||||
@@ -44,49 +55,36 @@ class Trainer:
|
|||||||
if method:
|
if method:
|
||||||
method(context)
|
method(context)
|
||||||
|
|
||||||
def train(self, checkpoint: Optional[Checkpoint] = None):
|
def _trainer_loop(self, checkpoint: Optional[Checkpoint] = None):
|
||||||
config = self.train_config
|
cfg = self.train_config
|
||||||
spawn_parallel_fn(
|
context = TrainContextBuilder(cfg).with_checkpoint(checkpoint).build()
|
||||||
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,
|
|
||||||
device_ids=config.device_ids,
|
|
||||||
checkpoint=checkpoint,
|
|
||||||
)
|
|
||||||
|
|
||||||
def _train_impl(self, checkpoint: Optional[Checkpoint] = None) -> Checkpoint:
|
|
||||||
context = self._build_context(checkpoint)
|
|
||||||
self._call_callbacks("on_train_begin", context)
|
self._call_callbacks("on_train_begin", context)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
context.model.train()
|
context.model.train()
|
||||||
# 1.epoch
|
grad_accum_steps = cfg.grad_accum_steps
|
||||||
for epoch in range(context.epoch, self.train_config.n_epoch):
|
|
||||||
|
for epoch in range(context.epoch, cfg.n_epoch):
|
||||||
context.epoch = epoch
|
context.epoch = epoch
|
||||||
self._call_callbacks("on_epoch_begin", context)
|
self._call_callbacks("on_epoch_begin", context)
|
||||||
|
|
||||||
for batch in context.dataloader:
|
for batch in context.dataloader:
|
||||||
if context.iteration % self.train_config.accumulation_steps == 0:
|
self._call_callbacks("on_batch_begin", context)
|
||||||
# 2. step
|
loss = context.strategy(batch)
|
||||||
|
context.loss = loss.item()
|
||||||
|
stand_loss = loss / grad_accum_steps
|
||||||
|
stand_loss.backward()
|
||||||
|
context.iteration += 1
|
||||||
|
self._call_callbacks("on_batch_end", context)
|
||||||
|
|
||||||
|
if context.iteration % grad_accum_steps == 0:
|
||||||
self._call_callbacks("on_step_begin", context)
|
self._call_callbacks("on_step_begin", context)
|
||||||
context.optimizer.step()
|
context.optimizer.step()
|
||||||
context.optimizer.zero_grad()
|
context.optimizer.zero_grad()
|
||||||
self._call_callbacks("on_step_end", context)
|
self._call_callbacks("on_step_end", context)
|
||||||
|
|
||||||
# 3. batch
|
if context.scheduler:
|
||||||
self._call_callbacks("on_batch_begin", context)
|
context.scheduler.step()
|
||||||
loss = context.strategy(batch)
|
|
||||||
context.loss = loss.item()
|
|
||||||
context.iteration += 1
|
|
||||||
|
|
||||||
# to make the loss normalized by accumulation steps
|
|
||||||
stand_loss = loss / self.train_config.accumulation_steps
|
|
||||||
stand_loss.backward()
|
|
||||||
|
|
||||||
self._call_callbacks("on_batch_end", context)
|
|
||||||
|
|
||||||
self._call_callbacks("on_epoch_end", context)
|
self._call_callbacks("on_epoch_end", context)
|
||||||
|
|
||||||
@@ -96,3 +94,16 @@ class Trainer:
|
|||||||
raise
|
raise
|
||||||
finally:
|
finally:
|
||||||
self._call_callbacks("on_train_end", context)
|
self._call_callbacks("on_train_end", context)
|
||||||
|
|
||||||
|
def train(self, checkpoint: Optional[Checkpoint] = None):
|
||||||
|
cfg = self.train_config
|
||||||
|
spawn_parallel_fn(
|
||||||
|
self._trainer_loop,
|
||||||
|
backend=cfg.backend,
|
||||||
|
world_size=cfg.nprocs,
|
||||||
|
master_addr=cfg.master_addr,
|
||||||
|
master_port=cfg.master_port,
|
||||||
|
device_type=cfg.device_type,
|
||||||
|
start_method=cfg.start_method,
|
||||||
|
checkpoint=checkpoint,
|
||||||
|
)
|
||||||
|
|||||||
+8
-6
@@ -1,12 +1,13 @@
|
|||||||
services:
|
services:
|
||||||
server:
|
server:
|
||||||
build: .
|
build:
|
||||||
image: astrai:latest
|
context: .
|
||||||
|
dockerfile: Dockerfile
|
||||||
|
user: "${UID:-1000}:${GID:-1000}"
|
||||||
ports:
|
ports:
|
||||||
- "8000:8000"
|
- "8000:8000"
|
||||||
volumes:
|
volumes:
|
||||||
- ./params:/app/params:ro
|
- ./params:/app/params:ro
|
||||||
- ./checkpoints:/app/checkpoints
|
|
||||||
command: python -m scripts.tools.server --port 8000 --device cuda
|
command: python -m scripts.tools.server --port 8000 --device cuda
|
||||||
deploy:
|
deploy:
|
||||||
resources:
|
resources:
|
||||||
@@ -25,13 +26,14 @@ services:
|
|||||||
|
|
||||||
server-cpu:
|
server-cpu:
|
||||||
profiles: [cpu]
|
profiles: [cpu]
|
||||||
build: .
|
build:
|
||||||
image: astrai:latest
|
context: .
|
||||||
|
dockerfile: Dockerfile
|
||||||
|
user: "${UID:-1000}:${GID:-1000}"
|
||||||
ports:
|
ports:
|
||||||
- "8000:8000"
|
- "8000:8000"
|
||||||
volumes:
|
volumes:
|
||||||
- ./params:/app/params:ro
|
- ./params:/app/params:ro
|
||||||
- ./checkpoints:/app/checkpoints
|
|
||||||
command: python -m scripts.tools.server --port 8000 --device cpu
|
command: python -m scripts.tools.server --port 8000 --device cpu
|
||||||
healthcheck:
|
healthcheck:
|
||||||
test: ["CMD", "curl", "-f", "http://localhost:8000/health"]
|
test: ["CMD", "curl", "-f", "http://localhost:8000/health"]
|
||||||
|
|||||||
@@ -11,7 +11,6 @@ PARAMETER_ROOT = Path(PROJECT_ROOT, "params")
|
|||||||
|
|
||||||
|
|
||||||
def generate_text():
|
def generate_text():
|
||||||
# Load model from pretrained
|
|
||||||
model = AutoModel.from_pretrained(PARAMETER_ROOT)
|
model = AutoModel.from_pretrained(PARAMETER_ROOT)
|
||||||
tokenizer = AutoTokenizer.from_pretrained(PARAMETER_ROOT)
|
tokenizer = AutoTokenizer.from_pretrained(PARAMETER_ROOT)
|
||||||
model.to(device="cuda", dtype=torch.bfloat16)
|
model.to(device="cuda", dtype=torch.bfloat16)
|
||||||
@@ -22,16 +21,15 @@ def generate_text():
|
|||||||
model=model,
|
model=model,
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
)
|
)
|
||||||
response = engine.generate(
|
for token in engine.generate(
|
||||||
prompt=query,
|
prompt=query,
|
||||||
stream=False,
|
stream=True,
|
||||||
max_tokens=2048,
|
max_tokens=2048,
|
||||||
temperature=0.8,
|
temperature=0.8,
|
||||||
top_p=0.95,
|
top_p=0.95,
|
||||||
top_k=50,
|
top_k=50,
|
||||||
)
|
):
|
||||||
|
print(token, end="", flush=True)
|
||||||
print(response)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
@@ -24,12 +24,23 @@ def batch_generate():
|
|||||||
"请问什么是显卡",
|
"请问什么是显卡",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
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(
|
engine = InferenceEngine(
|
||||||
model=model,
|
model=model,
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
)
|
)
|
||||||
responses = engine.generate(
|
responses = engine.generate(
|
||||||
prompt=inputs,
|
prompt=prompts,
|
||||||
stream=False,
|
stream=False,
|
||||||
max_tokens=2048,
|
max_tokens=2048,
|
||||||
temperature=0.8,
|
temperature=0.8,
|
||||||
|
|||||||
+8
-1
@@ -16,6 +16,7 @@ NC='\033[0m' # No Color
|
|||||||
IMAGE_NAME="astrai"
|
IMAGE_NAME="astrai"
|
||||||
IMAGE_TAG="latest"
|
IMAGE_TAG="latest"
|
||||||
REGISTRY=""
|
REGISTRY=""
|
||||||
|
CONTAINER_ID=""
|
||||||
|
|
||||||
# Print colored messages
|
# Print colored messages
|
||||||
print_info() {
|
print_info() {
|
||||||
@@ -175,6 +176,10 @@ main() {
|
|||||||
PORT="$2"
|
PORT="$2"
|
||||||
shift 2
|
shift 2
|
||||||
;;
|
;;
|
||||||
|
--container)
|
||||||
|
CONTAINER_ID="$2"
|
||||||
|
shift 2
|
||||||
|
;;
|
||||||
--gpu)
|
--gpu)
|
||||||
GPU=true
|
GPU=true
|
||||||
shift
|
shift
|
||||||
@@ -197,6 +202,7 @@ main() {
|
|||||||
echo " --dockerfile FILE Dockerfile path (default: Dockerfile)"
|
echo " --dockerfile FILE Dockerfile path (default: Dockerfile)"
|
||||||
echo " --context PATH Build context (default: .)"
|
echo " --context PATH Build context (default: .)"
|
||||||
echo " --port PORT Port for run (default: 8000)"
|
echo " --port PORT Port for run (default: 8000)"
|
||||||
|
echo " --container ID Container ID for logs"
|
||||||
echo " --gpu Enable GPU support"
|
echo " --gpu Enable GPU support"
|
||||||
echo " --help Show this help message"
|
echo " --help Show this help message"
|
||||||
echo ""
|
echo ""
|
||||||
@@ -205,6 +211,7 @@ main() {
|
|||||||
echo " $0 build --tag v1.0.0"
|
echo " $0 build --tag v1.0.0"
|
||||||
echo " $0 run --port 8080"
|
echo " $0 run --port 8080"
|
||||||
echo " $0 run --gpu"
|
echo " $0 run --gpu"
|
||||||
|
echo " $0 logs --container abc123"
|
||||||
echo " $0 push --registry ghcr.io/username"
|
echo " $0 push --registry ghcr.io/username"
|
||||||
exit 0
|
exit 0
|
||||||
;;
|
;;
|
||||||
@@ -237,7 +244,7 @@ main() {
|
|||||||
show_info
|
show_info
|
||||||
;;
|
;;
|
||||||
logs)
|
logs)
|
||||||
show_logs "$2"
|
show_logs "$CONTAINER_ID"
|
||||||
;;
|
;;
|
||||||
"")
|
"")
|
||||||
print_error "No command specified. Use --help for usage"
|
print_error "No command specified. Use --help for usage"
|
||||||
|
|||||||
+27
-18
@@ -1,13 +1,13 @@
|
|||||||
"""Benchmark Transformer with PagedCache (replaces old persistent_key_values)."""
|
"""Benchmark AutoRegressiveLM with KVCache"""
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, Dict
|
from typing import Any, Dict
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
|
||||||
|
|
||||||
from astrai.inference.cache import PagedCache
|
from astrai.config import AutoRegressiveLMConfig
|
||||||
from astrai.model.transformer import ModelConfig, Transformer
|
from astrai.inference import KVCache
|
||||||
|
from astrai.model.transformer import AutoRegressiveLM
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -21,7 +21,7 @@ class BenchmarkResult:
|
|||||||
class GenerationBenchmark:
|
class GenerationBenchmark:
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: ModelConfig,
|
config: AutoRegressiveLMConfig,
|
||||||
device: str = "cuda",
|
device: str = "cuda",
|
||||||
dtype: torch.dtype = torch.bfloat16,
|
dtype: torch.dtype = torch.bfloat16,
|
||||||
page_size: int = 128,
|
page_size: int = 128,
|
||||||
@@ -29,11 +29,11 @@ class GenerationBenchmark:
|
|||||||
self.config = config
|
self.config = config
|
||||||
self.device = device
|
self.device = device
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
self.model = Transformer(config).to(device=device, dtype=dtype)
|
self.model = AutoRegressiveLM(config).to(device=device, dtype=dtype)
|
||||||
self.model.eval()
|
self.model.eval()
|
||||||
head_dim = config.dim // config.n_heads
|
head_dim = config.dim // config.n_heads
|
||||||
n_pages = (config.max_len * 4 + page_size - 1) // page_size
|
n_pages = (config.max_len * 4 + page_size - 1) // page_size
|
||||||
self._page_cache = PagedCache(
|
self._page_cache = KVCache(
|
||||||
config.n_layers,
|
config.n_layers,
|
||||||
n_pages,
|
n_pages,
|
||||||
page_size,
|
page_size,
|
||||||
@@ -60,9 +60,6 @@ class GenerationBenchmark:
|
|||||||
)
|
)
|
||||||
return prompt_ids, gen_ids
|
return prompt_ids, gen_ids
|
||||||
|
|
||||||
def _make_mask(self, batch_size: int, seq_len: int) -> Tensor:
|
|
||||||
return torch.ones(batch_size, seq_len, dtype=torch.bool, device=self.device)
|
|
||||||
|
|
||||||
@torch.inference_mode()
|
@torch.inference_mode()
|
||||||
def run_prefill_benchmark(
|
def run_prefill_benchmark(
|
||||||
self,
|
self,
|
||||||
@@ -133,7 +130,12 @@ class GenerationBenchmark:
|
|||||||
)
|
)
|
||||||
|
|
||||||
n_pages = (prompt_length + gen_length + page_size - 1) // page_size
|
n_pages = (prompt_length + gen_length + page_size - 1) // page_size
|
||||||
pages = self._page_cache.alloc_n(n_pages * batch_size)
|
total = n_pages * batch_size
|
||||||
|
pages = []
|
||||||
|
for _ in range(total):
|
||||||
|
p = self._page_cache._pool.alloc()
|
||||||
|
assert p >= 0, "OOM"
|
||||||
|
pages.append(p)
|
||||||
page_table = torch.tensor(
|
page_table = torch.tensor(
|
||||||
[pages[i * n_pages : (i + 1) * n_pages] for i in range(batch_size)],
|
[pages[i * n_pages : (i + 1) * n_pages] for i in range(batch_size)],
|
||||||
dtype=torch.long,
|
dtype=torch.long,
|
||||||
@@ -144,8 +146,11 @@ class GenerationBenchmark:
|
|||||||
_ = self.model(
|
_ = self.model(
|
||||||
prompt_ids,
|
prompt_ids,
|
||||||
paged_cache=cv,
|
paged_cache=cv,
|
||||||
start_pos=0,
|
position_ids=torch.arange(
|
||||||
input_mask=self._make_mask(batch_size, prompt_length),
|
prompt_length, dtype=torch.long, device=self.device
|
||||||
|
)
|
||||||
|
.unsqueeze(0)
|
||||||
|
.expand(batch_size, -1),
|
||||||
)
|
)
|
||||||
|
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
@@ -161,8 +166,12 @@ class GenerationBenchmark:
|
|||||||
_ = self.model(
|
_ = self.model(
|
||||||
input_token,
|
input_token,
|
||||||
paged_cache=cv,
|
paged_cache=cv,
|
||||||
start_pos=current_pos,
|
position_ids=torch.full(
|
||||||
input_mask=self._make_mask(batch_size, 1),
|
(batch_size, 1),
|
||||||
|
current_pos,
|
||||||
|
dtype=torch.long,
|
||||||
|
device=self.device,
|
||||||
|
),
|
||||||
)
|
)
|
||||||
current_pos += 1
|
current_pos += 1
|
||||||
end.record()
|
end.record()
|
||||||
@@ -172,7 +181,7 @@ class GenerationBenchmark:
|
|||||||
total_time += trial_time
|
total_time += trial_time
|
||||||
|
|
||||||
for idx in pages:
|
for idx in pages:
|
||||||
self._page_cache.free(idx)
|
self._page_cache._pool.free(idx)
|
||||||
|
|
||||||
print(
|
print(
|
||||||
f" Trial {trial + 1}/{num_trials}: {gen_length} tokens in {trial_time:.3f}s "
|
f" Trial {trial + 1}/{num_trials}: {gen_length} tokens in {trial_time:.3f}s "
|
||||||
@@ -207,7 +216,7 @@ def print_benchmark_result(result: BenchmarkResult):
|
|||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
config = ModelConfig(
|
config = AutoRegressiveLMConfig(
|
||||||
vocab_size=10000,
|
vocab_size=10000,
|
||||||
dim=1536,
|
dim=1536,
|
||||||
n_heads=24,
|
n_heads=24,
|
||||||
@@ -221,7 +230,7 @@ if __name__ == "__main__":
|
|||||||
benchmark = GenerationBenchmark(config)
|
benchmark = GenerationBenchmark(config)
|
||||||
|
|
||||||
print("=" * 80)
|
print("=" * 80)
|
||||||
print("Running Transformer Generation Benchmark (PagedCache)")
|
print("Running AutoRegressiveLM Generation Benchmark (KVCache)")
|
||||||
print("=" * 80)
|
print("=" * 80)
|
||||||
|
|
||||||
prefill_result = benchmark.run_prefill_benchmark(
|
prefill_result = benchmark.run_prefill_benchmark(
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ from astrai.tokenize import AutoTokenizer
|
|||||||
|
|
||||||
|
|
||||||
def processor(
|
def processor(
|
||||||
model_dir: str,
|
param_path: str,
|
||||||
input_json_file: str,
|
input_json_file: str,
|
||||||
output_json_file: str,
|
output_json_file: str,
|
||||||
temperature: float,
|
temperature: float,
|
||||||
@@ -18,14 +18,17 @@ def processor(
|
|||||||
question_key: str,
|
question_key: str,
|
||||||
response_key: str,
|
response_key: str,
|
||||||
max_tokens: int,
|
max_tokens: int,
|
||||||
|
batch_size: int,
|
||||||
):
|
):
|
||||||
# Load model and tokenizer
|
# Load model and tokenizer
|
||||||
model = AutoModel.from_pretrained(model_dir)
|
model = AutoModel.from_pretrained(param_path)
|
||||||
tokenizer = AutoTokenizer.from_pretrained(model_dir)
|
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||||
model.to(device="cuda", dtype=torch.bfloat16)
|
model.to(device="cuda", dtype=torch.bfloat16)
|
||||||
|
|
||||||
# Create inference engine
|
# Create inference engine
|
||||||
engine = InferenceEngine(model=model, tokenizer=tokenizer)
|
engine = InferenceEngine(
|
||||||
|
model=model, tokenizer=tokenizer, max_batch_size=batch_size
|
||||||
|
)
|
||||||
|
|
||||||
with open(input_json_file, "r", encoding="utf-8") as f:
|
with open(input_json_file, "r", encoding="utf-8") as f:
|
||||||
input_data = [json.loads(line) for line in f]
|
input_data = [json.loads(line) for line in f]
|
||||||
@@ -72,7 +75,7 @@ if __name__ == "__main__":
|
|||||||
parser = argparse.ArgumentParser(description="Run generate with a Khaosz model.")
|
parser = argparse.ArgumentParser(description="Run generate with a Khaosz model.")
|
||||||
|
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--model_dir", type=str, required=True, help="Path to the model directory."
|
"--param_path", type=str, required=True, help="Path to the model directory."
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--input_json_file",
|
"--input_json_file",
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ from pathlib import Path
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.inference.server import run_server
|
from astrai.inference import run_server
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
|
|||||||
+68
-37
@@ -8,16 +8,16 @@ import torch.nn as nn
|
|||||||
import torch.optim as optim
|
import torch.optim as optim
|
||||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||||
|
|
||||||
from astrai.config import ModelConfig, TrainConfig
|
from astrai.config import AutoRegressiveLMConfig, TrainConfig
|
||||||
from astrai.dataset import DatasetFactory
|
from astrai.dataset import DatasetFactory
|
||||||
from astrai.model import Transformer
|
from astrai.model import AutoRegressiveLM
|
||||||
from astrai.parallel import get_rank
|
from astrai.parallel import get_rank
|
||||||
from astrai.trainer import SchedulerFactory, Trainer
|
from astrai.trainer import SchedulerFactory, Trainer
|
||||||
|
|
||||||
|
|
||||||
def parse_args() -> argparse.Namespace:
|
def parse_args() -> argparse.Namespace:
|
||||||
|
|
||||||
parser = argparse.ArgumentParser(description="Train the Transformer model.")
|
parser = argparse.ArgumentParser(description="Train the AutoRegressiveLM model.")
|
||||||
|
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--train_type",
|
"--train_type",
|
||||||
@@ -42,18 +42,20 @@ def parse_args() -> argparse.Namespace:
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--n_epoch", type=int, default=1, help="Number of epochs to train."
|
"--n_epoch", type=int, default=1, help="Number of epochs to train."
|
||||||
)
|
)
|
||||||
parser.add_argument("--group_size", type=int, default=4, help="GRPO group size.")
|
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--accumulation_steps",
|
"--batch_per_device", type=int, default=1, help="Batch size per GPU."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--grad_accum_steps",
|
||||||
type=int,
|
type=int,
|
||||||
default=1,
|
default=1,
|
||||||
help="Number of iterations between each optimizer step.",
|
help="Number of iterations between each optimizer step.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--warmup_steps",
|
"--warmup_ratio",
|
||||||
type=int,
|
type=float,
|
||||||
default=1000,
|
default=0.05,
|
||||||
help="Number of iters between warnings.",
|
help="Fraction of total steps used for LR warmup.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--max_lr", type=float, default=3e-4, help="Max learning rate for training."
|
"--max_lr", type=float, default=3e-4, help="Max learning rate for training."
|
||||||
@@ -68,13 +70,13 @@ def parse_args() -> argparse.Namespace:
|
|||||||
"--adamw_beta1",
|
"--adamw_beta1",
|
||||||
type=float,
|
type=float,
|
||||||
default=0.9,
|
default=0.9,
|
||||||
help="Beta values for AdamW optimizer.",
|
help="Beta1 for AdamW optimizer.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--adamw_beta2",
|
"--adamw_beta2",
|
||||||
type=float,
|
type=float,
|
||||||
default=0.95,
|
default=0.95,
|
||||||
help="Beta values for AdamW optimizer.",
|
help="Beta2 for AdamW optimizer.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--adamw_weight_decay",
|
"--adamw_weight_decay",
|
||||||
@@ -98,27 +100,23 @@ def parse_args() -> argparse.Namespace:
|
|||||||
"--window_size",
|
"--window_size",
|
||||||
type=int,
|
type=int,
|
||||||
default=None,
|
default=None,
|
||||||
help="the max length of the input sequence.",
|
help="Max length of the input sequence.",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--stride", type=int, default=None, help="the step size of the input sequence."
|
"--stride", type=int, default=None, help="Step size of the input sequence."
|
||||||
)
|
)
|
||||||
parser.add_argument("--dpo_beta", type=float, default=0.1, help="DPO beta value.")
|
parser.add_argument("--dpo_beta", type=float, default=0.1, help="DPO beta value.")
|
||||||
parser.add_argument("--group_size", type=int, default=4, help="GRPO group size.")
|
parser.add_argument("--group_size", type=int, default=4, help="GRPO group size.")
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--on_policy",
|
"--grpo_clip_eps", type=float, default=0.2, help="GRPO clipping epsilon."
|
||||||
action="store_true",
|
|
||||||
default=False,
|
|
||||||
help="Enable on-policy GRPO mode.",
|
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--grpo_kl_coef", type=float, default=0.01, help="GRPO KL penalty coefficient."
|
"--grpo_kl_coef", type=float, default=0.01, help="GRPO KL penalty coefficient."
|
||||||
)
|
)
|
||||||
parser.add_argument("--group_size", type=int, default=4, help="GRPO group size.")
|
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--label_smoothing",
|
"--label_smoothing",
|
||||||
type=float,
|
type=float,
|
||||||
default=0.1,
|
default=0.05,
|
||||||
help="cross_entropy function label smoothing parameter",
|
help="cross_entropy function label smoothing parameter",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -134,7 +132,6 @@ def parse_args() -> argparse.Namespace:
|
|||||||
default="checkpoint",
|
default="checkpoint",
|
||||||
help="Directory to save checkpoints.",
|
help="Directory to save checkpoints.",
|
||||||
)
|
)
|
||||||
parser.add_argument("--group_size", type=int, default=4, help="GRPO group size.")
|
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--grpo_sync_interval",
|
"--grpo_sync_interval",
|
||||||
type=int,
|
type=int,
|
||||||
@@ -152,6 +149,13 @@ def parse_args() -> argparse.Namespace:
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--device_type", type=str, default="cuda", help="Device type to use."
|
"--device_type", type=str, default="cuda", help="Device type to use."
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--start_method",
|
||||||
|
type=str,
|
||||||
|
default="spawn",
|
||||||
|
choices=["spawn", "fork", "forkserver"],
|
||||||
|
help="Multiprocessing start method.",
|
||||||
|
)
|
||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
@@ -160,18 +164,20 @@ def parse_args() -> argparse.Namespace:
|
|||||||
|
|
||||||
def ddp_wrap(model: nn.Module):
|
def ddp_wrap(model: nn.Module):
|
||||||
local_rank = get_rank()
|
local_rank = get_rank()
|
||||||
model = model.to(device=f"cuda:{local_rank}", dtype=torch.bfloat16)
|
|
||||||
ddp_model = DDP(
|
ddp_model = DDP(
|
||||||
model,
|
model,
|
||||||
device_ids=[local_rank],
|
device_ids=[local_rank],
|
||||||
output_device=local_rank,
|
output_device=local_rank,
|
||||||
|
static_graph=True,
|
||||||
find_unused_parameters=False,
|
find_unused_parameters=False,
|
||||||
|
gradient_as_bucket_view=True,
|
||||||
|
broadcast_buffers=False,
|
||||||
)
|
)
|
||||||
return ddp_model
|
return ddp_model
|
||||||
|
|
||||||
|
|
||||||
def create_optimizer(model: nn.Module, **kwargs) -> optim.Optimizer:
|
def create_optimizer(model: nn.Module, **kwargs) -> optim.Optimizer:
|
||||||
return optim.AdamW(model.parameters(), **kwargs)
|
return optim.AdamW(model.parameters(), fused=True, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
def create_scheduler(
|
def create_scheduler(
|
||||||
@@ -181,7 +187,26 @@ def create_scheduler(
|
|||||||
|
|
||||||
|
|
||||||
def prepare_checkpoint(model: nn.Module) -> dict:
|
def prepare_checkpoint(model: nn.Module) -> dict:
|
||||||
return model.module.state_dict()
|
if isinstance(model, DDP):
|
||||||
|
return model.module.state_dict()
|
||||||
|
return model.state_dict()
|
||||||
|
|
||||||
|
|
||||||
|
def compute_total_steps(
|
||||||
|
dataset_len: int,
|
||||||
|
n_epoch: int,
|
||||||
|
batch_per_device: int,
|
||||||
|
nprocs: int,
|
||||||
|
grad_accum_steps: int,
|
||||||
|
) -> int:
|
||||||
|
|
||||||
|
def ceil_div(a: int, b: int) -> int:
|
||||||
|
return (a + b - 1) // b
|
||||||
|
|
||||||
|
samples_per_replica = ceil_div(dataset_len, nprocs)
|
||||||
|
batches_per_replica = ceil_div(samples_per_replica, batch_per_device)
|
||||||
|
total_steps = (batches_per_replica // grad_accum_steps) * n_epoch
|
||||||
|
return total_steps
|
||||||
|
|
||||||
|
|
||||||
def train(
|
def train(
|
||||||
@@ -190,11 +215,11 @@ def train(
|
|||||||
data_root_path: str,
|
data_root_path: str,
|
||||||
max_lr: float,
|
max_lr: float,
|
||||||
n_epoch: int,
|
n_epoch: int,
|
||||||
batch_size: int,
|
batch_per_device: int,
|
||||||
start_epoch: int,
|
start_epoch: int,
|
||||||
start_batch: int,
|
start_batch: int,
|
||||||
accumulation_steps: int,
|
grad_accum_steps: int,
|
||||||
warmup_steps: int,
|
warmup_ratio: float,
|
||||||
ckpt_interval: int,
|
ckpt_interval: int,
|
||||||
ckpt_dir: str,
|
ckpt_dir: str,
|
||||||
dpo_beta: float,
|
dpo_beta: float,
|
||||||
@@ -214,21 +239,20 @@ def train(
|
|||||||
stride: int,
|
stride: int,
|
||||||
nprocs: int,
|
nprocs: int,
|
||||||
device_type: str,
|
device_type: str,
|
||||||
|
start_method: str,
|
||||||
):
|
):
|
||||||
assert train_type in ["seq", "sft", "dpo", "grpo"]
|
assert train_type in ["seq", "sft", "dpo", "grpo"]
|
||||||
assert os.path.exists(param_path)
|
assert os.path.exists(param_path)
|
||||||
|
|
||||||
# Load config
|
# Load config
|
||||||
config = ModelConfig()
|
|
||||||
config_path = os.path.join(param_path, "config.json")
|
config_path = os.path.join(param_path, "config.json")
|
||||||
if os.path.exists(config_path):
|
config = AutoRegressiveLMConfig.from_file(config_path)
|
||||||
config.load(config_path)
|
|
||||||
|
|
||||||
if window_size is None:
|
if window_size is None:
|
||||||
window_size = config.max_len
|
window_size = config.max_len
|
||||||
|
|
||||||
# Create bare Transformer (for training, no tokenizer needed)
|
# Create bare AutoRegressiveLM (for training, no tokenizer needed)
|
||||||
model = Transformer(config)
|
model = AutoRegressiveLM(config)
|
||||||
|
|
||||||
# Load weights if available
|
# Load weights if available
|
||||||
weights_path = os.path.join(param_path, "model.safetensors")
|
weights_path = os.path.join(param_path, "model.safetensors")
|
||||||
@@ -236,8 +260,10 @@ def train(
|
|||||||
state_dict = st.load_file(weights_path)
|
state_dict = st.load_file(weights_path)
|
||||||
model.load_state_dict(state_dict, strict=False)
|
model.load_state_dict(state_dict, strict=False)
|
||||||
|
|
||||||
|
model = model.to(dtype=torch.bfloat16)
|
||||||
|
|
||||||
strategy_kwargs = {
|
strategy_kwargs = {
|
||||||
"dpo_beta": dpo_beta,
|
"beta": dpo_beta,
|
||||||
"label_smoothing": label_smoothing,
|
"label_smoothing": label_smoothing,
|
||||||
"clip_eps": grpo_clip_eps,
|
"clip_eps": grpo_clip_eps,
|
||||||
"kl_coef": grpo_kl_coef,
|
"kl_coef": grpo_kl_coef,
|
||||||
@@ -261,13 +287,17 @@ def train(
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
total_steps = len(dataset) * n_epoch // (batch_size * nprocs)
|
total_steps = compute_total_steps(
|
||||||
|
len(dataset), n_epoch, batch_per_device, nprocs, grad_accum_steps
|
||||||
|
)
|
||||||
|
warmup_steps = int(warmup_ratio * total_steps)
|
||||||
|
|
||||||
scheduler_fn = partial(
|
scheduler_fn = partial(
|
||||||
create_scheduler,
|
create_scheduler,
|
||||||
**{
|
**{
|
||||||
"schedule_type": "cosine",
|
"schedule_type": "cosine",
|
||||||
"warmup_steps": warmup_steps,
|
"warmup_steps": min(warmup_steps, total_steps),
|
||||||
"lr_decay_steps": total_steps - warmup_steps,
|
"lr_decay_steps": total_steps - min(warmup_steps, total_steps),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -279,11 +309,11 @@ def train(
|
|||||||
scheduler_fn=scheduler_fn,
|
scheduler_fn=scheduler_fn,
|
||||||
ckpt_dir=ckpt_dir,
|
ckpt_dir=ckpt_dir,
|
||||||
n_epoch=n_epoch,
|
n_epoch=n_epoch,
|
||||||
batch_size=batch_size,
|
batch_per_device=batch_per_device,
|
||||||
start_epoch=start_epoch,
|
start_epoch=start_epoch,
|
||||||
start_batch=start_batch,
|
start_batch=start_batch,
|
||||||
ckpt_interval=ckpt_interval,
|
ckpt_interval=ckpt_interval,
|
||||||
accumulation_steps=accumulation_steps,
|
grad_accum_steps=grad_accum_steps,
|
||||||
max_grad_norm=max_grad_norm,
|
max_grad_norm=max_grad_norm,
|
||||||
random_seed=random_seed,
|
random_seed=random_seed,
|
||||||
num_workers=num_workers,
|
num_workers=num_workers,
|
||||||
@@ -292,6 +322,7 @@ def train(
|
|||||||
parallel_wrapper=ddp_wrap,
|
parallel_wrapper=ddp_wrap,
|
||||||
state_dict_fn=prepare_checkpoint,
|
state_dict_fn=prepare_checkpoint,
|
||||||
device_type=device_type,
|
device_type=device_type,
|
||||||
|
start_method=start_method,
|
||||||
extra_kwargs=strategy_kwargs,
|
extra_kwargs=strategy_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
+59
-77
@@ -3,18 +3,22 @@ import os
|
|||||||
import shutil
|
import shutil
|
||||||
import tempfile
|
import tempfile
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
import pytest
|
import pytest
|
||||||
import safetensors.torch as st
|
|
||||||
import torch
|
import torch
|
||||||
from tokenizers import Tokenizer, models, pre_tokenizers, trainers
|
from tokenizers import Tokenizer, models, pre_tokenizers, trainers
|
||||||
from torch.utils.data import Dataset
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
from astrai.config.model_config import ModelConfig
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
from astrai.model.transformer import Transformer
|
from astrai.model.transformer import AutoRegressiveLM
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
|
||||||
|
def pytest_configure(config):
|
||||||
|
config.addinivalue_line("markers", "slow: marks tests as slow")
|
||||||
|
config.addinivalue_line("markers", "integration: integration tests")
|
||||||
|
config.addinivalue_line("markers", "unit: fast unit tests")
|
||||||
|
|
||||||
|
|
||||||
def create_test_tokenizer(vocab_size: int = 1000) -> AutoTokenizer:
|
def create_test_tokenizer(vocab_size: int = 1000) -> AutoTokenizer:
|
||||||
"""Create a simple tokenizer for testing purposes."""
|
"""Create a simple tokenizer for testing purposes."""
|
||||||
tokenizer = Tokenizer(models.BPE())
|
tokenizer = Tokenizer(models.BPE())
|
||||||
@@ -22,7 +26,6 @@ def create_test_tokenizer(vocab_size: int = 1000) -> AutoTokenizer:
|
|||||||
trainer = trainers.BpeTrainer(
|
trainer = trainers.BpeTrainer(
|
||||||
vocab_size=vocab_size, min_frequency=1, special_tokens=["<unk>", "<pad>"]
|
vocab_size=vocab_size, min_frequency=1, special_tokens=["<unk>", "<pad>"]
|
||||||
)
|
)
|
||||||
# Train on empty iterator with single character
|
|
||||||
tokenizer.train_from_iterator([chr(i) for i in range(256)], trainer)
|
tokenizer.train_from_iterator([chr(i) for i in range(256)], trainer)
|
||||||
auto_tokenizer = AutoTokenizer()
|
auto_tokenizer = AutoTokenizer()
|
||||||
auto_tokenizer._tokenizer = tokenizer
|
auto_tokenizer._tokenizer = tokenizer
|
||||||
@@ -34,7 +37,7 @@ class RandomDataset(Dataset):
|
|||||||
"""Random dataset for testing purposes."""
|
"""Random dataset for testing purposes."""
|
||||||
|
|
||||||
def __init__(self, length=None, max_length=64, vocab_size=1000):
|
def __init__(self, length=None, max_length=64, vocab_size=1000):
|
||||||
self.length = length or int(np.random.randint(100, 200))
|
self.length = length or int(torch.randint(100, 200, (1,)).item())
|
||||||
self.max_length = max_length
|
self.max_length = max_length
|
||||||
self.vocab_size = vocab_size
|
self.vocab_size = vocab_size
|
||||||
|
|
||||||
@@ -52,7 +55,7 @@ class MultiTurnDataset(Dataset):
|
|||||||
"""Multi-turn dataset with loss mask for SFT training tests."""
|
"""Multi-turn dataset with loss mask for SFT training tests."""
|
||||||
|
|
||||||
def __init__(self, length=None, max_length=64, vocab_size=1000):
|
def __init__(self, length=None, max_length=64, vocab_size=1000):
|
||||||
self.length = length or int(np.random.randint(100, 200))
|
self.length = length or int(torch.randint(100, 200, (1,)).item())
|
||||||
self.max_length = max_length
|
self.max_length = max_length
|
||||||
self.vocab_size = vocab_size
|
self.vocab_size = vocab_size
|
||||||
|
|
||||||
@@ -93,46 +96,65 @@ class EarlyStoppingDataset(Dataset):
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture(scope="session")
|
||||||
def base_test_env(request: pytest.FixtureRequest):
|
def test_tokenizer():
|
||||||
"""Create base test environment with randomly configured model and tokenizer"""
|
"""Session-scoped tokenizer, created once for the entire test run."""
|
||||||
func_name = request.function.__name__
|
return create_test_tokenizer()
|
||||||
test_dir = tempfile.mkdtemp(prefix=f"{func_name}_")
|
|
||||||
config_path = os.path.join(test_dir, "config.json")
|
|
||||||
|
|
||||||
n_dim_choices = [8, 16, 32]
|
|
||||||
n_head_choices = [2, 4]
|
|
||||||
|
|
||||||
dim = int(np.random.choice(n_dim_choices))
|
@pytest.fixture(scope="session")
|
||||||
n_heads = int(np.random.choice(n_head_choices))
|
def test_model():
|
||||||
n_kv_heads = n_heads // 2
|
"""Session-scoped small AutoRegressiveLM model, created once."""
|
||||||
dim_ffn = dim * 2
|
config = AutoRegressiveLMConfig(
|
||||||
|
vocab_size=1000,
|
||||||
|
dim=8,
|
||||||
|
n_heads=2,
|
||||||
|
n_kv_heads=1,
|
||||||
|
dim_ffn=16,
|
||||||
|
max_len=64,
|
||||||
|
n_layers=2,
|
||||||
|
norm_eps=1e-5,
|
||||||
|
)
|
||||||
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
model = AutoRegressiveLM(config).to(device=device)
|
||||||
|
|
||||||
config = {
|
return {
|
||||||
"vocab_size": 1000,
|
"model": model,
|
||||||
"dim": dim,
|
"device": device,
|
||||||
"n_heads": n_heads,
|
"config": config,
|
||||||
"n_kv_heads": n_kv_heads,
|
|
||||||
"dim_ffn": dim_ffn,
|
|
||||||
"max_len": 1024,
|
|
||||||
"n_layers": 4,
|
|
||||||
"norm_eps": 1e-5,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def base_test_env(test_model, test_tokenizer):
|
||||||
|
"""Function-scoped test environment with isolated temp directory.
|
||||||
|
|
||||||
|
Composes session-scoped model and tokenizer with a per-test temp dir.
|
||||||
|
"""
|
||||||
|
test_dir = tempfile.mkdtemp()
|
||||||
|
config_path = os.path.join(test_dir, "config.json")
|
||||||
with open(config_path, "w") as f:
|
with open(config_path, "w") as f:
|
||||||
json.dump(config, f)
|
json.dump(
|
||||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
{
|
||||||
transformer_config = ModelConfig().load(config_path)
|
"vocab_size": 1000,
|
||||||
model = Transformer(transformer_config).to(device=device)
|
"dim": 8,
|
||||||
tokenizer = create_test_tokenizer()
|
"n_heads": 2,
|
||||||
|
"n_kv_heads": 1,
|
||||||
|
"dim_ffn": 16,
|
||||||
|
"max_len": 64,
|
||||||
|
"n_layers": 2,
|
||||||
|
"norm_eps": 1e-5,
|
||||||
|
},
|
||||||
|
f,
|
||||||
|
)
|
||||||
|
|
||||||
yield {
|
yield {
|
||||||
"device": device,
|
"device": test_model["device"],
|
||||||
"test_dir": str(test_dir),
|
"test_dir": str(test_dir),
|
||||||
"config_path": config_path,
|
"config_path": config_path,
|
||||||
"transformer_config": transformer_config,
|
"transformer_config": test_model["config"],
|
||||||
"model": model,
|
"model": test_model["model"],
|
||||||
"tokenizer": tokenizer,
|
"tokenizer": test_tokenizer,
|
||||||
}
|
}
|
||||||
|
|
||||||
shutil.rmtree(test_dir)
|
shutil.rmtree(test_dir)
|
||||||
@@ -154,43 +176,3 @@ def multi_turn_dataset():
|
|||||||
def early_stopping_dataset():
|
def early_stopping_dataset():
|
||||||
dataset = EarlyStoppingDataset()
|
dataset = EarlyStoppingDataset()
|
||||||
yield dataset
|
yield dataset
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
def test_env(request: pytest.FixtureRequest):
|
|
||||||
"""Create a test environment with saved model and tokenizer files."""
|
|
||||||
|
|
||||||
func_name = request.function.__name__
|
|
||||||
test_dir = tempfile.mkdtemp(prefix=f"{func_name}_")
|
|
||||||
config_path = os.path.join(test_dir, "config.json")
|
|
||||||
tokenizer_path = os.path.join(test_dir, "tokenizer.json")
|
|
||||||
model_path = os.path.join(test_dir, "model.safetensors")
|
|
||||||
|
|
||||||
config = {
|
|
||||||
"vocab_size": 1000,
|
|
||||||
"dim": 128,
|
|
||||||
"n_heads": 4,
|
|
||||||
"n_kv_heads": 2,
|
|
||||||
"dim_ffn": 256,
|
|
||||||
"max_len": 64,
|
|
||||||
"n_layers": 2,
|
|
||||||
"norm_eps": 1e-5,
|
|
||||||
}
|
|
||||||
with open(config_path, "w") as f:
|
|
||||||
json.dump(config, f)
|
|
||||||
|
|
||||||
tokenizer = create_test_tokenizer(vocab_size=config["vocab_size"])
|
|
||||||
tokenizer.save(tokenizer_path)
|
|
||||||
|
|
||||||
transformer_config = ModelConfig().load(config_path)
|
|
||||||
model = Transformer(transformer_config)
|
|
||||||
st.save_file(model.state_dict(), model_path)
|
|
||||||
|
|
||||||
yield {
|
|
||||||
"test_dir": test_dir,
|
|
||||||
"model": model,
|
|
||||||
"tokenizer": tokenizer,
|
|
||||||
"transformer_config": transformer_config,
|
|
||||||
}
|
|
||||||
|
|
||||||
shutil.rmtree(test_dir)
|
|
||||||
|
|||||||
@@ -35,6 +35,33 @@ def test_single_process():
|
|||||||
assert loaded_checkpoint.iteration == 30
|
assert loaded_checkpoint.iteration == 30
|
||||||
|
|
||||||
|
|
||||||
|
def test_checkpoint_with_extra():
|
||||||
|
"""Verify extra keys are saved as individual .pt files and loaded back."""
|
||||||
|
model = torch.nn.Linear(10, 5)
|
||||||
|
optimizer = AdamW(model.parameters(), lr=1e-3)
|
||||||
|
optimizer.step()
|
||||||
|
|
||||||
|
extra = {
|
||||||
|
"optimizer": optimizer.state_dict(),
|
||||||
|
"scheduler": {"last_epoch": 5},
|
||||||
|
}
|
||||||
|
checkpoint = Checkpoint(
|
||||||
|
state_dict=model.state_dict(), epoch=1, iteration=10, extra=extra
|
||||||
|
)
|
||||||
|
|
||||||
|
with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
|
checkpoint.save(tmpdir)
|
||||||
|
|
||||||
|
import os
|
||||||
|
|
||||||
|
assert os.path.exists(os.path.join(tmpdir, "optimizer.pt"))
|
||||||
|
assert os.path.exists(os.path.join(tmpdir, "scheduler.pt"))
|
||||||
|
|
||||||
|
loaded = Checkpoint.load(tmpdir)
|
||||||
|
assert loaded.extra["scheduler"]["last_epoch"] == 5
|
||||||
|
assert "state" in loaded.extra["optimizer"]
|
||||||
|
|
||||||
|
|
||||||
def simple_training():
|
def simple_training():
|
||||||
model = torch.nn.Linear(10, 5)
|
model = torch.nn.Linear(10, 5)
|
||||||
optimizer = AdamW(model.parameters(), lr=1e-3)
|
optimizer = AdamW(model.parameters(), lr=1e-3)
|
||||||
|
|||||||
+279
-4
@@ -1,8 +1,20 @@
|
|||||||
|
import json
|
||||||
|
import os
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.dataset.dataset import DatasetFactory
|
from astrai.dataset.dataset import DatasetFactory, SEQDataset
|
||||||
from astrai.serialization import save_h5
|
from astrai.dataset.storage import (
|
||||||
|
BaseSegmentFetcher,
|
||||||
|
H5Storage,
|
||||||
|
MultiSegmentFetcher,
|
||||||
|
StorageFactory,
|
||||||
|
detect_format,
|
||||||
|
load_json,
|
||||||
|
save_h5,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_dataset_loader_random_paths(base_test_env):
|
def test_dataset_loader_random_paths(base_test_env):
|
||||||
@@ -64,7 +76,7 @@ def test_dpo_strategy_with_random_data(base_test_env):
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert dpo_dataset is not None
|
assert dpo_dataset is not None
|
||||||
assert hasattr(dpo_dataset, "fetcher")
|
assert dpo_dataset.storage is not None
|
||||||
assert len(dpo_dataset) > 0
|
assert len(dpo_dataset) > 0
|
||||||
|
|
||||||
# Test that we can get DPO items without errors
|
# Test that we can get DPO items without errors
|
||||||
@@ -100,7 +112,7 @@ def test_sft_dataset_with_random_data(base_test_env):
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert sft_dataset is not None
|
assert sft_dataset is not None
|
||||||
assert hasattr(sft_dataset, "fetcher")
|
assert sft_dataset.storage is not None
|
||||||
assert len(sft_dataset) > 0
|
assert len(sft_dataset) > 0
|
||||||
|
|
||||||
# Test that we can get SFT items without errors
|
# Test that we can get SFT items without errors
|
||||||
@@ -143,3 +155,266 @@ def test_dataset_with_custom_stride(base_test_env):
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert len(dataset) > len(default_stride_dataset)
|
assert len(dataset) > len(default_stride_dataset)
|
||||||
|
|
||||||
|
|
||||||
|
# ============== JSON Storage Tests (raw text + tokenizer) ==============
|
||||||
|
|
||||||
|
|
||||||
|
def _make_tokenizer_fn(tokenizer):
|
||||||
|
"""Wrap tokenizer.encode() as a str -> List[int] callable."""
|
||||||
|
return lambda text: tokenizer.encode(text, add_special_tokens=False)
|
||||||
|
|
||||||
|
|
||||||
|
def test_seq_dataset_from_json_text(base_test_env):
|
||||||
|
"""Test loading SEQ dataset from raw-text JSON with tokenizer"""
|
||||||
|
tokenizer = base_test_env["tokenizer"]
|
||||||
|
tokenizer_fn = _make_tokenizer_fn(tokenizer)
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
data_dir = os.path.join(test_dir, "json_text")
|
||||||
|
os.makedirs(data_dir, exist_ok=True)
|
||||||
|
|
||||||
|
texts = [
|
||||||
|
"hello world this is a test sentence for tokenizer",
|
||||||
|
"another sentence with different words and tokens",
|
||||||
|
"machine learning is fascinating and powerful",
|
||||||
|
]
|
||||||
|
|
||||||
|
json_path = os.path.join(data_dir, "seq_data.json")
|
||||||
|
with open(json_path, "w", encoding="utf-8") as f:
|
||||||
|
json.dump({"sequence": texts}, f, ensure_ascii=False)
|
||||||
|
|
||||||
|
dataset = DatasetFactory.load(
|
||||||
|
train_type="seq",
|
||||||
|
load_path=data_dir,
|
||||||
|
window_size=16,
|
||||||
|
tokenizer=tokenizer_fn,
|
||||||
|
)
|
||||||
|
assert dataset is not None
|
||||||
|
assert len(dataset) > 0
|
||||||
|
assert dataset.count > 0
|
||||||
|
assert "sequence" in dataset.keys
|
||||||
|
|
||||||
|
item = dataset[0]
|
||||||
|
assert "input_ids" in item
|
||||||
|
assert "target_ids" in item
|
||||||
|
assert item["input_ids"].shape[0] == 16
|
||||||
|
|
||||||
|
|
||||||
|
def test_sft_dataset_from_json_text(base_test_env):
|
||||||
|
"""Test loading SFT dataset from raw-text JSON with tokenizer"""
|
||||||
|
tokenizer = base_test_env["tokenizer"]
|
||||||
|
tokenizer_fn = _make_tokenizer_fn(tokenizer)
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
data_dir = os.path.join(test_dir, "json_sft")
|
||||||
|
os.makedirs(data_dir, exist_ok=True)
|
||||||
|
|
||||||
|
texts = [
|
||||||
|
"user asks a question about the weather",
|
||||||
|
"assistant provides a helpful response to the user",
|
||||||
|
]
|
||||||
|
|
||||||
|
json_path = os.path.join(data_dir, "sft_data.json")
|
||||||
|
with open(json_path, "w", encoding="utf-8") as f:
|
||||||
|
json.dump(
|
||||||
|
{"sequence": texts, "loss_mask": texts},
|
||||||
|
f,
|
||||||
|
ensure_ascii=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
dataset = DatasetFactory.load(
|
||||||
|
train_type="sft",
|
||||||
|
load_path=data_dir,
|
||||||
|
window_size=16,
|
||||||
|
tokenizer=tokenizer_fn,
|
||||||
|
)
|
||||||
|
assert dataset is not None
|
||||||
|
assert len(dataset) > 0
|
||||||
|
|
||||||
|
item = dataset[0]
|
||||||
|
assert "loss_mask" in item
|
||||||
|
|
||||||
|
|
||||||
|
def test_json_storage_explicit_tokenizer(base_test_env):
|
||||||
|
"""Test explicit JSON storage with tokenizer"""
|
||||||
|
tokenizer = base_test_env["tokenizer"]
|
||||||
|
tokenizer_fn = _make_tokenizer_fn(tokenizer)
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
data_dir = os.path.join(test_dir, "json_explicit")
|
||||||
|
os.makedirs(data_dir, exist_ok=True)
|
||||||
|
|
||||||
|
texts = ["abcdefghijklmnopqrstuvwxyz" * 10]
|
||||||
|
|
||||||
|
json_path = os.path.join(data_dir, "data.json")
|
||||||
|
with open(json_path, "w", encoding="utf-8") as f:
|
||||||
|
json.dump({"sequence": texts}, f, ensure_ascii=False)
|
||||||
|
|
||||||
|
token_count = len(tokenizer_fn(texts[0]))
|
||||||
|
|
||||||
|
dataset = DatasetFactory.load(
|
||||||
|
train_type="seq",
|
||||||
|
load_path=data_dir,
|
||||||
|
window_size=32,
|
||||||
|
storage_type="json",
|
||||||
|
tokenizer=tokenizer_fn,
|
||||||
|
)
|
||||||
|
assert dataset is not None
|
||||||
|
assert len(dataset) > 0
|
||||||
|
assert dataset.count == token_count
|
||||||
|
|
||||||
|
|
||||||
|
def test_dataset_count_property(base_test_env):
|
||||||
|
"""Test the count property returns correct raw token count"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
|
||||||
|
seq_length = 200
|
||||||
|
dummy_data = {
|
||||||
|
"sequence": [torch.randint(0, 1000, (seq_length,), dtype=torch.int64)],
|
||||||
|
}
|
||||||
|
|
||||||
|
save_h5(test_dir, "count_test_data", dummy_data)
|
||||||
|
|
||||||
|
dataset = DatasetFactory.load(
|
||||||
|
train_type="seq",
|
||||||
|
load_path=test_dir,
|
||||||
|
window_size=64,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert dataset.count == seq_length
|
||||||
|
assert dataset.count > len(dataset) # raw tokens > windows
|
||||||
|
assert len(dataset) == (seq_length - 1 - 64) // 64 + 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_empty_dataset_count():
|
||||||
|
"""Test count returns 0 when no data is loaded"""
|
||||||
|
dataset = SEQDataset(window_size=64, stride=32)
|
||||||
|
assert dataset.count == 0
|
||||||
|
assert dataset.keys == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_dataset_too_short_for_window(base_test_env):
|
||||||
|
"""Dataset shorter than window_size returns __len__ == 0"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
seq_length = 30
|
||||||
|
save_h5(
|
||||||
|
test_dir,
|
||||||
|
"short",
|
||||||
|
{"sequence": [torch.randint(0, 1000, (seq_length,), dtype=torch.int64)]},
|
||||||
|
)
|
||||||
|
dataset = DatasetFactory.load("seq", test_dir, window_size=64)
|
||||||
|
assert len(dataset) == 0
|
||||||
|
assert dataset.count == seq_length
|
||||||
|
|
||||||
|
|
||||||
|
def test_unloaded_dataset_getitem_raises():
|
||||||
|
"""__getitem__ without load() should fail clearly"""
|
||||||
|
dataset = SEQDataset(window_size=64, stride=32)
|
||||||
|
with pytest.raises(RuntimeError, match="not loaded"):
|
||||||
|
dataset.get_index(0)
|
||||||
|
|
||||||
|
|
||||||
|
def test_unloaded_dataset_len():
|
||||||
|
"""__len__ without load() returns 0"""
|
||||||
|
dataset = SEQDataset(window_size=64, stride=32)
|
||||||
|
assert len(dataset) == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_base_segment_fetcher_empty():
|
||||||
|
"""BaseSegmentFetcher with empty segments list"""
|
||||||
|
fetcher = BaseSegmentFetcher([])
|
||||||
|
assert len(fetcher) == 0
|
||||||
|
with pytest.raises(ValueError, match="out of bounds"):
|
||||||
|
fetcher.fetch_data(0, 1)
|
||||||
|
|
||||||
|
|
||||||
|
def test_base_segment_fetcher_begin_equals_end(base_test_env):
|
||||||
|
"""fetch_data with begin == end returns empty tensor"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
dummy = {"sequence": [torch.randint(0, 1000, (100,), dtype=torch.int64)]}
|
||||||
|
save_h5(test_dir, "empty_fetch", dummy)
|
||||||
|
|
||||||
|
dataset = DatasetFactory.load("seq", test_dir, window_size=32)
|
||||||
|
fetcher = dataset.storage._fetcher.multi_fetchers["sequence"]
|
||||||
|
result = fetcher.fetch_data(10, 10)
|
||||||
|
assert result.numel() == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_multi_segment_fetcher_empty_dict():
|
||||||
|
"""MultiSegmentFetcher with empty dict has __len__ == 0"""
|
||||||
|
fetcher = MultiSegmentFetcher({})
|
||||||
|
assert len(fetcher) == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_storage_fetch_before_load():
|
||||||
|
"""BaseStorage.fetch before load raises RuntimeError"""
|
||||||
|
storage = H5Storage()
|
||||||
|
with pytest.raises(RuntimeError, match="not loaded"):
|
||||||
|
storage.fetch(0, 10, "sequence")
|
||||||
|
|
||||||
|
|
||||||
|
def test_detect_format_nonexistent_path():
|
||||||
|
"""detect_format raises FileNotFoundError for bad path"""
|
||||||
|
with pytest.raises(FileNotFoundError, match="No supported"):
|
||||||
|
detect_format("/nonexistent/path/xyz")
|
||||||
|
|
||||||
|
|
||||||
|
def test_detect_format_unsupported_file(base_test_env):
|
||||||
|
"""detect_format raises ValueError for unsupported file extension"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
path = os.path.join(test_dir, "data.txt")
|
||||||
|
with open(path, "w") as f:
|
||||||
|
f.write("hello")
|
||||||
|
with pytest.raises(ValueError, match="Unsupported"):
|
||||||
|
detect_format(path)
|
||||||
|
|
||||||
|
|
||||||
|
def test_create_storage_invalid_type():
|
||||||
|
"""StorageFactory.create raises ValueError for unknown type"""
|
||||||
|
with pytest.raises(ValueError, match="Unknown component"):
|
||||||
|
StorageFactory.create("parquet")
|
||||||
|
|
||||||
|
|
||||||
|
def test_json_pretokenized_without_tokenizer(base_test_env):
|
||||||
|
"""Pre-tokenized JSON (List[List[int]]) loads without tokenizer"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
data_dir = os.path.join(test_dir, "json_pretok")
|
||||||
|
os.makedirs(data_dir, exist_ok=True)
|
||||||
|
|
||||||
|
json_path = os.path.join(data_dir, "data.json")
|
||||||
|
with open(json_path, "w", encoding="utf-8") as f:
|
||||||
|
json.dump({"sequence": [[1, 2, 3, 4, 5], [6, 7, 8, 9, 10]]}, f)
|
||||||
|
|
||||||
|
dataset = DatasetFactory.load("seq", data_dir, window_size=4, storage_type="json")
|
||||||
|
assert len(dataset) > 0
|
||||||
|
assert dataset.count == 10
|
||||||
|
|
||||||
|
item = dataset[0]
|
||||||
|
assert item["input_ids"].tolist() == [1, 2, 3, 4]
|
||||||
|
assert item["target_ids"].tolist() == [2, 3, 4, 5]
|
||||||
|
|
||||||
|
|
||||||
|
def test_load_json_skips_config_file(base_test_env):
|
||||||
|
"""load_json skips scalar-value config files"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
with open(os.path.join(test_dir, "config.json"), "w") as f:
|
||||||
|
json.dump({"vocab_size": 1000, "dim": 16}, f)
|
||||||
|
|
||||||
|
with open(os.path.join(test_dir, "data.json"), "w") as f:
|
||||||
|
json.dump({"sequence": [[1, 2, 3, 4, 5]]}, f)
|
||||||
|
|
||||||
|
result = load_json(test_dir)
|
||||||
|
assert "sequence" in result
|
||||||
|
assert "vocab_size" not in result
|
||||||
|
assert len(result["sequence"]) == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_base_segment_fetcher_multi_segment():
|
||||||
|
"""fetch_data across multiple segment boundaries"""
|
||||||
|
segs = [
|
||||||
|
torch.tensor([1, 2, 3]),
|
||||||
|
torch.tensor([4, 5, 6, 7]),
|
||||||
|
torch.tensor([8, 9]),
|
||||||
|
]
|
||||||
|
fetcher = BaseSegmentFetcher(segs)
|
||||||
|
assert len(fetcher) == 9
|
||||||
|
result = fetcher.fetch_data(2, 7)
|
||||||
|
assert result.tolist() == [3, 4, 5, 6, 7]
|
||||||
|
|||||||
@@ -5,12 +5,20 @@ from unittest.mock import MagicMock
|
|||||||
import pytest
|
import pytest
|
||||||
from fastapi.testclient import TestClient
|
from fastapi.testclient import TestClient
|
||||||
|
|
||||||
from astrai.inference.server import app
|
from astrai.inference import app
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def client():
|
def client():
|
||||||
"""Provide a test client for the FastAPI app."""
|
"""Provide a test client for the FastAPI app."""
|
||||||
|
app.state.server_config = {
|
||||||
|
"device": "cpu",
|
||||||
|
"dtype": "bfloat16",
|
||||||
|
"param_path": None,
|
||||||
|
"max_batch_size": 1,
|
||||||
|
"_test": True,
|
||||||
|
}
|
||||||
|
app.state.engine = None
|
||||||
return TestClient(app)
|
return TestClient(app)
|
||||||
|
|
||||||
|
|
||||||
@@ -39,7 +47,7 @@ def mock_engine():
|
|||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def loaded_model(mock_engine, monkeypatch):
|
def loaded_model(client, mock_engine):
|
||||||
"""Simulate that the engine is loaded."""
|
"""Simulate that the engine is loaded."""
|
||||||
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
|
app.state.engine = mock_engine
|
||||||
return mock_engine
|
return mock_engine
|
||||||
|
|||||||
@@ -0,0 +1,279 @@
|
|||||||
|
"""Unit tests for inference cache components."""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.inference import (
|
||||||
|
Allocator,
|
||||||
|
KVCache,
|
||||||
|
PagePool,
|
||||||
|
PrefixCache,
|
||||||
|
Storage,
|
||||||
|
TaskTable,
|
||||||
|
page_hash,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def make_pool(n_pages: int, page_size: int) -> PagePool:
|
||||||
|
return PagePool(Allocator(n_pages), PrefixCache(page_size))
|
||||||
|
|
||||||
|
|
||||||
|
def test_page_hash_full_page():
|
||||||
|
token_ids = list(range(256))
|
||||||
|
h = page_hash(token_ids, 0, 64)
|
||||||
|
assert isinstance(h, int)
|
||||||
|
assert h >= 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_page_hash_different_page_differs():
|
||||||
|
token_ids = list(range(256))
|
||||||
|
assert page_hash(token_ids, 0, 64) != page_hash(token_ids, 1, 64)
|
||||||
|
|
||||||
|
|
||||||
|
def test_page_pool_alloc_free_cycle():
|
||||||
|
pool = make_pool(4, 64)
|
||||||
|
a = pool.alloc()
|
||||||
|
b = pool.alloc()
|
||||||
|
assert a != b
|
||||||
|
pool.free(a)
|
||||||
|
pool.free(b)
|
||||||
|
c = pool.alloc()
|
||||||
|
assert c in (a, b)
|
||||||
|
|
||||||
|
|
||||||
|
def test_page_pool_alloc_when_full():
|
||||||
|
pool = make_pool(2, 64)
|
||||||
|
pool.alloc()
|
||||||
|
pool.alloc()
|
||||||
|
assert pool.alloc() == -1
|
||||||
|
|
||||||
|
|
||||||
|
def test_page_pool_lru_eviction():
|
||||||
|
pool = make_pool(2, 64)
|
||||||
|
p0 = pool.alloc()
|
||||||
|
p1 = pool.alloc()
|
||||||
|
pool.record(p0, list(range(64)), 0)
|
||||||
|
pool.record(p1, list(range(64, 128)), 0)
|
||||||
|
pool.free(p0)
|
||||||
|
pool.free(p1)
|
||||||
|
pool.alloc()
|
||||||
|
assert p0 in pool._alloc._lru or p1 in pool._alloc._lru
|
||||||
|
|
||||||
|
|
||||||
|
def test_page_pool_inc_ref_and_free():
|
||||||
|
pool = make_pool(2, 64)
|
||||||
|
p = pool.alloc()
|
||||||
|
pool.inc_ref(p)
|
||||||
|
assert pool._alloc._refs[p] == 2
|
||||||
|
pool.free(p)
|
||||||
|
assert pool._alloc._refs[p] == 1
|
||||||
|
pool.free(p)
|
||||||
|
assert pool._alloc._refs[p] == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_page_pool_keep_cached_realloc():
|
||||||
|
"""Free mask has priority over LRU; cached page returned only when no free pages."""
|
||||||
|
pool = make_pool(3, 64)
|
||||||
|
p0 = pool.alloc()
|
||||||
|
p1 = pool.alloc()
|
||||||
|
p2 = pool.alloc()
|
||||||
|
for p in (p0, p1, p2):
|
||||||
|
pool.record(p, [p] * 64, 0)
|
||||||
|
pool.free(p0)
|
||||||
|
pool.free(p1)
|
||||||
|
pool.free(p2)
|
||||||
|
assert pool.alloc() == p0
|
||||||
|
|
||||||
|
|
||||||
|
def test_prefix_cache_lookup_returns_hits():
|
||||||
|
token_ids = list(range(256))
|
||||||
|
pool = make_pool(16, 64)
|
||||||
|
pages = [pool.alloc() for _ in range(4)]
|
||||||
|
for i, p in enumerate(pages):
|
||||||
|
pool.record(p, token_ids, i)
|
||||||
|
pool.free(p)
|
||||||
|
hits = pool.lookup(token_ids)
|
||||||
|
assert hits == pages
|
||||||
|
|
||||||
|
|
||||||
|
def test_prefix_cache_lookup_stops_at_first_miss():
|
||||||
|
token_ids = list(range(256))
|
||||||
|
pool = make_pool(16, 64)
|
||||||
|
p0 = pool.alloc()
|
||||||
|
pool.record(p0, token_ids, 0)
|
||||||
|
pool.free(p0)
|
||||||
|
p1 = pool.alloc()
|
||||||
|
pool.record(p1, [99] * 64, 1)
|
||||||
|
pool.free(p1)
|
||||||
|
hits = pool.lookup(token_ids)
|
||||||
|
assert len(hits) == 1
|
||||||
|
assert hits[0] == p0
|
||||||
|
|
||||||
|
|
||||||
|
def test_prefix_cache_ignores_partial_last_page():
|
||||||
|
token_ids = list(range(100))
|
||||||
|
pool = make_pool(16, 64)
|
||||||
|
p = pool.alloc()
|
||||||
|
pool.record(p, token_ids, 0)
|
||||||
|
pool.free(p)
|
||||||
|
hits = pool.lookup(token_ids)
|
||||||
|
assert len(hits) == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_prefix_cache_on_evict_clears_mappings():
|
||||||
|
pool = make_pool(4, 64)
|
||||||
|
p = pool.alloc()
|
||||||
|
pool.record(p, list(range(64)), 0)
|
||||||
|
pool.free(p)
|
||||||
|
assert p in pool._prefix._page_to_hash
|
||||||
|
pool._prefix.evict(p)
|
||||||
|
assert p not in pool._prefix._page_to_hash
|
||||||
|
|
||||||
|
|
||||||
|
def test_prefix_cache_has_page():
|
||||||
|
pool = make_pool(4, 64)
|
||||||
|
p = pool.alloc()
|
||||||
|
assert p not in pool._prefix._page_to_hash
|
||||||
|
pool.record(p, list(range(64)), 0)
|
||||||
|
pool.free(p)
|
||||||
|
assert p in pool._prefix._page_to_hash
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_table_set_get():
|
||||||
|
table = TaskTable(page_size=64)
|
||||||
|
table.set("task1", [0, 1, 2], 128)
|
||||||
|
assert table.get("task1") == [0, 1, 2]
|
||||||
|
assert table.get_cached("task1") == 128
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_table_get_missing():
|
||||||
|
table = TaskTable(page_size=64)
|
||||||
|
assert table.get("nonexistent") == []
|
||||||
|
assert table.get_cached("nonexistent") == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_table_pop():
|
||||||
|
table = TaskTable(page_size=64)
|
||||||
|
table.set("task1", [0, 1], 64)
|
||||||
|
pages, cached = table.pop("task1")
|
||||||
|
assert pages == [0, 1]
|
||||||
|
assert cached == 64
|
||||||
|
assert table.get("task1") == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_kv_cache_task_extend_allocates():
|
||||||
|
cache = KVCache(
|
||||||
|
n_layers=1,
|
||||||
|
n_pages=8,
|
||||||
|
page_size=64,
|
||||||
|
n_kv_heads=2,
|
||||||
|
head_dim=8,
|
||||||
|
device=torch.device("cpu"),
|
||||||
|
dtype=torch.float32,
|
||||||
|
)
|
||||||
|
cache._table.set("task1", [], 0)
|
||||||
|
ok = cache.task_extend("task1", 200)
|
||||||
|
assert ok
|
||||||
|
assert len(cache._table.get("task1")) == 4
|
||||||
|
|
||||||
|
|
||||||
|
def test_kv_cache_task_extend_fails_when_pool_full():
|
||||||
|
cache = KVCache(
|
||||||
|
n_layers=1,
|
||||||
|
n_pages=2,
|
||||||
|
page_size=64,
|
||||||
|
n_kv_heads=2,
|
||||||
|
head_dim=8,
|
||||||
|
device=torch.device("cpu"),
|
||||||
|
dtype=torch.float32,
|
||||||
|
)
|
||||||
|
cache._table.set("task1", [0, 1], 0)
|
||||||
|
ok = cache.task_extend("task1", 300)
|
||||||
|
assert not ok
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_table_table_tensor():
|
||||||
|
table = TaskTable(page_size=64)
|
||||||
|
table.set("a", [0, 1], 0)
|
||||||
|
table.set("b", [2, 3, 4], 0)
|
||||||
|
t = table.table_tensor(["a", "b"], torch.device("cpu"))
|
||||||
|
assert t.shape == (2, 3)
|
||||||
|
assert t[0].tolist() == [0, 1, -1]
|
||||||
|
assert t[1].tolist() == [2, 3, 4]
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_table_table_tensor_empty_input():
|
||||||
|
table = TaskTable(page_size=64)
|
||||||
|
t = table.table_tensor([], torch.device("cpu"))
|
||||||
|
assert t.numel() == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_storage_write_gather_single_page():
|
||||||
|
storage = Storage(
|
||||||
|
n_layers=2,
|
||||||
|
n_pages=8,
|
||||||
|
page_size=4,
|
||||||
|
n_kv_heads=2,
|
||||||
|
head_dim=8,
|
||||||
|
device=torch.device("cpu"),
|
||||||
|
dtype=torch.float32,
|
||||||
|
)
|
||||||
|
page_table = torch.tensor([[0]], dtype=torch.long)
|
||||||
|
k = torch.randn(1, 2, 2, 8)
|
||||||
|
v = torch.randn(1, 2, 2, 8)
|
||||||
|
|
||||||
|
storage.write(0, page_table, 0, k, v)
|
||||||
|
gk, gv = storage.gather(0, page_table, 2)
|
||||||
|
assert torch.allclose(gk, k)
|
||||||
|
|
||||||
|
|
||||||
|
def test_storage_write_cross_page():
|
||||||
|
storage = Storage(
|
||||||
|
n_layers=1,
|
||||||
|
n_pages=8,
|
||||||
|
page_size=4,
|
||||||
|
n_kv_heads=2,
|
||||||
|
head_dim=8,
|
||||||
|
device=torch.device("cpu"),
|
||||||
|
dtype=torch.float32,
|
||||||
|
)
|
||||||
|
page_table = torch.tensor([[0, 1]], dtype=torch.long)
|
||||||
|
k = torch.randn(1, 8, 2, 8)
|
||||||
|
v = torch.randn(1, 8, 2, 8)
|
||||||
|
|
||||||
|
storage.write(0, page_table, 0, k, v)
|
||||||
|
gk, gv = storage.gather(0, page_table, 8)
|
||||||
|
assert torch.allclose(gk, k)
|
||||||
|
|
||||||
|
|
||||||
|
def test_storage_gather_truncates_to_total_len():
|
||||||
|
storage = Storage(
|
||||||
|
n_layers=1,
|
||||||
|
n_pages=8,
|
||||||
|
page_size=4,
|
||||||
|
n_kv_heads=2,
|
||||||
|
head_dim=8,
|
||||||
|
device=torch.device("cpu"),
|
||||||
|
dtype=torch.float32,
|
||||||
|
)
|
||||||
|
page_table = torch.tensor([[0, 1]], dtype=torch.long)
|
||||||
|
k = torch.randn(1, 6, 2, 8)
|
||||||
|
v = torch.randn(1, 6, 2, 8)
|
||||||
|
storage.write(0, page_table, 0, k, v)
|
||||||
|
|
||||||
|
gk, gv = storage.gather(0, page_table, 5)
|
||||||
|
assert gk.shape == (1, 5, 2, 8)
|
||||||
|
|
||||||
|
|
||||||
|
def test_storage_gather_clamps_negative_padding():
|
||||||
|
storage = Storage(
|
||||||
|
n_layers=1,
|
||||||
|
n_pages=8,
|
||||||
|
page_size=4,
|
||||||
|
n_kv_heads=2,
|
||||||
|
head_dim=8,
|
||||||
|
device=torch.device("cpu"),
|
||||||
|
dtype=torch.float32,
|
||||||
|
)
|
||||||
|
page_table = torch.tensor([[0, -1]], dtype=torch.long)
|
||||||
|
gk, gv = storage.gather(0, page_table, 4)
|
||||||
|
assert gk.shape == (1, 4, 2, 8)
|
||||||
@@ -0,0 +1,181 @@
|
|||||||
|
"""Unit tests for GenerateResult accumulator and InferenceEngine.generate()."""
|
||||||
|
|
||||||
|
import threading
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
from astrai.inference import STOP
|
||||||
|
from astrai.inference.engine import GenerateResult
|
||||||
|
|
||||||
|
|
||||||
|
def test_result_append_single():
|
||||||
|
r = GenerateResult(count=1)
|
||||||
|
r.append("hello", 0)
|
||||||
|
assert r.results[0] == "hello"
|
||||||
|
|
||||||
|
|
||||||
|
def test_result_append_multiple_tasks():
|
||||||
|
r = GenerateResult(count=3)
|
||||||
|
r.append("a", 0)
|
||||||
|
r.append("b", 1)
|
||||||
|
r.append("c", 2)
|
||||||
|
assert r.results[0] == "a"
|
||||||
|
assert r.results[1] == "b"
|
||||||
|
assert r.results[2] == "c"
|
||||||
|
|
||||||
|
|
||||||
|
def test_result_stop_marks_complete():
|
||||||
|
r = GenerateResult(count=2)
|
||||||
|
r.append("text", 0)
|
||||||
|
r.append(STOP, 0)
|
||||||
|
r.append("more", 1)
|
||||||
|
assert r._done[0] is True
|
||||||
|
assert r._done[1] is False
|
||||||
|
assert r._completed == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_result_stop_does_not_double_count():
|
||||||
|
r = GenerateResult(count=1)
|
||||||
|
r.append(STOP, 0)
|
||||||
|
r.append(STOP, 0)
|
||||||
|
assert r._completed == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_result_pop_all_returns_and_clears():
|
||||||
|
r = GenerateResult(count=2)
|
||||||
|
r.append("a", 0)
|
||||||
|
r.append("b", 1)
|
||||||
|
out = r.pop_all()
|
||||||
|
assert len(out) == 2
|
||||||
|
assert out[0] == (0, "a")
|
||||||
|
assert out[1] == (1, "b")
|
||||||
|
assert r.pop_all() == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_result_wait_blocks_until_data():
|
||||||
|
r = GenerateResult(count=1)
|
||||||
|
|
||||||
|
def delayed_append():
|
||||||
|
import time
|
||||||
|
|
||||||
|
time.sleep(0.05)
|
||||||
|
r.append("delayed", 0)
|
||||||
|
|
||||||
|
t = threading.Thread(target=delayed_append)
|
||||||
|
t.start()
|
||||||
|
ok = r.wait(timeout=5.0)
|
||||||
|
t.join()
|
||||||
|
assert ok
|
||||||
|
assert r.results[0] == "delayed"
|
||||||
|
|
||||||
|
|
||||||
|
def test_result_wait_timeout():
|
||||||
|
r = GenerateResult(count=1)
|
||||||
|
ok = r.wait(timeout=0.01)
|
||||||
|
assert not ok
|
||||||
|
|
||||||
|
|
||||||
|
def test_result_wait_completion_non_streaming():
|
||||||
|
r = GenerateResult(count=2)
|
||||||
|
|
||||||
|
def finish_later():
|
||||||
|
import time
|
||||||
|
|
||||||
|
time.sleep(0.05)
|
||||||
|
r.append(STOP, 0)
|
||||||
|
time.sleep(0.05)
|
||||||
|
r.append(STOP, 1)
|
||||||
|
|
||||||
|
t = threading.Thread(target=finish_later)
|
||||||
|
t.start()
|
||||||
|
r.wait_completion()
|
||||||
|
t.join()
|
||||||
|
assert r._completed == 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_result_get_results():
|
||||||
|
r = GenerateResult(count=2)
|
||||||
|
r.append("hello", 0)
|
||||||
|
r.append("world", 1)
|
||||||
|
results = r.get_results()
|
||||||
|
assert results == ["hello", "world"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_engine_generate_non_streaming_single():
|
||||||
|
from astrai.inference.engine import InferenceEngine
|
||||||
|
|
||||||
|
mock_model = MagicMock()
|
||||||
|
mock_tokenizer = MagicMock()
|
||||||
|
mock_tokenizer.encode.return_value = [1, 2, 3]
|
||||||
|
mock_tokenizer.decode.return_value = "response"
|
||||||
|
mock_tokenizer.stop_ids = [0]
|
||||||
|
|
||||||
|
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
|
||||||
|
instance = MockSched.return_value
|
||||||
|
|
||||||
|
def fake_add(prompt, **kw):
|
||||||
|
cb = kw["stream_callback"]
|
||||||
|
cb("response")
|
||||||
|
cb(STOP)
|
||||||
|
|
||||||
|
instance.add_task.side_effect = fake_add
|
||||||
|
instance.remove_task.return_value = []
|
||||||
|
|
||||||
|
eng = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=1)
|
||||||
|
result = eng.generate("hello")
|
||||||
|
assert result == "response"
|
||||||
|
|
||||||
|
|
||||||
|
def test_engine_generate_streaming_yields_tokens():
|
||||||
|
from astrai.inference.engine import InferenceEngine
|
||||||
|
|
||||||
|
mock_model = MagicMock()
|
||||||
|
mock_tokenizer = MagicMock()
|
||||||
|
mock_tokenizer.encode.return_value = [1, 2, 3]
|
||||||
|
mock_tokenizer.decode.return_value = "tok"
|
||||||
|
mock_tokenizer.stop_ids = [0]
|
||||||
|
|
||||||
|
callbacks_saved = []
|
||||||
|
|
||||||
|
def capture_cb(prompt, **kw):
|
||||||
|
callbacks_saved.append(kw.get("stream_callback"))
|
||||||
|
|
||||||
|
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
|
||||||
|
instance = MockSched.return_value
|
||||||
|
instance.add_task.side_effect = capture_cb
|
||||||
|
instance.remove_task.return_value = []
|
||||||
|
|
||||||
|
eng = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=1)
|
||||||
|
gen = eng.generate("hello", stream=True)
|
||||||
|
|
||||||
|
cb = callbacks_saved[0]
|
||||||
|
cb("t1")
|
||||||
|
cb("t2")
|
||||||
|
cb(STOP)
|
||||||
|
|
||||||
|
tokens = list(gen)
|
||||||
|
assert tokens == ["t1", "t2"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_engine_generate_non_streaming_batch():
|
||||||
|
from astrai.inference.engine import InferenceEngine
|
||||||
|
|
||||||
|
mock_model = MagicMock()
|
||||||
|
mock_tokenizer = MagicMock()
|
||||||
|
mock_tokenizer.encode.return_value = [1, 2, 3]
|
||||||
|
mock_tokenizer.decode.return_value = "r"
|
||||||
|
mock_tokenizer.stop_ids = [0]
|
||||||
|
|
||||||
|
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
|
||||||
|
instance = MockSched.return_value
|
||||||
|
|
||||||
|
def fake_add(prompt, **kw):
|
||||||
|
cb = kw["stream_callback"]
|
||||||
|
cb("r")
|
||||||
|
cb(STOP)
|
||||||
|
|
||||||
|
instance.add_task.side_effect = fake_add
|
||||||
|
instance.remove_task.return_value = []
|
||||||
|
|
||||||
|
eng = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=2)
|
||||||
|
results = eng.generate(["hello", "world"])
|
||||||
|
assert results == ["r", "r"]
|
||||||
@@ -0,0 +1,127 @@
|
|||||||
|
"""Unit tests for inference sampling strategies."""
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.inference.sample import (
|
||||||
|
SamplingPipeline,
|
||||||
|
TemperatureStrategy,
|
||||||
|
TopKStrategy,
|
||||||
|
TopPStrategy,
|
||||||
|
sample,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_temperature_scalar():
|
||||||
|
logits = torch.tensor([[1.0, 2.0, 3.0]])
|
||||||
|
s = TemperatureStrategy(0.5)
|
||||||
|
result = s.apply(logits.clone())
|
||||||
|
assert torch.allclose(result, logits / 0.5)
|
||||||
|
|
||||||
|
|
||||||
|
def test_temperature_skip_when_one():
|
||||||
|
logits = torch.tensor([[1.0, 2.0, 3.0]])
|
||||||
|
s = TemperatureStrategy(1.0)
|
||||||
|
result = s.apply(logits.clone())
|
||||||
|
assert torch.equal(result, logits)
|
||||||
|
|
||||||
|
|
||||||
|
def test_temperature_per_sample_tensor():
|
||||||
|
logits = torch.tensor([[1.0, 2.0], [3.0, 4.0]])
|
||||||
|
s = TemperatureStrategy(torch.tensor([0.5, 0.5]))
|
||||||
|
result = s.apply(logits.clone())
|
||||||
|
assert torch.allclose(result, logits / 0.5)
|
||||||
|
|
||||||
|
|
||||||
|
def test_top_k_keeps_top():
|
||||||
|
logits = torch.tensor([[0.1, 0.5, 0.3, 0.9, 0.2]])
|
||||||
|
s = TopKStrategy(top_k=2)
|
||||||
|
result = s.apply(logits.clone(), filter_value=-1e9)
|
||||||
|
kept = (result > -1e9).sum().item()
|
||||||
|
assert kept == 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_top_k_skip_when_zero():
|
||||||
|
logits = torch.tensor([[1.0, 2.0, 3.0]])
|
||||||
|
s = TopKStrategy(top_k=0)
|
||||||
|
result = s.apply(logits.clone())
|
||||||
|
assert torch.equal(result, logits)
|
||||||
|
|
||||||
|
|
||||||
|
def test_top_k_batch_tensor():
|
||||||
|
"""Each row respects its own top_k."""
|
||||||
|
logits = torch.tensor([[0.1, 0.5, 0.3], [0.9, 0.2, 0.1]])
|
||||||
|
s = TopKStrategy(top_k=torch.tensor([2, 1]))
|
||||||
|
result = s.apply(logits.clone(), filter_value=-1e9)
|
||||||
|
assert (result[0] > -1e9).sum() == 2
|
||||||
|
assert (result[1] > -1e9).sum() == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_top_p_nucleus_filtering():
|
||||||
|
logits = torch.tensor([[10.0, 1.0, 1.0, 1.0, 1.0]])
|
||||||
|
s = TopPStrategy(top_p=0.5)
|
||||||
|
result = s.apply(logits.clone(), filter_value=-1e9)
|
||||||
|
kept = (result > -1e9).sum().item()
|
||||||
|
assert kept >= 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_top_p_skip_when_one():
|
||||||
|
logits = torch.tensor([[1.0, 2.0, 3.0]])
|
||||||
|
s = TopPStrategy(top_p=1.0)
|
||||||
|
result = s.apply(logits.clone())
|
||||||
|
assert torch.equal(result, logits)
|
||||||
|
|
||||||
|
|
||||||
|
def test_top_p_filter_all_except_max_when_zero():
|
||||||
|
logits = torch.tensor([[0.1, 0.5, 0.3, 0.9, 0.2]])
|
||||||
|
s = TopPStrategy(top_p=0.0)
|
||||||
|
result = s.apply(logits.clone(), filter_value=-1e9)
|
||||||
|
kept = (result > -1e9).sum().item()
|
||||||
|
assert kept == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_sampling_pipeline_composes_strategies():
|
||||||
|
logits = torch.tensor([[1.0, 2.0, 3.0, 4.0, 5.0]])
|
||||||
|
pipeline = SamplingPipeline(
|
||||||
|
[
|
||||||
|
TemperatureStrategy(0.8),
|
||||||
|
TopKStrategy(3),
|
||||||
|
TopPStrategy(0.95),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
result = pipeline.apply(logits.clone(), filter_value=-1e9)
|
||||||
|
kept = (result > -1e9).sum().item()
|
||||||
|
assert 1 <= kept <= 3
|
||||||
|
|
||||||
|
|
||||||
|
def test_sampling_pipeline_sample_returns_valid_token():
|
||||||
|
logits = torch.tensor([[1.0, 2.0, 3.0, 4.0, 5.0]])
|
||||||
|
pipeline = SamplingPipeline(
|
||||||
|
[
|
||||||
|
TemperatureStrategy(0.8),
|
||||||
|
TopKStrategy(3),
|
||||||
|
TopPStrategy(0.95),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
tokens = pipeline.sample(logits)
|
||||||
|
assert tokens.shape == (1,)
|
||||||
|
assert 0 <= tokens[0] < logits.size(-1)
|
||||||
|
|
||||||
|
|
||||||
|
def test_module_sample_shortcut():
|
||||||
|
logits = torch.tensor([[1.0, 2.0, 3.0, 4.0, 5.0]])
|
||||||
|
tokens = sample(logits, temperature=0.8, top_k=3, top_p=0.95)
|
||||||
|
assert tokens.shape == (1,)
|
||||||
|
assert 0 <= tokens[0] < logits.size(-1)
|
||||||
|
|
||||||
|
|
||||||
|
def test_module_sample_batch():
|
||||||
|
logits = torch.tensor(
|
||||||
|
[
|
||||||
|
[1.0, 2.0, 3.0, 4.0, 5.0],
|
||||||
|
[5.0, 4.0, 3.0, 2.0, 1.0],
|
||||||
|
]
|
||||||
|
)
|
||||||
|
tokens = sample(logits, temperature=0.8, top_k=3, top_p=0.95)
|
||||||
|
assert tokens.shape == (2,)
|
||||||
|
for t in tokens:
|
||||||
|
assert 0 <= t < logits.size(-1)
|
||||||
@@ -1,12 +1,12 @@
|
|||||||
"""Tests for scheduler concurrency."""
|
"""Tests for scheduler concurrency."""
|
||||||
|
|
||||||
import threading
|
import threading
|
||||||
import time
|
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
from astrai.inference.scheduler import InferenceScheduler
|
from astrai.inference import InferenceScheduler
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
@@ -19,6 +19,9 @@ def mock_model_and_tokenizer():
|
|||||||
mock_model.config.dim = 128
|
mock_model.config.dim = 128
|
||||||
mock_model.config.n_layers = 2
|
mock_model.config.n_layers = 2
|
||||||
mock_model.config.max_len = 100
|
mock_model.config.max_len = 100
|
||||||
|
mock_model.parameters.return_value = iter(
|
||||||
|
[MagicMock(dtype=torch.float32, device=torch.device("cpu"))]
|
||||||
|
)
|
||||||
|
|
||||||
mock_tokenizer = MagicMock()
|
mock_tokenizer = MagicMock()
|
||||||
mock_tokenizer.encode.return_value = [1, 2, 3, 4, 5]
|
mock_tokenizer.encode.return_value = [1, 2, 3, 4, 5]
|
||||||
@@ -33,8 +36,8 @@ def test_scheduler_concurrent_add_task(mock_model_and_tokenizer):
|
|||||||
"""Test concurrent add_task operations."""
|
"""Test concurrent add_task operations."""
|
||||||
mock_model, mock_tokenizer = mock_model_and_tokenizer
|
mock_model, mock_tokenizer = mock_model_and_tokenizer
|
||||||
|
|
||||||
with patch("astrai.inference.scheduler.AutoModel"):
|
with patch("astrai.inference.core.scheduler.AutoModel"):
|
||||||
with patch("astrai.inference.scheduler.AutoTokenizer"):
|
with patch("astrai.inference.core.scheduler.AutoTokenizer"):
|
||||||
scheduler = InferenceScheduler(
|
scheduler = InferenceScheduler(
|
||||||
model=mock_model,
|
model=mock_model,
|
||||||
tokenizer=mock_tokenizer,
|
tokenizer=mock_tokenizer,
|
||||||
@@ -59,14 +62,11 @@ def test_scheduler_concurrent_add_task(mock_model_and_tokenizer):
|
|||||||
for t in threads:
|
for t in threads:
|
||||||
t.start()
|
t.start()
|
||||||
|
|
||||||
# Let some tasks be processed
|
|
||||||
time.sleep(0.1)
|
|
||||||
|
|
||||||
scheduler.stop()
|
|
||||||
|
|
||||||
for t in threads:
|
for t in threads:
|
||||||
t.join()
|
t.join()
|
||||||
|
|
||||||
|
scheduler.stop()
|
||||||
|
|
||||||
assert len(results["errors"]) == 0, f"Errors: {results['errors']}"
|
assert len(results["errors"]) == 0, f"Errors: {results['errors']}"
|
||||||
assert len(results["task_ids"]) == 50
|
assert len(results["task_ids"]) == 50
|
||||||
|
|
||||||
@@ -75,8 +75,8 @@ def test_scheduler_concurrent_add_remove_task(mock_model_and_tokenizer):
|
|||||||
"""Test concurrent add and remove task operations."""
|
"""Test concurrent add and remove task operations."""
|
||||||
mock_model, mock_tokenizer = mock_model_and_tokenizer
|
mock_model, mock_tokenizer = mock_model_and_tokenizer
|
||||||
|
|
||||||
with patch("astrai.inference.scheduler.AutoModel"):
|
with patch("astrai.inference.core.scheduler.AutoModel"):
|
||||||
with patch("astrai.inference.scheduler.AutoTokenizer"):
|
with patch("astrai.inference.core.scheduler.AutoTokenizer"):
|
||||||
scheduler = InferenceScheduler(
|
scheduler = InferenceScheduler(
|
||||||
model=mock_model,
|
model=mock_model,
|
||||||
tokenizer=mock_tokenizer,
|
tokenizer=mock_tokenizer,
|
||||||
@@ -85,19 +85,21 @@ def test_scheduler_concurrent_add_remove_task(mock_model_and_tokenizer):
|
|||||||
)
|
)
|
||||||
|
|
||||||
results = {"added": [], "removed": [], "errors": []}
|
results = {"added": [], "removed": [], "errors": []}
|
||||||
|
add_ready = threading.Event()
|
||||||
|
|
||||||
def add_worker():
|
def add_worker():
|
||||||
try:
|
try:
|
||||||
for i in range(20):
|
for i in range(20):
|
||||||
task_id = scheduler.add_task(f"prompt {i}")
|
task_id = scheduler.add_task(f"prompt {i}")
|
||||||
results["added"].append(task_id)
|
results["added"].append(task_id)
|
||||||
time.sleep(0.001)
|
if len(results["added"]) >= 10:
|
||||||
|
add_ready.set()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
results["errors"].append(f"Add: {str(e)}")
|
results["errors"].append(f"Add: {str(e)}")
|
||||||
|
|
||||||
def remove_worker():
|
def remove_worker():
|
||||||
try:
|
try:
|
||||||
time.sleep(0.05) # Wait for some tasks to be added
|
add_ready.wait(timeout=5.0)
|
||||||
for task_id in results["added"][:10]:
|
for task_id in results["added"][:10]:
|
||||||
scheduler.remove_task(task_id)
|
scheduler.remove_task(task_id)
|
||||||
results["removed"].append(task_id)
|
results["removed"].append(task_id)
|
||||||
@@ -110,11 +112,9 @@ def test_scheduler_concurrent_add_remove_task(mock_model_and_tokenizer):
|
|||||||
add_thread.start()
|
add_thread.start()
|
||||||
remove_thread.start()
|
remove_thread.start()
|
||||||
|
|
||||||
time.sleep(0.2)
|
|
||||||
scheduler.stop()
|
|
||||||
|
|
||||||
add_thread.join()
|
add_thread.join()
|
||||||
remove_thread.join()
|
remove_thread.join()
|
||||||
|
scheduler.stop()
|
||||||
|
|
||||||
assert len(results["errors"]) == 0, f"Errors: {results['errors']}"
|
assert len(results["errors"]) == 0, f"Errors: {results['errors']}"
|
||||||
assert len(results["added"]) == 20
|
assert len(results["added"]) == 20
|
||||||
@@ -124,8 +124,8 @@ def test_scheduler_concurrent_get_stats(mock_model_and_tokenizer):
|
|||||||
"""Test concurrent get_stats operations."""
|
"""Test concurrent get_stats operations."""
|
||||||
mock_model, mock_tokenizer = mock_model_and_tokenizer
|
mock_model, mock_tokenizer = mock_model_and_tokenizer
|
||||||
|
|
||||||
with patch("astrai.inference.scheduler.AutoModel"):
|
with patch("astrai.inference.core.scheduler.AutoModel"):
|
||||||
with patch("astrai.inference.scheduler.AutoTokenizer"):
|
with patch("astrai.inference.core.scheduler.AutoTokenizer"):
|
||||||
scheduler = InferenceScheduler(
|
scheduler = InferenceScheduler(
|
||||||
model=mock_model,
|
model=mock_model,
|
||||||
tokenizer=mock_tokenizer,
|
tokenizer=mock_tokenizer,
|
||||||
@@ -134,21 +134,24 @@ def test_scheduler_concurrent_get_stats(mock_model_and_tokenizer):
|
|||||||
)
|
)
|
||||||
|
|
||||||
results = {"stats": [], "errors": []}
|
results = {"stats": [], "errors": []}
|
||||||
|
started = threading.Event()
|
||||||
|
stats_done = threading.Event()
|
||||||
|
|
||||||
def add_tasks():
|
def add_tasks():
|
||||||
try:
|
try:
|
||||||
for i in range(20):
|
for i in range(20):
|
||||||
scheduler.add_task(f"prompt {i}")
|
scheduler.add_task(f"prompt {i}")
|
||||||
time.sleep(0.001)
|
started.set()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
results["errors"].append(f"Add: {str(e)}")
|
results["errors"].append(f"Add: {str(e)}")
|
||||||
|
|
||||||
def get_stats():
|
def get_stats():
|
||||||
try:
|
try:
|
||||||
|
started.wait(timeout=5.0)
|
||||||
for _ in range(50):
|
for _ in range(50):
|
||||||
stats = scheduler.get_stats()
|
stats = scheduler.get_stats()
|
||||||
results["stats"].append(stats)
|
results["stats"].append(stats)
|
||||||
time.sleep(0.001)
|
stats_done.set()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
results["errors"].append(f"Get stats: {str(e)}")
|
results["errors"].append(f"Get stats: {str(e)}")
|
||||||
|
|
||||||
@@ -158,16 +161,15 @@ def test_scheduler_concurrent_get_stats(mock_model_and_tokenizer):
|
|||||||
add_thread.start()
|
add_thread.start()
|
||||||
stats_thread.start()
|
stats_thread.start()
|
||||||
|
|
||||||
time.sleep(0.3)
|
add_thread.join()
|
||||||
|
stats_done.wait(timeout=5.0)
|
||||||
scheduler.stop()
|
scheduler.stop()
|
||||||
|
|
||||||
add_thread.join()
|
|
||||||
stats_thread.join()
|
stats_thread.join()
|
||||||
|
|
||||||
assert len(results["errors"]) == 0, f"Errors: {results['errors']}"
|
assert len(results["errors"]) == 0, f"Errors: {results['errors']}"
|
||||||
assert len(results["stats"]) == 50
|
assert len(results["stats"]) == 50
|
||||||
|
|
||||||
# Verify stats are consistent
|
|
||||||
for stats in results["stats"]:
|
for stats in results["stats"]:
|
||||||
assert "total_tasks" in stats
|
assert "total_tasks" in stats
|
||||||
assert stats["total_tasks"] >= 0
|
assert stats["total_tasks"] >= 0
|
||||||
@@ -2,10 +2,12 @@
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
from astrai.inference import app
|
||||||
|
|
||||||
def test_health_no_model(client, monkeypatch):
|
|
||||||
|
def test_health_no_model(client):
|
||||||
"""GET /health should return 200 even when engine not loaded."""
|
"""GET /health should return 200 even when engine not loaded."""
|
||||||
monkeypatch.setattr("astrai.inference.server._state.engine", None)
|
app.state.engine = None
|
||||||
response = client.get("/health")
|
response = client.get("/health")
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
data = response.json()
|
data = response.json()
|
||||||
@@ -22,15 +24,14 @@ def test_health_with_model(client, loaded_model):
|
|||||||
assert data["model_loaded"] is True
|
assert data["model_loaded"] is True
|
||||||
|
|
||||||
|
|
||||||
def test_chat_completions_non_stream(client, loaded_model, monkeypatch):
|
def test_chat_completions_non_stream(client, loaded_model):
|
||||||
"""POST /v1/chat/completions with stream=false returns OpenAI-style JSON."""
|
"""POST /v1/chat/completions with stream=false returns OpenAI-style JSON."""
|
||||||
|
|
||||||
async def async_gen():
|
async def async_gen():
|
||||||
yield "Assistant reply"
|
yield "Assistant reply"
|
||||||
|
|
||||||
mock_engine = loaded_model
|
app.state.engine = loaded_model
|
||||||
mock_engine.generate_async.return_value = async_gen()
|
loaded_model.generate_async.return_value = async_gen()
|
||||||
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/v1/chat/completions",
|
"/v1/chat/completions",
|
||||||
json={
|
json={
|
||||||
@@ -48,16 +49,15 @@ def test_chat_completions_non_stream(client, loaded_model, monkeypatch):
|
|||||||
assert "prompt_tokens" in data["usage"]
|
assert "prompt_tokens" in data["usage"]
|
||||||
|
|
||||||
|
|
||||||
def test_chat_completions_stream(client, loaded_model, monkeypatch):
|
def test_chat_completions_stream(client, loaded_model):
|
||||||
"""POST /v1/chat/completions with stream=true returns SSE stream."""
|
"""POST /v1/chat/completions with stream=true returns SSE stream."""
|
||||||
|
|
||||||
async def async_gen():
|
async def async_gen():
|
||||||
yield "cumulative1"
|
yield "cumulative1"
|
||||||
yield "cumulative2"
|
yield "cumulative2"
|
||||||
|
|
||||||
mock_engine = loaded_model
|
app.state.engine = loaded_model
|
||||||
mock_engine.generate_async.return_value = async_gen()
|
loaded_model.generate_async.return_value = async_gen()
|
||||||
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/v1/chat/completions",
|
"/v1/chat/completions",
|
||||||
json={
|
json={
|
||||||
@@ -77,15 +77,14 @@ def test_chat_completions_stream(client, loaded_model, monkeypatch):
|
|||||||
assert any("[DONE]" in line for line in lines)
|
assert any("[DONE]" in line for line in lines)
|
||||||
|
|
||||||
|
|
||||||
def test_messages_non_stream(client, loaded_model, monkeypatch):
|
def test_messages_non_stream(client, loaded_model):
|
||||||
"""POST /v1/messages with stream=false returns Anthropic-style JSON."""
|
"""POST /v1/messages with stream=false returns Anthropic-style JSON."""
|
||||||
|
|
||||||
async def async_gen():
|
async def async_gen():
|
||||||
yield "Assistant reply"
|
yield "Assistant reply"
|
||||||
|
|
||||||
mock_engine = loaded_model
|
app.state.engine = loaded_model
|
||||||
mock_engine.generate_async.return_value = async_gen()
|
loaded_model.generate_async.return_value = async_gen()
|
||||||
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/v1/messages",
|
"/v1/messages",
|
||||||
json={
|
json={
|
||||||
@@ -105,16 +104,15 @@ def test_messages_non_stream(client, loaded_model, monkeypatch):
|
|||||||
assert "input_tokens" in data["usage"]
|
assert "input_tokens" in data["usage"]
|
||||||
|
|
||||||
|
|
||||||
def test_messages_stream(client, loaded_model, monkeypatch):
|
def test_messages_stream(client, loaded_model):
|
||||||
"""POST /v1/messages with stream=true returns Anthropic SSE stream."""
|
"""POST /v1/messages with stream=true returns Anthropic SSE stream."""
|
||||||
|
|
||||||
async def async_gen():
|
async def async_gen():
|
||||||
yield "cumulative1"
|
yield "cumulative1"
|
||||||
yield "cumulative2"
|
yield "cumulative2"
|
||||||
|
|
||||||
mock_engine = loaded_model
|
app.state.engine = loaded_model
|
||||||
mock_engine.generate_async.return_value = async_gen()
|
loaded_model.generate_async.return_value = async_gen()
|
||||||
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/v1/messages",
|
"/v1/messages",
|
||||||
json={
|
json={
|
||||||
@@ -137,15 +135,14 @@ def test_messages_stream(client, loaded_model, monkeypatch):
|
|||||||
assert "message_stop" in content
|
assert "message_stop" in content
|
||||||
|
|
||||||
|
|
||||||
def test_messages_with_system(client, loaded_model, monkeypatch):
|
def test_messages_with_system(client, loaded_model):
|
||||||
"""POST /v1/messages with system prompt."""
|
"""POST /v1/messages with system prompt."""
|
||||||
|
|
||||||
async def async_gen():
|
async def async_gen():
|
||||||
yield "Reply"
|
yield "Reply"
|
||||||
|
|
||||||
mock_engine = loaded_model
|
app.state.engine = loaded_model
|
||||||
mock_engine.generate_async.return_value = async_gen()
|
loaded_model.generate_async.return_value = async_gen()
|
||||||
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/v1/messages",
|
"/v1/messages",
|
||||||
json={
|
json={
|
||||||
@@ -160,5 +157,60 @@ def test_messages_with_system(client, loaded_model, monkeypatch):
|
|||||||
assert data["type"] == "message"
|
assert data["type"] == "message"
|
||||||
|
|
||||||
|
|
||||||
|
def test_chat_completions_stop_sequence(client, loaded_model):
|
||||||
|
"""POST /v1/chat/completions with stop parameter truncates at stop sequence."""
|
||||||
|
|
||||||
|
async def async_gen():
|
||||||
|
yield "Hello"
|
||||||
|
yield "X"
|
||||||
|
yield "world"
|
||||||
|
|
||||||
|
app.state.engine = loaded_model
|
||||||
|
loaded_model.generate_async.return_value = async_gen()
|
||||||
|
response = client.post(
|
||||||
|
"/v1/chat/completions",
|
||||||
|
json={
|
||||||
|
"messages": [{"role": "user", "content": "Hello"}],
|
||||||
|
"max_tokens": 100,
|
||||||
|
"stream": False,
|
||||||
|
"stop": ["X"],
|
||||||
|
},
|
||||||
|
)
|
||||||
|
assert response.status_code == 200
|
||||||
|
data = response.json()
|
||||||
|
content = data["choices"][0]["message"]["content"]
|
||||||
|
assert "X" in content
|
||||||
|
assert "world" not in content
|
||||||
|
|
||||||
|
|
||||||
|
def test_chat_completions_stop_sequence_stream(client, loaded_model):
|
||||||
|
"""POST /v1/chat/completions with stop parameter truncates SSE stream."""
|
||||||
|
|
||||||
|
async def async_gen():
|
||||||
|
yield "Hello"
|
||||||
|
yield "X"
|
||||||
|
yield "world"
|
||||||
|
|
||||||
|
app.state.engine = loaded_model
|
||||||
|
loaded_model.generate_async.return_value = async_gen()
|
||||||
|
response = client.post(
|
||||||
|
"/v1/chat/completions",
|
||||||
|
json={
|
||||||
|
"messages": [{"role": "user", "content": "Hello"}],
|
||||||
|
"max_tokens": 100,
|
||||||
|
"stream": True,
|
||||||
|
"stop": ["X"],
|
||||||
|
},
|
||||||
|
headers={"Accept": "text/event-stream"},
|
||||||
|
)
|
||||||
|
assert response.status_code == 200
|
||||||
|
content = response.content.decode("utf-8")
|
||||||
|
assert "Hello" in content
|
||||||
|
assert "world" not in content
|
||||||
|
assert any(
|
||||||
|
"finish_reason" in line for line in content.split("\n") if "stop" in line
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
pytest.main([__file__, "-v"])
|
pytest.main([__file__, "-v"])
|
||||||
|
|||||||
@@ -0,0 +1,170 @@
|
|||||||
|
"""Unit tests for Task and TaskManager."""
|
||||||
|
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
from astrai.inference import STOP, Task, TaskManager, TaskStatus
|
||||||
|
|
||||||
|
|
||||||
|
def _make_mock_tokenizer():
|
||||||
|
t = MagicMock()
|
||||||
|
t.encode.return_value = [1, 2, 3, 4, 5]
|
||||||
|
t.stop_ids = [0]
|
||||||
|
return t
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_default_status_is_pending():
|
||||||
|
task = Task("id1", [1, 2, 3])
|
||||||
|
assert task.status == TaskStatus.PENDING
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_next_pos():
|
||||||
|
task = Task("id1", [1, 2, 3])
|
||||||
|
task.input_tokens = 5
|
||||||
|
assert task.next_pos == 5
|
||||||
|
task.output_ids.append(4)
|
||||||
|
assert task.next_pos == 6
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_is_finished_max_tokens():
|
||||||
|
task = Task("id1", [1, 2, 3], max_tokens=2)
|
||||||
|
task.output_tokens = 2
|
||||||
|
assert task.is_finished([])
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_is_finished_stop_id():
|
||||||
|
task = Task("id1", [1, 2, 3])
|
||||||
|
task.output_ids = [5, 0]
|
||||||
|
assert task.is_finished([0])
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_is_finished_not_yet():
|
||||||
|
task = Task("id1", [1, 2, 3], max_tokens=10)
|
||||||
|
task.output_ids = [1, 2]
|
||||||
|
assert not task.is_finished([0])
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_manager_add_task():
|
||||||
|
tm = TaskManager(tokenizer=_make_mock_tokenizer())
|
||||||
|
tid = tm.add_task("hello")
|
||||||
|
assert tid.startswith("task_")
|
||||||
|
assert tm._total_tasks == 1
|
||||||
|
assert len(tm.waiting_queue) == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_manager_add_task_too_long_immediate_stop():
|
||||||
|
t = _make_mock_tokenizer()
|
||||||
|
t.encode.return_value = list(range(9000))
|
||||||
|
cb_calls = []
|
||||||
|
|
||||||
|
tm = TaskManager(tokenizer=t, max_seq_len=16)
|
||||||
|
tm.add_task("long", stream_callback=lambda tok: cb_calls.append(tok))
|
||||||
|
assert cb_calls[0] is STOP
|
||||||
|
assert len(tm.waiting_queue) == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_manager_remove_task():
|
||||||
|
tm = TaskManager(tokenizer=_make_mock_tokenizer())
|
||||||
|
tid = tm.add_task("test")
|
||||||
|
tm.remove_task(tid)
|
||||||
|
assert len(tm.waiting_queue) == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_manager_remove_active_task():
|
||||||
|
tm = TaskManager(tokenizer=_make_mock_tokenizer())
|
||||||
|
tid = tm.add_task("test")
|
||||||
|
tasks = tm.pull_candidates(1)
|
||||||
|
tm.activate(tasks[0])
|
||||||
|
assert len(tm.active_tasks) == 1
|
||||||
|
removed = tm.remove_task(tid)
|
||||||
|
assert len(removed) == 1
|
||||||
|
assert len(tm.active_tasks) == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_manager_pull_candidates_fifo():
|
||||||
|
tm = TaskManager(tokenizer=_make_mock_tokenizer())
|
||||||
|
tm.add_task("a")
|
||||||
|
tm.add_task("b")
|
||||||
|
tm.add_task("c")
|
||||||
|
pulled = tm.pull_candidates(2)
|
||||||
|
assert len(pulled) == 2
|
||||||
|
assert pulled[0].prompt_ids == [1, 2, 3, 4, 5]
|
||||||
|
assert len(tm.waiting_queue) == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_manager_activate():
|
||||||
|
tm = TaskManager(tokenizer=_make_mock_tokenizer())
|
||||||
|
tm.add_task("test")
|
||||||
|
task = tm.pull_candidates(1)[0]
|
||||||
|
tm.activate(task)
|
||||||
|
assert task.status == TaskStatus.RUNNING
|
||||||
|
assert task in tm.active_tasks
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_manager_return_to_waiting():
|
||||||
|
tm = TaskManager(tokenizer=_make_mock_tokenizer())
|
||||||
|
tm.add_task("a")
|
||||||
|
tm.add_task("b")
|
||||||
|
t1 = tm.pull_candidates(1)[0]
|
||||||
|
tm.return_to_waiting([t1])
|
||||||
|
assert len(tm.waiting_queue) == 2
|
||||||
|
assert tm.waiting_queue[0] == t1
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_manager_remove_finished_aborted():
|
||||||
|
tm = TaskManager(tokenizer=_make_mock_tokenizer())
|
||||||
|
tm.add_task("test")
|
||||||
|
task = tm.pull_candidates(1)[0]
|
||||||
|
tm.activate(task)
|
||||||
|
task.status = TaskStatus.ABORTED
|
||||||
|
finished = tm.remove_finished_tasks([0])
|
||||||
|
assert len(finished) == 1
|
||||||
|
assert len(tm.active_tasks) == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_manager_remove_finished_stop_id():
|
||||||
|
tm = TaskManager(tokenizer=_make_mock_tokenizer())
|
||||||
|
tm.add_task("test")
|
||||||
|
task = tm.pull_candidates(1)[0]
|
||||||
|
tm.activate(task)
|
||||||
|
task.output_ids = [0]
|
||||||
|
task.output_tokens = 1
|
||||||
|
finished = tm.remove_finished_tasks([0])
|
||||||
|
assert len(finished) == 1
|
||||||
|
assert task.status == TaskStatus.FINISHED
|
||||||
|
assert len(tm.active_tasks) == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_manager_has_work():
|
||||||
|
tm = TaskManager(tokenizer=_make_mock_tokenizer())
|
||||||
|
assert not tm.has_work()
|
||||||
|
tm.add_task("test")
|
||||||
|
assert tm.has_work()
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_manager_wake():
|
||||||
|
import threading
|
||||||
|
|
||||||
|
tm = TaskManager(tokenizer=_make_mock_tokenizer())
|
||||||
|
called = threading.Event()
|
||||||
|
|
||||||
|
def waiter():
|
||||||
|
tm.wait_for_tasks(timeout=5.0)
|
||||||
|
called.set()
|
||||||
|
|
||||||
|
t = threading.Thread(target=waiter)
|
||||||
|
t.start()
|
||||||
|
import time
|
||||||
|
|
||||||
|
time.sleep(0.05)
|
||||||
|
tm.wake()
|
||||||
|
t.join(timeout=2.0)
|
||||||
|
assert called.is_set()
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_manager_get_stats():
|
||||||
|
tm = TaskManager(tokenizer=_make_mock_tokenizer())
|
||||||
|
tm.add_task("test")
|
||||||
|
stats = tm.get_stats()
|
||||||
|
assert stats["total_tasks"] == 1
|
||||||
|
assert stats["waiting_queue"] == 1
|
||||||
|
assert stats["active_tasks"] == 0
|
||||||
@@ -0,0 +1,166 @@
|
|||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.config.model_config import EncoderConfig
|
||||||
|
from astrai.model.encoder import EmbeddingEncoder
|
||||||
|
|
||||||
|
TINY_CONFIG = dict(
|
||||||
|
vocab_size=128,
|
||||||
|
dim=8,
|
||||||
|
n_heads=2,
|
||||||
|
n_kv_heads=1,
|
||||||
|
dim_ffn=16,
|
||||||
|
max_len=64,
|
||||||
|
n_layers=2,
|
||||||
|
norm_eps=1e-5,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_encoder_forward_mean():
|
||||||
|
config = EncoderConfig(**TINY_CONFIG)
|
||||||
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
model = EmbeddingEncoder(config).to(device=device)
|
||||||
|
model.eval()
|
||||||
|
|
||||||
|
batch_size, seq_len = 2, 8
|
||||||
|
input_ids = torch.randint(
|
||||||
|
0, config.vocab_size, (batch_size, seq_len), device=device
|
||||||
|
)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
output = model(input_ids)
|
||||||
|
|
||||||
|
assert output.shape == (batch_size, config.dim)
|
||||||
|
assert not torch.isnan(output).any()
|
||||||
|
|
||||||
|
|
||||||
|
def test_encoder_forward_cls():
|
||||||
|
config = EncoderConfig(**{**TINY_CONFIG, "pooling_type": "cls"})
|
||||||
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
model = EmbeddingEncoder(config).to(device=device)
|
||||||
|
model.eval()
|
||||||
|
|
||||||
|
batch_size, seq_len = 2, 8
|
||||||
|
input_ids = torch.randint(
|
||||||
|
0, config.vocab_size, (batch_size, seq_len), device=device
|
||||||
|
)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
output = model(input_ids)
|
||||||
|
|
||||||
|
assert output.shape == (batch_size, config.dim)
|
||||||
|
assert not torch.isnan(output).any()
|
||||||
|
|
||||||
|
|
||||||
|
def test_encoder_forward_last():
|
||||||
|
config = EncoderConfig(**{**TINY_CONFIG, "pooling_type": "last"})
|
||||||
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
model = EmbeddingEncoder(config).to(device=device)
|
||||||
|
model.eval()
|
||||||
|
|
||||||
|
batch_size, seq_len = 2, 8
|
||||||
|
input_ids = torch.randint(
|
||||||
|
0, config.vocab_size, (batch_size, seq_len), device=device
|
||||||
|
)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
output = model(input_ids)
|
||||||
|
|
||||||
|
assert output.shape == (batch_size, config.dim)
|
||||||
|
assert not torch.isnan(output).any()
|
||||||
|
|
||||||
|
|
||||||
|
def test_encoder_forward_with_padding():
|
||||||
|
config = EncoderConfig(**TINY_CONFIG)
|
||||||
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
model = EmbeddingEncoder(config).to(device=device)
|
||||||
|
model.eval()
|
||||||
|
|
||||||
|
batch_size, seq_len = 2, 8
|
||||||
|
input_ids = torch.randint(
|
||||||
|
0, config.vocab_size, (batch_size, seq_len), device=device
|
||||||
|
)
|
||||||
|
input_mask = torch.ones(batch_size, seq_len, dtype=torch.bool, device=device)
|
||||||
|
input_mask[:, 4:] = False
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
output = model(input_ids, input_mask=input_mask)
|
||||||
|
|
||||||
|
assert output.shape == (batch_size, config.dim)
|
||||||
|
assert not torch.isnan(output).any()
|
||||||
|
|
||||||
|
|
||||||
|
def test_encoder_normalize():
|
||||||
|
config = EncoderConfig(
|
||||||
|
**{**TINY_CONFIG, "pooling_type": "mean", "normalize_embeddings": True}
|
||||||
|
)
|
||||||
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
model = EmbeddingEncoder(config).to(device=device)
|
||||||
|
model.eval()
|
||||||
|
|
||||||
|
batch_size, seq_len = 2, 8
|
||||||
|
input_ids = torch.randint(
|
||||||
|
0, config.vocab_size, (batch_size, seq_len), device=device
|
||||||
|
)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
output = model(input_ids)
|
||||||
|
|
||||||
|
norms = output.norm(p=2, dim=-1)
|
||||||
|
assert torch.allclose(norms, torch.ones_like(norms), atol=1e-4)
|
||||||
|
|
||||||
|
|
||||||
|
def test_encoder_register():
|
||||||
|
from astrai.model.automodel import AutoModel
|
||||||
|
|
||||||
|
assert AutoModel.is_registered("embedding")
|
||||||
|
cls = AutoModel.get_component_class("embedding")
|
||||||
|
assert cls is EmbeddingEncoder
|
||||||
|
|
||||||
|
|
||||||
|
def test_encoder_from_transformer_checkpoint():
|
||||||
|
config = EncoderConfig(**TINY_CONFIG)
|
||||||
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
model = EmbeddingEncoder(config).to(device=device)
|
||||||
|
|
||||||
|
state_dict = model.state_dict()
|
||||||
|
state_dict["lm_head.weight"] = torch.randn(
|
||||||
|
config.vocab_size, config.dim, device=device
|
||||||
|
)
|
||||||
|
|
||||||
|
new_model = EmbeddingEncoder(config).to(device=device)
|
||||||
|
new_model.load_state_dict(state_dict, strict=True)
|
||||||
|
|
||||||
|
for key in model.state_dict():
|
||||||
|
assert torch.equal(new_model.state_dict()[key], model.state_dict()[key])
|
||||||
|
|
||||||
|
|
||||||
|
def test_encoder_save_load():
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import tempfile
|
||||||
|
|
||||||
|
import safetensors.torch as st
|
||||||
|
|
||||||
|
test_dir = tempfile.mkdtemp(prefix="encoder_test_")
|
||||||
|
config_path = os.path.join(test_dir, "config.json")
|
||||||
|
weights_path = os.path.join(test_dir, "model.safetensors")
|
||||||
|
|
||||||
|
try:
|
||||||
|
config_data = {**TINY_CONFIG, "pooling_type": "mean"}
|
||||||
|
with open(config_path, "w") as f:
|
||||||
|
json.dump(config_data, f)
|
||||||
|
|
||||||
|
config = EncoderConfig.from_file(config_path)
|
||||||
|
original = EmbeddingEncoder(config)
|
||||||
|
st.save_file(original.state_dict(), weights_path)
|
||||||
|
|
||||||
|
loaded = EmbeddingEncoder(config)
|
||||||
|
loaded.load_state_dict(st.load_file(weights_path))
|
||||||
|
|
||||||
|
for key in original.state_dict():
|
||||||
|
assert torch.equal(original.state_dict()[key], loaded.state_dict()[key])
|
||||||
|
finally:
|
||||||
|
if os.path.exists(test_dir):
|
||||||
|
for f in os.listdir(test_dir):
|
||||||
|
os.remove(os.path.join(test_dir, f))
|
||||||
|
os.rmdir(test_dir)
|
||||||
@@ -0,0 +1,108 @@
|
|||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
|
from astrai.model.transformer import AutoRegressiveLM
|
||||||
|
|
||||||
|
TINY_CONFIG = dict(
|
||||||
|
vocab_size=128,
|
||||||
|
dim=8,
|
||||||
|
n_heads=2,
|
||||||
|
n_kv_heads=1,
|
||||||
|
dim_ffn=16,
|
||||||
|
max_len=64,
|
||||||
|
n_layers=2,
|
||||||
|
norm_eps=1e-5,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
CONFIGS = [
|
||||||
|
pytest.param(
|
||||||
|
{**TINY_CONFIG, "attn_type": "gqa", "ffn_type": "mlp"},
|
||||||
|
id="gqa_mlp",
|
||||||
|
),
|
||||||
|
pytest.param(
|
||||||
|
{
|
||||||
|
**TINY_CONFIG,
|
||||||
|
"attn_type": "mla",
|
||||||
|
"ffn_type": "mlp",
|
||||||
|
"kv_lora_rank": 4,
|
||||||
|
"qk_nope_head_dim": 2,
|
||||||
|
"qk_rope_head_dim": 2,
|
||||||
|
},
|
||||||
|
id="mla_mlp",
|
||||||
|
),
|
||||||
|
pytest.param(
|
||||||
|
{
|
||||||
|
**TINY_CONFIG,
|
||||||
|
"attn_type": "gqa",
|
||||||
|
"ffn_type": "moe",
|
||||||
|
"n_routed_experts": 4,
|
||||||
|
"n_shared_experts": 1,
|
||||||
|
"n_activated_experts": 2,
|
||||||
|
"topk_method": "greedy",
|
||||||
|
},
|
||||||
|
id="gqa_moe",
|
||||||
|
),
|
||||||
|
pytest.param(
|
||||||
|
{
|
||||||
|
**TINY_CONFIG,
|
||||||
|
"attn_type": "gqa",
|
||||||
|
"ffn_type": "mlp",
|
||||||
|
"rope_theta": 100000.0,
|
||||||
|
},
|
||||||
|
id="gqa_rope_theta",
|
||||||
|
),
|
||||||
|
pytest.param(
|
||||||
|
{**TINY_CONFIG, "attn_type": "gqa", "ffn_type": "mlp", "use_qk_norm": True},
|
||||||
|
id="gqa_qk_norm",
|
||||||
|
),
|
||||||
|
pytest.param(
|
||||||
|
{**TINY_CONFIG, "attn_type": "gqa", "ffn_type": "mlp", "tie_weight": True},
|
||||||
|
id="gqa_tie_weight",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("config_kwargs", CONFIGS)
|
||||||
|
def test_model_forward(config_kwargs):
|
||||||
|
config = AutoRegressiveLMConfig(**config_kwargs)
|
||||||
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
model = AutoRegressiveLM(config).to(device=device)
|
||||||
|
model.eval()
|
||||||
|
|
||||||
|
batch_size, seq_len = 2, 8
|
||||||
|
input_ids = torch.randint(
|
||||||
|
0, config.vocab_size, (batch_size, seq_len), device=device
|
||||||
|
)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
output = model(input_ids)
|
||||||
|
|
||||||
|
assert "logits" in output
|
||||||
|
assert "hidden_states" in output
|
||||||
|
assert output["logits"].shape == (batch_size, seq_len, config.vocab_size)
|
||||||
|
assert output["hidden_states"].shape == (batch_size, seq_len, config.dim)
|
||||||
|
assert not torch.isnan(output["logits"]).any()
|
||||||
|
assert not torch.isnan(output["hidden_states"]).any()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("config_kwargs", CONFIGS)
|
||||||
|
def test_model_forward_with_padding(config_kwargs):
|
||||||
|
config = AutoRegressiveLMConfig(**config_kwargs)
|
||||||
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
model = AutoRegressiveLM(config).to(device=device)
|
||||||
|
model.eval()
|
||||||
|
|
||||||
|
batch_size, seq_len = 2, 8
|
||||||
|
input_ids = torch.randint(
|
||||||
|
0, config.vocab_size, (batch_size, seq_len), device=device
|
||||||
|
)
|
||||||
|
input_mask = torch.ones(batch_size, seq_len, dtype=torch.bool, device=device)
|
||||||
|
input_mask[:, 4:] = False
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
output = model(input_ids, input_mask=input_mask)
|
||||||
|
|
||||||
|
assert output["logits"].shape == (batch_size, seq_len, config.vocab_size)
|
||||||
|
assert not torch.isnan(output["logits"]).any()
|
||||||
@@ -6,8 +6,8 @@ import pytest
|
|||||||
import safetensors.torch as st
|
import safetensors.torch as st
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.config.model_config import ModelConfig
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
from astrai.model.transformer import Transformer
|
from astrai.model.transformer import AutoRegressiveLM
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
@@ -17,10 +17,10 @@ def transformer_test_env():
|
|||||||
|
|
||||||
config = {
|
config = {
|
||||||
"vocab_size": 1000,
|
"vocab_size": 1000,
|
||||||
"dim": 128,
|
"dim": 8,
|
||||||
"n_heads": 4,
|
"n_heads": 2,
|
||||||
"n_kv_heads": 2,
|
"n_kv_heads": 1,
|
||||||
"dim_ffn": 256,
|
"dim_ffn": 16,
|
||||||
"max_len": 64,
|
"max_len": 64,
|
||||||
"n_layers": 2,
|
"n_layers": 2,
|
||||||
"norm_eps": 1e-5,
|
"norm_eps": 1e-5,
|
||||||
@@ -50,8 +50,8 @@ def test_tie_weight_init(transformer_test_env):
|
|||||||
with open(config_path, "w") as f:
|
with open(config_path, "w") as f:
|
||||||
json.dump(config_data, f)
|
json.dump(config_data, f)
|
||||||
|
|
||||||
config = ModelConfig().load(config_path)
|
config = AutoRegressiveLMConfig.from_file(config_path)
|
||||||
model = Transformer(config)
|
model = AutoRegressiveLM(config)
|
||||||
|
|
||||||
assert torch.equal(model.lm_head.weight, model.embed_tokens.weight)
|
assert torch.equal(model.lm_head.weight, model.embed_tokens.weight)
|
||||||
assert model.lm_head.weight.data_ptr() == model.embed_tokens.weight.data_ptr()
|
assert model.lm_head.weight.data_ptr() == model.embed_tokens.weight.data_ptr()
|
||||||
@@ -68,8 +68,8 @@ def test_tie_weight_init(transformer_test_env):
|
|||||||
with open(config_path, "w") as f:
|
with open(config_path, "w") as f:
|
||||||
json.dump(config_data, f)
|
json.dump(config_data, f)
|
||||||
|
|
||||||
config = ModelConfig().load(config_path)
|
config = AutoRegressiveLMConfig.from_file(config_path)
|
||||||
model = Transformer(config)
|
model = AutoRegressiveLM(config)
|
||||||
|
|
||||||
assert not torch.equal(model.lm_head.weight, model.embed_tokens.weight)
|
assert not torch.equal(model.lm_head.weight, model.embed_tokens.weight)
|
||||||
assert model.lm_head.weight.data_ptr() != model.embed_tokens.weight.data_ptr()
|
assert model.lm_head.weight.data_ptr() != model.embed_tokens.weight.data_ptr()
|
||||||
@@ -94,13 +94,13 @@ def test_model_save_load_with_tie_weight(transformer_test_env):
|
|||||||
with open(config_path, "w") as f:
|
with open(config_path, "w") as f:
|
||||||
json.dump(config_data, f)
|
json.dump(config_data, f)
|
||||||
|
|
||||||
config = ModelConfig().load(config_path)
|
config = AutoRegressiveLMConfig.from_file(config_path)
|
||||||
original_model = Transformer(config)
|
original_model = AutoRegressiveLM(config)
|
||||||
|
|
||||||
st.save_file(original_model.state_dict(), model_path)
|
st.save_file(original_model.state_dict(), model_path)
|
||||||
|
|
||||||
loaded_config = ModelConfig().load(config_path)
|
loaded_config = AutoRegressiveLMConfig.from_file(config_path)
|
||||||
model = Transformer(loaded_config)
|
model = AutoRegressiveLM(loaded_config)
|
||||||
model.load_state_dict(st.load_file(model_path))
|
model.load_state_dict(st.load_file(model_path))
|
||||||
|
|
||||||
assert torch.equal(model.lm_head.weight, model.embed_tokens.weight)
|
assert torch.equal(model.lm_head.weight, model.embed_tokens.weight)
|
||||||
@@ -112,8 +112,8 @@ def test_model_save_load_with_tie_weight(transformer_test_env):
|
|||||||
with open(config_path, "w") as f:
|
with open(config_path, "w") as f:
|
||||||
json.dump(config_data, f)
|
json.dump(config_data, f)
|
||||||
|
|
||||||
loaded_config = ModelConfig().load(config_path)
|
loaded_config = AutoRegressiveLMConfig.from_file(config_path)
|
||||||
model = Transformer(loaded_config)
|
model = AutoRegressiveLM(loaded_config)
|
||||||
model.load_state_dict(st.load_file(model_path))
|
model.load_state_dict(st.load_file(model_path))
|
||||||
|
|
||||||
assert torch.equal(model.lm_head.weight, model.embed_tokens.weight)
|
assert torch.equal(model.lm_head.weight, model.embed_tokens.weight)
|
||||||
|
|||||||
@@ -31,8 +31,8 @@ def create_train_config(
|
|||||||
device: str,
|
device: str,
|
||||||
strategy: str = "seq",
|
strategy: str = "seq",
|
||||||
n_epoch: int = 1,
|
n_epoch: int = 1,
|
||||||
batch_size: int = 2,
|
batch_per_device: int = 2,
|
||||||
accumulation_steps: int = 1,
|
grad_accum_steps: int = 1,
|
||||||
max_grad_norm: float = 1.0,
|
max_grad_norm: float = 1.0,
|
||||||
ckpt_interval: int = 5,
|
ckpt_interval: int = 5,
|
||||||
random_seed: int = 42,
|
random_seed: int = 42,
|
||||||
@@ -47,8 +47,8 @@ def create_train_config(
|
|||||||
device: Device type ("cuda" or "cpu")
|
device: Device type ("cuda" or "cpu")
|
||||||
strategy: Training strategy type (default: "seq")
|
strategy: Training strategy type (default: "seq")
|
||||||
n_epoch: Number of epochs (default: 1)
|
n_epoch: Number of epochs (default: 1)
|
||||||
batch_size: Batch size (default: 2)
|
batch_per_device: Batch size per device (default: 2)
|
||||||
accumulation_steps: Gradient accumulation steps (default: 1)
|
grad_accum_steps: Gradient accumulation steps (default: 1)
|
||||||
max_grad_norm: Maximum gradient norm for clipping (default: 1.0)
|
max_grad_norm: Maximum gradient norm for clipping (default: 1.0)
|
||||||
ckpt_interval: Checkpoint save interval in iterations (default: 5)
|
ckpt_interval: Checkpoint save interval in iterations (default: 5)
|
||||||
random_seed: Random seed for reproducibility (default: 42)
|
random_seed: Random seed for reproducibility (default: 42)
|
||||||
@@ -74,9 +74,9 @@ def create_train_config(
|
|||||||
scheduler_fn=scheduler_fn,
|
scheduler_fn=scheduler_fn,
|
||||||
ckpt_dir=test_dir,
|
ckpt_dir=test_dir,
|
||||||
n_epoch=n_epoch,
|
n_epoch=n_epoch,
|
||||||
batch_size=batch_size,
|
batch_per_device=batch_per_device,
|
||||||
ckpt_interval=ckpt_interval,
|
ckpt_interval=ckpt_interval,
|
||||||
accumulation_steps=accumulation_steps,
|
grad_accum_steps=grad_accum_steps,
|
||||||
max_grad_norm=max_grad_norm,
|
max_grad_norm=max_grad_norm,
|
||||||
random_seed=random_seed,
|
random_seed=random_seed,
|
||||||
device_type=device,
|
device_type=device,
|
||||||
|
|||||||
@@ -1,11 +1,130 @@
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.config.train_config import TrainConfig
|
from astrai.config.train_config import TrainConfig
|
||||||
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
from astrai.trainer.schedule import SchedulerFactory
|
from astrai.trainer.schedule import SchedulerFactory
|
||||||
from astrai.trainer.train_callback import TrainCallback
|
from astrai.trainer.train_callback import GradientCheckpointingCallback, TrainCallback
|
||||||
from astrai.trainer.trainer import Trainer
|
from astrai.trainer.trainer import Trainer
|
||||||
|
|
||||||
|
|
||||||
|
def test_gradient_checkpointing_enable_disable(test_model):
|
||||||
|
"""Enable wraps forward, _disable restores it."""
|
||||||
|
model = test_model["model"]
|
||||||
|
callback = GradientCheckpointingCallback(modules=[DecoderBlock])
|
||||||
|
|
||||||
|
originals = [layer.forward for layer in model.layers]
|
||||||
|
|
||||||
|
for layer in model.layers:
|
||||||
|
callback._enable(layer)
|
||||||
|
|
||||||
|
for layer in model.layers:
|
||||||
|
assert hasattr(layer, "_original_forward")
|
||||||
|
assert layer.forward is not originals[0]
|
||||||
|
|
||||||
|
for layer in model.layers:
|
||||||
|
callback._disable(layer)
|
||||||
|
|
||||||
|
for layer in model.layers:
|
||||||
|
assert not hasattr(layer, "_original_forward")
|
||||||
|
|
||||||
|
|
||||||
|
def test_gradient_checkpointing_empty_modules_noop(test_model):
|
||||||
|
"""modules=None should leave forwards untouched."""
|
||||||
|
model = test_model["model"]
|
||||||
|
callback = GradientCheckpointingCallback()
|
||||||
|
|
||||||
|
originals = [layer.forward for layer in model.layers]
|
||||||
|
|
||||||
|
for layer in model.layers:
|
||||||
|
callback._enable(layer)
|
||||||
|
|
||||||
|
for layer, orig in zip(model.layers, originals):
|
||||||
|
assert layer.forward is orig
|
||||||
|
|
||||||
|
|
||||||
|
def test_gradient_checkpointing_forward_unchanged(test_model):
|
||||||
|
"""Forward output unchanged after patching (no_grad)."""
|
||||||
|
model = test_model["model"]
|
||||||
|
device = test_model["device"]
|
||||||
|
callback = GradientCheckpointingCallback(modules=[DecoderBlock])
|
||||||
|
|
||||||
|
input_ids = torch.randint(0, 1000, (2, 32)).to(device)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
ref = model(input_ids)["logits"].clone()
|
||||||
|
|
||||||
|
for layer in model.layers:
|
||||||
|
callback._enable(layer)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
out = model(input_ids)["logits"]
|
||||||
|
|
||||||
|
assert torch.equal(ref, out)
|
||||||
|
|
||||||
|
|
||||||
|
def test_gradient_checkpointing_backward(test_model):
|
||||||
|
"""backward passes gradients through checkpointed layers."""
|
||||||
|
model = test_model["model"]
|
||||||
|
device = test_model["device"]
|
||||||
|
callback = GradientCheckpointingCallback(modules=[DecoderBlock])
|
||||||
|
|
||||||
|
for layer in model.layers:
|
||||||
|
callback._enable(layer)
|
||||||
|
|
||||||
|
input_ids = torch.randint(0, 1000, (2, 32)).to(device)
|
||||||
|
target_ids = torch.randint(0, 1000, (2, 32)).to(device)
|
||||||
|
|
||||||
|
logits = model(input_ids)["logits"]
|
||||||
|
loss = torch.nn.functional.cross_entropy(
|
||||||
|
logits.flatten(0, 1).float(), target_ids.flatten()
|
||||||
|
)
|
||||||
|
loss.backward()
|
||||||
|
|
||||||
|
for name, param in model.named_parameters():
|
||||||
|
if param.requires_grad:
|
||||||
|
assert param.grad is not None, f"{name} gradient is None"
|
||||||
|
|
||||||
|
for layer in model.layers:
|
||||||
|
callback._disable(layer)
|
||||||
|
|
||||||
|
model.zero_grad()
|
||||||
|
for name, p in model.named_parameters():
|
||||||
|
assert p.grad is None or p.grad.sum().item() == 0, f"{name} grad not zeroed"
|
||||||
|
|
||||||
|
|
||||||
|
def test_gradient_checkpointing_trainer_integration(base_test_env, random_dataset):
|
||||||
|
"""Gradient checkpointing runs end-to-end via Trainer."""
|
||||||
|
|
||||||
|
def optimizer_fn(model):
|
||||||
|
return torch.optim.AdamW(model.parameters())
|
||||||
|
|
||||||
|
def scheduler_fn(optim):
|
||||||
|
return SchedulerFactory.create(
|
||||||
|
optim, "cosine", warmup_steps=10, lr_decay_steps=10, min_rate=0.05
|
||||||
|
)
|
||||||
|
|
||||||
|
train_config = TrainConfig(
|
||||||
|
model=base_test_env["model"],
|
||||||
|
strategy="seq",
|
||||||
|
dataset=random_dataset,
|
||||||
|
optimizer_fn=optimizer_fn,
|
||||||
|
scheduler_fn=scheduler_fn,
|
||||||
|
ckpt_dir=base_test_env["test_dir"],
|
||||||
|
n_epoch=1,
|
||||||
|
batch_per_device=2,
|
||||||
|
ckpt_interval=3,
|
||||||
|
grad_accum_steps=1,
|
||||||
|
max_grad_norm=1.0,
|
||||||
|
random_seed=42,
|
||||||
|
device_type=base_test_env["device"],
|
||||||
|
gradient_checkpointing_modules=[DecoderBlock],
|
||||||
|
)
|
||||||
|
|
||||||
|
trainer = Trainer(train_config)
|
||||||
|
trainer.train()
|
||||||
|
# no crash = callback correctly enabled/disabled
|
||||||
|
|
||||||
|
|
||||||
def test_callback_integration(base_test_env, random_dataset):
|
def test_callback_integration(base_test_env, random_dataset):
|
||||||
"""Test that all callbacks are properly integrated"""
|
"""Test that all callbacks are properly integrated"""
|
||||||
|
|
||||||
@@ -25,9 +144,9 @@ def test_callback_integration(base_test_env, random_dataset):
|
|||||||
scheduler_fn=scheduler_fn,
|
scheduler_fn=scheduler_fn,
|
||||||
ckpt_dir=base_test_env["test_dir"],
|
ckpt_dir=base_test_env["test_dir"],
|
||||||
n_epoch=1,
|
n_epoch=1,
|
||||||
batch_size=2,
|
batch_per_device=2,
|
||||||
ckpt_interval=3,
|
ckpt_interval=3,
|
||||||
accumulation_steps=1,
|
grad_accum_steps=1,
|
||||||
max_grad_norm=1.0,
|
max_grad_norm=1.0,
|
||||||
random_seed=42,
|
random_seed=42,
|
||||||
device_type=base_test_env["device"],
|
device_type=base_test_env["device"],
|
||||||
|
|||||||
@@ -28,9 +28,9 @@ def test_early_stopping_simulation(base_test_env, early_stopping_dataset):
|
|||||||
dataset=early_stopping_dataset,
|
dataset=early_stopping_dataset,
|
||||||
ckpt_dir=base_test_env["test_dir"],
|
ckpt_dir=base_test_env["test_dir"],
|
||||||
n_epoch=2,
|
n_epoch=2,
|
||||||
batch_size=2,
|
batch_per_device=2,
|
||||||
ckpt_interval=1,
|
ckpt_interval=1,
|
||||||
accumulation_steps=2,
|
grad_accum_steps=2,
|
||||||
random_seed=np.random.randint(1e4),
|
random_seed=np.random.randint(1e4),
|
||||||
device_type=base_test_env["device"],
|
device_type=base_test_env["device"],
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -72,6 +72,7 @@ def test_schedule_factory_random_configs():
|
|||||||
|
|
||||||
# Test scheduler step functionality
|
# Test scheduler step functionality
|
||||||
initial_lr = scheduler.get_last_lr()
|
initial_lr = scheduler.get_last_lr()
|
||||||
|
optimizer.step()
|
||||||
scheduler.step()
|
scheduler.step()
|
||||||
new_lr = scheduler.get_last_lr()
|
new_lr = scheduler.get_last_lr()
|
||||||
|
|
||||||
@@ -112,6 +113,7 @@ def test_schedule_factory_edge_cases():
|
|||||||
|
|
||||||
# Test multiple steps
|
# Test multiple steps
|
||||||
for _ in range(10):
|
for _ in range(10):
|
||||||
|
optimizer.step()
|
||||||
scheduler.step()
|
scheduler.step()
|
||||||
|
|
||||||
|
|
||||||
@@ -136,6 +138,7 @@ def test_schedule_factory_state_persistence():
|
|||||||
|
|
||||||
# Take a few steps
|
# Take a few steps
|
||||||
for _ in range(5):
|
for _ in range(5):
|
||||||
|
optimizer.step()
|
||||||
scheduler.step()
|
scheduler.step()
|
||||||
|
|
||||||
# Save state
|
# Save state
|
||||||
|
|||||||
@@ -7,45 +7,45 @@ def test_different_batch_sizes(base_test_env, random_dataset, train_config_facto
|
|||||||
"""Test training with different batch sizes"""
|
"""Test training with different batch sizes"""
|
||||||
batch_sizes = [1, 2, 4, 8]
|
batch_sizes = [1, 2, 4, 8]
|
||||||
|
|
||||||
for batch_size in batch_sizes:
|
for batch_per_device in batch_sizes:
|
||||||
train_config = train_config_factory(
|
train_config = train_config_factory(
|
||||||
model=base_test_env["model"],
|
model=base_test_env["model"],
|
||||||
dataset=random_dataset,
|
dataset=random_dataset,
|
||||||
test_dir=base_test_env["test_dir"],
|
test_dir=base_test_env["test_dir"],
|
||||||
device=base_test_env["device"],
|
device=base_test_env["device"],
|
||||||
batch_size=batch_size,
|
batch_per_device=batch_per_device,
|
||||||
)
|
)
|
||||||
|
|
||||||
assert train_config.batch_size == batch_size
|
assert train_config.batch_per_device == batch_per_device
|
||||||
|
|
||||||
|
|
||||||
def test_gradient_accumulation(base_test_env, random_dataset, train_config_factory):
|
def test_gradient_accumulation(base_test_env, random_dataset, train_config_factory):
|
||||||
"""Test training with different gradient accumulation steps"""
|
"""Test training with different gradient accumulation steps"""
|
||||||
accumulation_steps_list = [1, 2, 4]
|
grad_accum_steps_list = [1, 2, 4]
|
||||||
|
|
||||||
for accumulation_steps in accumulation_steps_list:
|
for grad_accum_steps in grad_accum_steps_list:
|
||||||
train_config = train_config_factory(
|
train_config = train_config_factory(
|
||||||
model=base_test_env["model"],
|
model=base_test_env["model"],
|
||||||
dataset=random_dataset,
|
dataset=random_dataset,
|
||||||
test_dir=base_test_env["test_dir"],
|
test_dir=base_test_env["test_dir"],
|
||||||
device=base_test_env["device"],
|
device=base_test_env["device"],
|
||||||
batch_size=2,
|
batch_per_device=2,
|
||||||
accumulation_steps=accumulation_steps,
|
grad_accum_steps=grad_accum_steps,
|
||||||
)
|
)
|
||||||
|
|
||||||
trainer = Trainer(train_config)
|
trainer = Trainer(train_config)
|
||||||
trainer.train()
|
trainer.train()
|
||||||
|
|
||||||
assert train_config.accumulation_steps == accumulation_steps
|
assert train_config.grad_accum_steps == grad_accum_steps
|
||||||
|
|
||||||
|
|
||||||
def test_memory_efficient_training(base_test_env, random_dataset, train_config_factory):
|
def test_memory_efficient_training(base_test_env, random_dataset, train_config_factory):
|
||||||
"""Test training with memory-efficient configurations"""
|
"""Test training with memory-efficient configurations"""
|
||||||
# Test with smaller batch sizes and gradient checkpointing
|
# Test with smaller batch sizes and gradient checkpointing
|
||||||
small_batch_configs = [
|
small_batch_configs = [
|
||||||
{"batch_size": 1, "accumulation_steps": 8},
|
{"batch_per_device": 1, "grad_accum_steps": 8},
|
||||||
{"batch_size": 2, "accumulation_steps": 4},
|
{"batch_per_device": 2, "grad_accum_steps": 4},
|
||||||
{"batch_size": 4, "accumulation_steps": 2},
|
{"batch_per_device": 4, "grad_accum_steps": 2},
|
||||||
]
|
]
|
||||||
|
|
||||||
for config in small_batch_configs:
|
for config in small_batch_configs:
|
||||||
@@ -54,8 +54,9 @@ def test_memory_efficient_training(base_test_env, random_dataset, train_config_f
|
|||||||
dataset=random_dataset,
|
dataset=random_dataset,
|
||||||
test_dir=base_test_env["test_dir"],
|
test_dir=base_test_env["test_dir"],
|
||||||
device=base_test_env["device"],
|
device=base_test_env["device"],
|
||||||
batch_size=config["batch_size"],
|
batch_per_device=config["batch_per_device"],
|
||||||
accumulation_steps=config["accumulation_steps"],
|
grad_accum_steps=config["grad_accum_steps"],
|
||||||
)
|
)
|
||||||
|
|
||||||
assert train_config.accumulation_steps == config["accumulation_steps"]
|
assert train_config.grad_accum_steps == config["grad_accum_steps"]
|
||||||
|
assert train_config.batch_per_device == config["batch_per_device"]
|
||||||
|
|||||||
Reference in New Issue
Block a user