Compare commits
167
Commits
v1.3.7
..
663ef900fc
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
663ef900fc | ||
|
|
7d478a54db | ||
|
|
f3eaaef842 | ||
|
|
d655b65027 | ||
|
|
31c22dc043 | ||
|
|
17127f8b3c | ||
|
|
d7695b40e3 | ||
|
|
fc62890e70 | ||
|
|
f433672140 | ||
|
|
7e1e5b6e6a | ||
|
|
553a42702d | ||
|
|
b133fc9c07 | ||
|
|
b33250dc28 | ||
|
|
a74e5b91a3 | ||
|
|
28886e4241 | ||
|
|
9d3ccfdffc | ||
|
|
a24a7b4da5 | ||
|
|
f7df02f9a3 | ||
|
|
ee450686f3 | ||
|
|
2565755e45 | ||
|
|
d08a92c7bd | ||
|
|
a1ea26d367 | ||
|
|
c17aa0dc54 | ||
|
|
b12b24eadc | ||
|
|
cd14d53707 | ||
|
|
e220413035 | ||
|
|
84ed2327f5 | ||
|
|
b14f301730 | ||
|
|
0654b4b916 | ||
|
|
1f0be382ad | ||
|
|
bb175fda91 | ||
|
|
13998da15a | ||
|
|
57729fd92d | ||
|
|
2c7a71a9c0 | ||
|
|
3e0007fc91 | ||
|
|
b092316385 | ||
|
|
9bcd696580 | ||
|
|
8f89c82d55 | ||
|
|
21871197d7 | ||
|
|
4c35d36146 | ||
|
|
9aca62c26c | ||
|
|
b5cdea98ad | ||
|
|
69fecaf387 | ||
|
|
fd6d25ad86 | ||
|
|
2c3cef1c87 | ||
|
|
89ece26c25 | ||
|
|
2c0b5d0b5e | ||
|
|
a4ae7d17fb | ||
|
|
8a8550184f | ||
|
|
b8b439b713 | ||
|
|
41cd40363a | ||
|
|
d923ebe38d | ||
|
|
29b0423c4e | ||
|
|
88f8dca2c2 | ||
|
|
9027fdc546 | ||
|
|
cbd140340d | ||
|
|
988e01314d | ||
|
|
7ba43a7c6f | ||
|
|
dea59f7e1d | ||
|
|
85dc771460 | ||
|
|
2c5629b81d | ||
|
|
841a582b28 | ||
|
|
c8567a6f65 | ||
|
|
8035be9b1f | ||
|
|
e9b03f4fca | ||
|
|
fd65b9bc23 | ||
|
|
9ebaea840f | ||
|
|
6adc221c10 | ||
|
|
9e63cb9ed0 | ||
|
|
4225518cf3 | ||
|
|
c50adbaac0 | ||
|
|
536dbc0c9a | ||
|
|
4af7acd449 | ||
|
|
53ed52b4b8 | ||
|
|
f1cc7cedce | ||
|
|
ddc4bd1cf6 | ||
|
|
cc36530c73 | ||
|
|
11fa807cfc | ||
|
|
bcdd93e0eb | ||
|
|
579b8c3129 | ||
|
|
d7da51569f | ||
|
|
e8e228d035 | ||
|
|
2579658e15 | ||
|
|
f0cd0134c6 | ||
|
|
abb96996f8 | ||
|
|
bbe6ff2d8f | ||
|
|
db9b39b084 | ||
|
|
849e1e00a3 | ||
|
|
5416c2e8fb | ||
|
|
599a51f4f7 | ||
|
|
17d6eaa2f2 | ||
|
|
2d908639e9 | ||
|
|
c7158418dd | ||
|
|
4d3c9341c1 | ||
|
|
4e508afa2d | ||
|
|
8999ca89b8 | ||
|
|
1adca39cd8 | ||
|
|
204873fa2f | ||
|
|
a5c1de6b1b | ||
|
|
27524ad085 | ||
|
|
27d1921d9c | ||
|
|
70c0e5de90 | ||
|
|
dfb151537b | ||
|
|
500c605fad | ||
|
|
dc9faca3b1 | ||
|
|
aabb0d83e9 | ||
|
|
44579ea6dc | ||
|
|
0f1fcb079f | ||
|
|
84d4769163 | ||
|
|
bf09a35c95 | ||
|
|
6715461a36 | ||
|
|
b4587c5d08 | ||
|
|
88ec63121d | ||
|
|
01d2da2893 | ||
|
|
25d4ea3f91 | ||
|
|
39985840c7 | ||
|
|
b1adc40cfb | ||
|
|
7348bac6ab | ||
|
|
8ab7564d02 | ||
|
|
d096b6e29e | ||
|
|
d88a41f8f1 | ||
|
|
376e9eba80 | ||
|
|
a62c2e11a2 | ||
|
|
a4e5a8c81c | ||
|
|
3e234c46f6 | ||
|
|
7a04b1f8ce | ||
|
|
a30e3d5114 | ||
|
|
1818d06576 | ||
|
|
4e8d1ee24e | ||
|
|
fec376b0dd | ||
|
|
a2512f8a5a | ||
|
|
457e16ea3c | ||
|
|
daf627a6de | ||
|
|
445378667f | ||
|
|
6ae1828449 | ||
|
|
e7b18b7c03 | ||
|
|
9e31d4ef2b | ||
|
|
52aa4d01d5 | ||
|
|
986be957ec | ||
|
|
cf9c60841b | ||
|
|
31bc7f5c2a | ||
|
|
3057741de9 | ||
|
|
acd1103bd0 | ||
|
|
dc7d2cfbca | ||
|
|
b36a78c612 | ||
|
|
985d940db6 | ||
|
|
5e73ca20aa | ||
|
|
438dc10391 | ||
|
|
615ba5d8ef | ||
|
|
02a7cb9fa0 | ||
|
|
9fe2121743 | ||
|
|
0422d6d38e | ||
|
|
9b416c1bbb | ||
|
|
d6899100ac | ||
|
|
0deee48602 | ||
|
|
746a1475b2 | ||
|
|
01ce1fb9e3 | ||
|
|
14f83cbdac | ||
|
|
dbe5891201 | ||
|
|
2a65c3314c | ||
|
|
1c2ff05a6d | ||
|
|
31ae2deeba | ||
|
|
69207e2c57 | ||
|
|
138c5bcc08 | ||
|
|
a923e0a23a | ||
|
|
f521a30b22 | ||
|
|
d4451f6afb |
@@ -0,0 +1,71 @@
|
||||
name: Release
|
||||
|
||||
on:
|
||||
push:
|
||||
tags:
|
||||
- "v*"
|
||||
|
||||
jobs:
|
||||
build-pure:
|
||||
name: Build pure-Python wheel
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Build wheel (no CUDA)
|
||||
run: |
|
||||
pip wheel . --no-deps -w dist/
|
||||
|
||||
- uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: pure-wheel
|
||||
path: dist/*.whl
|
||||
|
||||
build-cuda-linux:
|
||||
name: Build CUDA wheel (Linux)
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Install torch (CUDA 12.8)
|
||||
run: |
|
||||
pip install torch --index-url https://download.pytorch.org/whl/cu128
|
||||
|
||||
- name: Setup CUDA
|
||||
uses: Jimver/cuda-toolkit@v0.2.35
|
||||
with:
|
||||
cuda: "12.8.0"
|
||||
|
||||
- name: Build wheel (with CUDA kernels)
|
||||
run: |
|
||||
CSRC_KERNELS=true pip wheel . --no-deps --no-build-isolation -w dist/
|
||||
|
||||
- uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: cuda-wheel-linux
|
||||
path: dist/*.whl
|
||||
|
||||
release:
|
||||
name: Attach wheels to release
|
||||
needs: [build-pure, build-cuda-linux]
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: write
|
||||
steps:
|
||||
- uses: actions/download-artifact@v4
|
||||
with:
|
||||
pattern: "*-wheel"
|
||||
merge-multiple: true
|
||||
|
||||
- name: Create release & upload assets
|
||||
uses: softprops/action-gh-release@v2
|
||||
with:
|
||||
files: ./*.whl
|
||||
tag_name: ${{ github.ref_name }}
|
||||
generate_release_notes: true
|
||||
+16
-3
@@ -5,8 +5,16 @@
|
||||
!*/
|
||||
|
||||
# Allow specific file types and root files
|
||||
!*.py
|
||||
!*.sh
|
||||
!astrai/**/*.py
|
||||
!scripts/**/*.py
|
||||
!tests/**/*.py
|
||||
!csrc/**/*.py
|
||||
|
||||
!csrc/**/*.cu
|
||||
!csrc/**/*.h
|
||||
!csrc/**/*.cuh
|
||||
|
||||
!scripts/**/*.sh
|
||||
|
||||
# Allow GitHub files
|
||||
!/.github/**
|
||||
@@ -20,4 +28,9 @@
|
||||
!/CONTRIBUTING.md
|
||||
!/LICENSE
|
||||
!/pyproject.toml
|
||||
!/README.md
|
||||
!/README.md
|
||||
# Allow extension modules (only source .py)
|
||||
!/astrai/extension/**/*.py
|
||||
|
||||
# Allow build files
|
||||
!/setup.py
|
||||
|
||||
+1
-1
@@ -5,7 +5,7 @@ Thank you for your interest in contributing! This document provides step-by-step
|
||||
## Quick Start
|
||||
|
||||
```bash
|
||||
git clone https://github.com/your-username/AstrAI.git
|
||||
git clone https://github.com/ViperEkura/AstrAI.git
|
||||
cd AstrAI
|
||||
pip install -e ".[dev]" # install with dev dependencies (pytest, ruff)
|
||||
```
|
||||
|
||||
+1
-1
@@ -23,7 +23,7 @@ COPY astrai/ ./astrai/
|
||||
COPY pyproject.toml .
|
||||
RUN pip install --no-cache-dir --upgrade pip \
|
||||
&& pip install --no-cache-dir . \
|
||||
--extra-index-url https://download.pytorch.org/whl/cu126
|
||||
--extra-index-url https://download.pytorch.org/whl/cu128
|
||||
|
||||
# Production stage
|
||||
FROM ubuntu:24.04 AS production
|
||||
|
||||
@@ -9,9 +9,9 @@
|
||||
<div align="center">
|
||||
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
||||
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license">
|
||||
<img src="https://img.shields.io/github/v/release/ViperEkura/AstrAI?color=76bad9" alt="release">
|
||||
<img src="https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fapi.github.com%2Frepos%2FViperEkura%2FAstrAI&query=%24.stargazers_count&label=stars&suffix=%20stars&color=76bad9" alt="stars">
|
||||
<img src="https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fapi.github.com%2Frepos%2FViperEkura%2FAstrAI&query=%24.forks_count&label=forks&suffix=%20forks&color=76bad9" alt="forks">
|
||||
<img src="https://img.shields.io/github/v/tag/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release">
|
||||
<img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars">
|
||||
<img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks">
|
||||
</div>
|
||||
<br>
|
||||
|
||||
@@ -20,7 +20,7 @@
|
||||
<a href="assets/docs/README-zh-CN.md">中文</a> •
|
||||
<a href="https://github.com/ViperEkura/AstrAI/issues">Issue Tracker</a> •
|
||||
<a href="https://github.com/ViperEkura/AstrAI/discussions">Discussions</a> •
|
||||
<a href="https://huggingface.co/ViperEk/">HuggingFace</a>
|
||||
<a href="https://huggingface.co/ViperEkura">HuggingFace</a>
|
||||
</div>
|
||||
|
||||
<br>
|
||||
@@ -28,7 +28,8 @@
|
||||
## 📖 Table of Contents
|
||||
|
||||
- [Features](#features)
|
||||
- [Quick Start](#quick-start)
|
||||
- [Getting Started](#getting-started)
|
||||
- [Demo](#demo)
|
||||
- [Documentation](#documentation)
|
||||
- [Contributing](#contributing)
|
||||
- [Community](#community)
|
||||
@@ -49,39 +50,51 @@
|
||||
- 🤗 **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.
|
||||
|
||||
### Quick Start
|
||||
### Getting Started
|
||||
|
||||
#### Installation
|
||||
End-to-end walkthrough in 5 steps:
|
||||
|
||||
**1. Install**
|
||||
|
||||
```bash
|
||||
git clone https://github.com/ViperEkura/AstrAI.git
|
||||
cd AstrAI
|
||||
pip install -e .
|
||||
pip install -e . # pure PyTorch (no CUDA kernels)
|
||||
# CSRC_KERNELS=true pip install -e . --no-build-isolation # optional: fused CUDA kernels
|
||||
# pip install -e ".[dev]" # dev dependencies (pytest, ruff)
|
||||
```
|
||||
|
||||
For development dependencies:
|
||||
**2. Download model**
|
||||
|
||||
```bash
|
||||
pip install -e ".[dev]"
|
||||
python scripts/demo/download.py # downloads 1B checkpoint to params/
|
||||
```
|
||||
|
||||
#### Download Pre-trained Model
|
||||
**3. Preprocess data**
|
||||
|
||||
Download pre-trained model weights (1B bilingual checkpoint) to `params/`:
|
||||
Create `pretrain.json` (preprocessing config for `seq` strategy):
|
||||
|
||||
```json
|
||||
{
|
||||
"version": 1,
|
||||
"input": {"sections": [{"field": "text", "action": "train"}]},
|
||||
"preprocessing": {"max_seq_len": 2048},
|
||||
"output": {"storage_format": "bin"}
|
||||
}
|
||||
```
|
||||
|
||||
```bash
|
||||
python scripts/demo/download.py
|
||||
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c pretrain.json
|
||||
```
|
||||
|
||||
Or download manually from [HuggingFace](https://huggingface.co/ViperEk/KHAOSZ) into `params/`.
|
||||
|
||||
#### Train a Model
|
||||
**4. Train**
|
||||
|
||||
```bash
|
||||
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||
|
||||
nohup python scripts/tools/train.py \
|
||||
--nprocs=4 \
|
||||
--parallel_mode=ddp \
|
||||
--train_type=seq \
|
||||
--data_root_path=/path/to/dataset \
|
||||
--param_path=/path/to/model \
|
||||
@@ -90,9 +103,7 @@ nohup python scripts/tools/train.py \
|
||||
--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 \
|
||||
--weight_decay=0.1 \
|
||||
--window_size=2048 \
|
||||
--ckpt_interval=10000 \
|
||||
--ckpt_dir=./checkpoint \
|
||||
@@ -101,15 +112,54 @@ nohup python scripts/tools/train.py \
|
||||
> out.log 2> err.log &
|
||||
```
|
||||
|
||||
Full reference at [Parameter Guide](assets/docs/params.md).
|
||||
**5. Serve & query**
|
||||
|
||||
#### Generate Text
|
||||
```bash
|
||||
# Terminal 1: start server
|
||||
python scripts/tools/server.py --param_path ./params --device cuda
|
||||
|
||||
# Terminal 2: query
|
||||
curl http://localhost:8000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||
```
|
||||
|
||||
### Demo
|
||||
|
||||
Check out the demos in the `scripts/demo/` folder:
|
||||
|
||||
```bash
|
||||
# Download model weights (required before running demos)
|
||||
python scripts/demo/download.py # model → params/
|
||||
|
||||
# Interactive streaming chat (multi-turn, maintains history)
|
||||
python scripts/demo/stream_chat.py
|
||||
# Type your message after >>, type !exit to quit
|
||||
|
||||
# Batch generation (5 hardcoded prompts, non-streaming)
|
||||
python scripts/demo/generate_batch.py
|
||||
|
||||
# Single-prompt autoregressive streaming
|
||||
python scripts/demo/generate_ar.py
|
||||
```
|
||||
|
||||
All generation demos use `temperature=0.8`, `top_p=0.95`, `top_k=50`, `max_tokens=2048` by default and require `params/` to contain model weights (run `download.py` first).
|
||||
|
||||
Watch a video walkthrough on [bilibili](https://www.bilibili.com/video/BV1fuLB6yEj6).
|
||||
|
||||
---
|
||||
|
||||
See [Documentation](#documentation) for full references beyond the examples above.
|
||||
|
||||
#### Text Generation
|
||||
|
||||
Batch generation from a JSONL file:
|
||||
|
||||
```bash
|
||||
python scripts/tools/generate.py \
|
||||
--param_path /path/to/model \
|
||||
--input_json_file /path/to/input.json \
|
||||
--output_json_file /path/to/output.json
|
||||
--param_path ./params \
|
||||
--input_json_file input.jsonl \
|
||||
--output_json_file output.jsonl
|
||||
```
|
||||
|
||||
#### Docker
|
||||
@@ -123,9 +173,6 @@ docker build -t astrai:latest .
|
||||
# Run with GPU support
|
||||
docker run --gpus all -it astrai:latest
|
||||
|
||||
# Run with specific GPUs
|
||||
docker run --gpus '"device=0,1"' -it astrai:latest
|
||||
|
||||
# Run inference server
|
||||
docker run --gpus all -p 8000:8000 astrai:latest \
|
||||
python -m scripts.tools.server --port 8000 --device cuda
|
||||
@@ -142,88 +189,42 @@ docker compose --profile cpu up -d
|
||||
|
||||
> **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`.
|
||||
|
||||
#### Start HTTP Server
|
||||
#### HTTP API Examples
|
||||
|
||||
Start the inference server with OpenAI and Anthropic-compatible HTTP API:
|
||||
Additional request examples beyond the [Getting Started](#getting-started) flow:
|
||||
|
||||
```bash
|
||||
python -m scripts.tools.server --port 8000 --device cuda
|
||||
```
|
||||
|
||||
Make requests:
|
||||
|
||||
```bash
|
||||
# OpenAI-compatible
|
||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"max_tokens": 512
|
||||
}'
|
||||
|
||||
# OpenAI-compatible streaming
|
||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"messages": [{"role": "user", "content": "Tell a story"}],
|
||||
"stream": true,
|
||||
"max_tokens": 500
|
||||
}'
|
||||
-d '{"messages":[{"role":"user","content":"Tell a story"}],"stream":true,"max_tokens":500}'
|
||||
|
||||
# Anthropic-compatible
|
||||
curl -X POST http://localhost:8000/v1/messages \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "astrai",
|
||||
"system": "You are a helpful assistant.",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"max_tokens": 512
|
||||
}'
|
||||
-d '{"model":"astrai","system":"You are a helpful assistant.","messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||
|
||||
# Anthropic-compatible streaming with stop sequences
|
||||
curl -X POST http://localhost:8000/v1/messages \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "astrai",
|
||||
"messages": [{"role": "user", "content": "Write a story"}],
|
||||
"max_tokens": 500,
|
||||
"stream": true,
|
||||
"stop_sequences": ["The end"]
|
||||
}'
|
||||
-d '{"model":"astrai","messages":[{"role":"user","content":"Write a story"}],"max_tokens":500,"stream":true,"stop_sequences":["The end"]}'
|
||||
|
||||
# Health check
|
||||
curl http://localhost:8000/health
|
||||
```
|
||||
|
||||
#### Demo
|
||||
|
||||
Check out the demos in the `scripts/demo/` folder:
|
||||
|
||||
```bash
|
||||
# Download pre‑processed data (required before running demos)
|
||||
python scripts/demo/download.py
|
||||
|
||||
# Interactive streaming chat
|
||||
python scripts/demo/stream_chat.py
|
||||
|
||||
# Batch generation
|
||||
python scripts/demo/generate_batch.py
|
||||
|
||||
# Auto‑regressive generation
|
||||
python scripts/demo/generate_ar.py
|
||||
```
|
||||
|
||||
Watch a video walkthrough on [bilibili](https://www.bilibili.com/video/BV1fuLB6yEj6).
|
||||
See [Inference Guide](assets/docs/inference.md) for SSE streaming format, error codes, and stats endpoint.
|
||||
|
||||
### Documentation
|
||||
|
||||
| Document | Description |
|
||||
|----------|-------------|
|
||||
| [Parameter Guide](./assets/docs/params.md) | Training & inference parameters |
|
||||
| [CLI Reference](./assets/docs/params.md) | Parameters for all CLI tools (train, server, generate, preprocess) |
|
||||
| [Architecture](./assets/docs/architecture.md) | System architecture, class diagram & design patterns |
|
||||
| [Training](./assets/docs/training.md) | Training loop, strategies & formulas |
|
||||
| [Inference](./assets/docs/inference.md) | KVCache, continuous batching, sampling & HTTP API |
|
||||
| [Data Flow](./assets/docs/dataflow.md) | Data pipeline, storage backends & dataset architecture |
|
||||
| [Preprocessing](./assets/docs/preprocessing.md) | Declarative JSON-driven data preprocessing |
|
||||
|
||||
### Contributing
|
||||
|
||||
@@ -240,7 +241,7 @@ For major changes, please open an issue first to discuss what you would like to
|
||||
|
||||
- **GitHub Issues**: [Issue Tracker](https://github.com/ViperEkura/AstrAI/issues)
|
||||
- **Discussions**: [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions)
|
||||
- **HuggingFace**: [Model Hub](https://huggingface.co/ViperEk)
|
||||
- **HuggingFace**: [Model Hub](https://huggingface.co/ViperEkura)
|
||||
|
||||
### License
|
||||
|
||||
|
||||
+82
-81
@@ -15,9 +15,9 @@
|
||||
<div align="center">
|
||||
<img src="https://img.shields.io/badge/python-3.12+-blue.svg" alt="python">
|
||||
<img src="https://img.shields.io/badge/license-GPL--3.0-blue.svg" alt="license">
|
||||
<img src="https://img.shields.io/github/v/release/ViperEkura/AstrAI?color=76bad9" alt="release">
|
||||
<img src="https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fapi.github.com%2Frepos%2FViperEkura%2FAstrAI&query=%24.stargazers_count&label=stars&suffix=%20stars&color=76bad9" alt="stars">
|
||||
<img src="https://img.shields.io/badge/dynamic/json?url=https%3A%2F%2Fapi.github.com%2Frepos%2FViperEkura%2FAstrAI&query=%24.forks_count&label=forks&suffix=%20forks&color=76bad9" alt="forks">
|
||||
<img src="https://img.shields.io/github/v/tag/ViperEkura/AstrAI?label=Release&color=76bad9" alt="release">
|
||||
<img src="https://img.shields.io/github/stars/ViperEkura/AstrAI?style=flat&label=Stars&color=76bad9" alt="stars">
|
||||
<img src="https://img.shields.io/github/forks/ViperEkura/AstrAI?style=flat&label=Forks&color=76bad9" alt="forks">
|
||||
</div>
|
||||
|
||||
<br>
|
||||
@@ -27,14 +27,15 @@
|
||||
<a href="#chinese">中文</a> •
|
||||
<a href="https://github.com/ViperEkura/AstrAI/issues">问题追踪</a> •
|
||||
<a href="https://github.com/ViperEkura/AstrAI/discussions">讨论区</a> •
|
||||
<a href="https://huggingface.co/ViperEk">HuggingFace</a>
|
||||
<a href="https://huggingface.co/ViperEkura">HuggingFace</a>
|
||||
</div>
|
||||
<br>
|
||||
|
||||
## 📖 目录
|
||||
|
||||
- [特性](#特性)
|
||||
- [快速开始](#快速开始)
|
||||
- [快速上手](#快速上手)
|
||||
- [演示](#演示)
|
||||
- [文档](#文档)
|
||||
- [贡献](#贡献)
|
||||
- [社区](#社区)
|
||||
@@ -55,39 +56,51 @@
|
||||
- 🤗 **HuggingFace 风格 API**: 类 HuggingFace 的 AutoModel/AutoTokenizer 接口,方便加载模型和分词器。
|
||||
- 🔌 **双 API 兼容**: 同时支持 OpenAI 和 Anthropic 聊天补全 API,开箱即用。
|
||||
|
||||
### 快速开始
|
||||
### 快速上手
|
||||
|
||||
#### 安装
|
||||
端到端演示,只需 5 步:
|
||||
|
||||
**1. 安装**
|
||||
|
||||
```bash
|
||||
git clone https://github.com/ViperEkura/AstrAI.git
|
||||
cd AstrAI
|
||||
pip install -e .
|
||||
pip install -e . # 纯 PyTorch(不含 CUDA 内核)
|
||||
# CSRC_KERNELS=true pip install -e . --no-build-isolation # 可选:融合 CUDA 内核加速
|
||||
# pip install -e ".[dev]" # 可选:开发依赖(pytest, ruff)
|
||||
```
|
||||
|
||||
安装开发依赖:
|
||||
**2. 下载模型**
|
||||
|
||||
```bash
|
||||
pip install -e ".[dev]"
|
||||
python scripts/demo/download.py # 下载 1B 检查点到 params/
|
||||
```
|
||||
|
||||
#### 下载预训练模型
|
||||
**3. 预处理数据**
|
||||
|
||||
下载预训练模型权重(1B 双语检查点)到 `params/` 目录:
|
||||
创建 `pretrain.json`(`seq` 策略的预处理配置):
|
||||
|
||||
```json
|
||||
{
|
||||
"version": 1,
|
||||
"input": {"sections": [{"field": "text", "action": "train"}]},
|
||||
"preprocessing": {"max_seq_len": 2048},
|
||||
"output": {"storage_format": "bin"}
|
||||
}
|
||||
```
|
||||
|
||||
```bash
|
||||
python scripts/demo/download.py
|
||||
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c pretrain.json
|
||||
```
|
||||
|
||||
或从 [HuggingFace](https://huggingface.co/ViperEk/KHAOSZ) 手动下载放入 `params/`。
|
||||
|
||||
#### 训练模型
|
||||
**4. 训练**
|
||||
|
||||
```bash
|
||||
export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||
|
||||
nohup python scripts/tools/train.py \
|
||||
--nprocs=4 \
|
||||
--parallel_mode=ddp \
|
||||
--train_type=seq \
|
||||
--data_root_path=/path/to/dataset \
|
||||
--param_path=/path/to/model \
|
||||
@@ -96,9 +109,7 @@ nohup python scripts/tools/train.py \
|
||||
--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 \
|
||||
--weight_decay=0.1 \
|
||||
--window_size=2048 \
|
||||
--ckpt_interval=10000 \
|
||||
--ckpt_dir=./checkpoint \
|
||||
@@ -107,15 +118,54 @@ nohup python scripts/tools/train.py \
|
||||
> out.log 2> err.log &
|
||||
```
|
||||
|
||||
完整参数列表见[参数说明](./params.md)。
|
||||
**5. 启动服务并调用**
|
||||
|
||||
```bash
|
||||
# 终端 1:启动服务
|
||||
python scripts/tools/server.py --param_path ./params --device cuda
|
||||
|
||||
# 终端 2:发起请求
|
||||
curl http://localhost:8000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"messages":[{"role":"user","content":"你好"}],"max_tokens":512}'
|
||||
```
|
||||
|
||||
### 演示
|
||||
|
||||
查看 `scripts/demo/` 文件夹中的演示:
|
||||
|
||||
```bash
|
||||
# 下载模型权重(运行演示前必需)
|
||||
python scripts/demo/download.py # model → params/
|
||||
|
||||
# 交互式流式聊天(多轮对话,保持历史记录)
|
||||
python scripts/demo/stream_chat.py
|
||||
# 在 >> 后输入消息,输入 !exit 退出
|
||||
|
||||
# 批量生成(5 条硬编码提示词,非流式)
|
||||
python scripts/demo/generate_batch.py
|
||||
|
||||
# 单条提示词自回归流式生成
|
||||
python scripts/demo/generate_ar.py
|
||||
```
|
||||
|
||||
所有生成演示默认使用 `temperature=0.8`、`top_p=0.95`、`top_k=50`、`max_tokens=2048`,需要 `params/` 目录包含模型权重(请先运行 `download.py`)。
|
||||
|
||||
观看 [bilibili](https://www.bilibili.com/video/BV1fuLB6yEj6) 上的视频演示。
|
||||
|
||||
---
|
||||
|
||||
更多选项请参考[文档](#文档)。
|
||||
|
||||
#### 文本生成
|
||||
|
||||
从 JSONL 文件批量生成:
|
||||
|
||||
```bash
|
||||
python scripts/tools/generate.py \
|
||||
--param_path /path/to/model \
|
||||
--input_json_file /path/to/input.json \
|
||||
--output_json_file /path/to/output.json
|
||||
--param_path ./params \
|
||||
--input_json_file input.jsonl \
|
||||
--output_json_file output.jsonl
|
||||
```
|
||||
|
||||
#### Docker
|
||||
@@ -129,9 +179,6 @@ docker build -t astrai:latest .
|
||||
# 启用 GPU 运行
|
||||
docker run --gpus all -it astrai:latest
|
||||
|
||||
# 指定特定 GPU
|
||||
docker run --gpus '"device=0,1"' -it astrai:latest
|
||||
|
||||
# 运行推理服务
|
||||
docker run --gpus all -p 8000:8000 astrai:latest \
|
||||
python -m scripts.tools.server --port 8000 --device cuda
|
||||
@@ -148,88 +195,42 @@ docker compose --profile cpu up -d
|
||||
|
||||
> **注意**: 必须使用 `--gpus all` 才能启用 CUDA 支持,否则 `torch.cuda.is_available()` 将返回 `False`。
|
||||
|
||||
#### 启动 HTTP 服务
|
||||
#### HTTP API 示例
|
||||
|
||||
启动推理服务器,支持 OpenAI 和 Anthropic 兼容的 HTTP API:
|
||||
除[快速上手](#快速上手)流程外,更多请求示例:
|
||||
|
||||
```bash
|
||||
python -m scripts.tools.server --port 8000 --device cuda
|
||||
```
|
||||
|
||||
发起请求:
|
||||
|
||||
```bash
|
||||
# OpenAI 兼容
|
||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"messages": [{"role": "user", "content": "你好"}],
|
||||
"max_tokens": 512
|
||||
}'
|
||||
|
||||
# OpenAI 兼容流式
|
||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"messages": [{"role": "user", "content": "讲个故事"}],
|
||||
"stream": true,
|
||||
"max_tokens": 500
|
||||
}'
|
||||
-d '{"messages":[{"role":"user","content":"讲个故事"}],"stream":true,"max_tokens":500}'
|
||||
|
||||
# Anthropic 兼容
|
||||
curl -X POST http://localhost:8000/v1/messages \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "astrai",
|
||||
"system": "你是一个乐于助人的助手。",
|
||||
"messages": [{"role": "user", "content": "你好"}],
|
||||
"max_tokens": 512
|
||||
}'
|
||||
-d '{"model":"astrai","system":"你是一个乐于助人的助手。","messages":[{"role":"user","content":"你好"}],"max_tokens":512}'
|
||||
|
||||
# Anthropic 兼容流式并设置停止序列
|
||||
curl -X POST http://localhost:8000/v1/messages \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "astrai",
|
||||
"messages": [{"role": "user", "content": "写个故事"}],
|
||||
"max_tokens": 500,
|
||||
"stream": true,
|
||||
"stop_sequences": ["结束"]
|
||||
}'
|
||||
-d '{"model":"astrai","messages":[{"role":"user","content":"写个故事"}],"max_tokens":500,"stream":true,"stop_sequences":["结束"]}'
|
||||
|
||||
# 健康检查
|
||||
curl http://localhost:8000/health
|
||||
```
|
||||
|
||||
#### 演示
|
||||
|
||||
查看 `scripts/demo/` 文件夹中的演示:
|
||||
|
||||
```bash
|
||||
# 下载预处理数据(运行演示前必需)
|
||||
python scripts/demo/download.py
|
||||
|
||||
# 交互式流式聊天
|
||||
python scripts/demo/stream_chat.py
|
||||
|
||||
# 批量生成
|
||||
python scripts/demo/generate_batch.py
|
||||
|
||||
# 自回归生成
|
||||
python scripts/demo/generate_ar.py
|
||||
```
|
||||
|
||||
观看 [bilibili](https://www.bilibili.com/video/BV1fuLB6yEj6) 上的视频演示。
|
||||
SSE 流式格式、错误码和统计端点详见[推理文档](./inference.md)。
|
||||
|
||||
### 文档
|
||||
|
||||
| 文档 | 说明 |
|
||||
|------|------|
|
||||
| [参数说明](./params.md) | 训练与推理参数配置 |
|
||||
| [CLI 参考](./params.md) | 所有 CLI 工具参数(训练、服务、生成、预处理) |
|
||||
| [架构文档](./architecture.md) | 系统架构、类图与设计模式 |
|
||||
| [训练文档](./training.md) | 训练循环、策略与公式 |
|
||||
| [推理文档](./inference.md) | KVCache、连续批处理、采样与 HTTP API |
|
||||
| [数据流程](./dataflow.md) | 数据管道、存储后端与数据集架构 |
|
||||
| [数据预处理](./preprocessing.md) | 声明式 JSON 驱动数据预处理 |
|
||||
|
||||
### 贡献
|
||||
|
||||
@@ -246,7 +247,7 @@ python scripts/demo/generate_ar.py
|
||||
|
||||
- **GitHub Issues**: [问题追踪](https://github.com/ViperEkura/AstrAI/issues)
|
||||
- **Discussions**: [GitHub 讨论区](https://github.com/ViperEkura/AstrAI/discussions)
|
||||
- **HuggingFace**: [模型中心](https://huggingface.co/ViperEk)
|
||||
- **HuggingFace**: [模型中心](https://huggingface.co/ViperEkura)
|
||||
|
||||
### 许可证
|
||||
|
||||
|
||||
+360
-153
@@ -1,5 +1,12 @@
|
||||
# AstrAI Architecture
|
||||
|
||||
## Contents
|
||||
|
||||
- [Class Diagram](#class-diagram) — Full Mermaid class diagram across 10+ namespaces
|
||||
- [Module Overview](#module-overview) — Component inventory per module
|
||||
- [Design Patterns](#design-patterns) — 13 documented patterns with classes
|
||||
- [Core Relationships](#core-relationships) — 11 key inter-component relationships
|
||||
|
||||
## Class Diagram
|
||||
|
||||
```mermaid
|
||||
@@ -8,62 +15,99 @@ classDiagram
|
||||
class BaseConfig {
|
||||
+to_dict() Dict
|
||||
+from_dict(d) Self
|
||||
+from_file(path) Self
|
||||
+to_file(path)
|
||||
}
|
||||
|
||||
class BaseModelConfig {
|
||||
+Optional[str] model_type
|
||||
+float neftune_alpha
|
||||
+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
|
||||
+Optional[int] vocab_size
|
||||
+Optional[int] dim
|
||||
+Optional[int] n_layers
|
||||
+Optional[float] norm_eps
|
||||
+Optional[int] dim_ffn
|
||||
+Optional[bool] tie_weight
|
||||
+Optional[dict] rope_scaling
|
||||
+int max_len
|
||||
+float rope_theta
|
||||
+Optional[int] max_len
|
||||
+Optional[float] rope_theta
|
||||
+str attn_type
|
||||
+int n_heads
|
||||
+int n_kv_heads
|
||||
+bool use_qk_norm
|
||||
+bool use_gated_attention
|
||||
+Optional[int] n_heads
|
||||
+Optional[int] n_kv_heads
|
||||
+Optional[bool] use_qk_norm
|
||||
+Optional[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[int] n_routed_experts
|
||||
+Optional[int] n_shared_experts
|
||||
+Optional[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[int] vocab_size
|
||||
+Optional[int] dim
|
||||
+Optional[int] n_layers
|
||||
+Optional[float] norm_eps
|
||||
+Optional[int] dim_ffn
|
||||
+Optional[int] max_len
|
||||
+Optional[float] rope_theta
|
||||
+str attn_type
|
||||
+Optional[int] n_heads
|
||||
+Optional[int] n_kv_heads
|
||||
+Optional[bool] use_qk_norm
|
||||
+str ffn_type
|
||||
+Optional[dict] rope_scaling
|
||||
+Optional[str] pooling_type
|
||||
+Optional[bool] normalize_embeddings
|
||||
}
|
||||
|
||||
class ConfigFactory {
|
||||
+Registry _registry
|
||||
+Dict _entries
|
||||
+register(name) decorator
|
||||
+load(raw) BaseConfig
|
||||
}
|
||||
|
||||
class InputConfig {
|
||||
+Optional[List[Dict]] sections
|
||||
+Optional[Dict[str, Dict]] sources
|
||||
}
|
||||
|
||||
class ProcessingConfig {
|
||||
+int max_seq_len
|
||||
+int min_chars
|
||||
+int max_chars
|
||||
+Optional[int] max_items
|
||||
+str packing_strategy
|
||||
+int max_packed_len
|
||||
+str truncation_mode
|
||||
}
|
||||
|
||||
class OutputConfig {
|
||||
+Optional[str] domain_key
|
||||
+str storage_format
|
||||
+int max_tokens_per_shard
|
||||
+Dict[str, str] dtype
|
||||
+str position_ids_mode
|
||||
}
|
||||
|
||||
class PipelineConfig {
|
||||
+int version
|
||||
+InputConfig input
|
||||
+dict mask
|
||||
+str mask_default
|
||||
+ProcessingConfig preprocessing
|
||||
+OutputConfig output
|
||||
+from_dict(d) Self
|
||||
}
|
||||
|
||||
class TrainConfig {
|
||||
+Callable[[], nn.Module] model_fn
|
||||
+str strategy
|
||||
@@ -73,14 +117,13 @@ classDiagram
|
||||
+int n_epoch
|
||||
+int batch_per_device
|
||||
+int grad_accum_steps
|
||||
+float max_grad_norm
|
||||
+Optional[float] max_grad_norm
|
||||
+list gradient_checkpointing_modules
|
||||
+int start_epoch
|
||||
+int start_batch
|
||||
+int start_samples
|
||||
+str ckpt_dir
|
||||
+int ckpt_interval
|
||||
+str log_dir
|
||||
+int log_interval
|
||||
+List[str] metrics
|
||||
+Optional[LoRAConfig] lora
|
||||
+int random_seed
|
||||
@@ -94,7 +137,9 @@ classDiagram
|
||||
+str start_method
|
||||
+str device_type
|
||||
+Optional[Dataset] val_dataset
|
||||
+Optional[float] val_split
|
||||
+int val_step
|
||||
+float neftune_alpha
|
||||
+str parallel_mode
|
||||
+dict executor_kwargs
|
||||
+dict extra_kwargs
|
||||
@@ -121,6 +166,13 @@ classDiagram
|
||||
+__getitem__(index) Dict
|
||||
}
|
||||
|
||||
class RecordDataset {
|
||||
+Optional[Callable] processor
|
||||
+load(load_path, storage_type)
|
||||
+__getitem__(index)
|
||||
+__len__()
|
||||
}
|
||||
|
||||
class DPODataset {
|
||||
+__getitem__(index) Dict
|
||||
}
|
||||
@@ -132,13 +184,26 @@ classDiagram
|
||||
class Store {
|
||||
+Dict[str, List[Tensor]] _data
|
||||
+Dict[str, List[int]] _cum
|
||||
+Dict[str, List[int]] _offsets
|
||||
+int _length
|
||||
+int _num_records
|
||||
+keys (property)
|
||||
+load(path)
|
||||
+fetch(begin, end, keys)
|
||||
+__len__()
|
||||
-_fetch_key(key, begin, end) Tensor
|
||||
-_normalize(raw)
|
||||
-_normalize(raw, offsets)
|
||||
}
|
||||
|
||||
class Streamable {
|
||||
<<mixin>>
|
||||
+fetch(begin, end, keys)
|
||||
-_fetch_stream_key(key, begin, end) Tensor
|
||||
}
|
||||
|
||||
class Recordable {
|
||||
<<mixin>>
|
||||
+num_records (property)
|
||||
+fetch_record(index, keys)
|
||||
-_fetch_record_key(key, index) Tensor
|
||||
}
|
||||
|
||||
class H5Store {
|
||||
@@ -150,22 +215,29 @@ classDiagram
|
||||
+load(path)
|
||||
}
|
||||
|
||||
class JsonlStore {
|
||||
+JsonlSource _source
|
||||
+Callable _processor
|
||||
+load(path, transform, processor)
|
||||
+fetch_record(index, keys)
|
||||
}
|
||||
|
||||
class ResumableDistributedSampler {
|
||||
+int epoch
|
||||
+int iter
|
||||
}
|
||||
|
||||
class StoreFactory {
|
||||
+Registry _registry
|
||||
+Dict _entries
|
||||
+register(name) decorator
|
||||
+create(storage_type) Store
|
||||
}
|
||||
|
||||
class DatasetFactory {
|
||||
+Registry _registry
|
||||
+Dict _entries
|
||||
+register(name) decorator
|
||||
+create(train_type, window_size, stride) BaseDataset
|
||||
+load(train_type, load_path, window_size, stride, storage_type) BaseDataset
|
||||
+load(train_type, load_path, window_size, stride, storage_type, tokenizer_path, max_len, store) BaseDataset
|
||||
}
|
||||
}
|
||||
|
||||
@@ -173,19 +245,20 @@ classDiagram
|
||||
class Checkpoint {
|
||||
+dict state_dict
|
||||
+int epoch
|
||||
+int iteration
|
||||
+int consumed_samples
|
||||
+dict extra
|
||||
+dict meta
|
||||
+dict config
|
||||
+save(save_dir)
|
||||
+load(save_dir, broadcast) Checkpoint
|
||||
+load_any(save_dir, broadcast) Optional[Checkpoint]
|
||||
}
|
||||
}
|
||||
|
||||
namespace model {
|
||||
class AutoModel {
|
||||
+BaseModelConfig config
|
||||
+Registry _registry
|
||||
+Dict _entries
|
||||
+register(name) decorator
|
||||
+get_component_class(name) Type
|
||||
+from_pretrained(path, disable_random_init, strict) nn.Module
|
||||
@@ -308,14 +381,57 @@ classDiagram
|
||||
|
||||
class Embedding {
|
||||
+Parameter weight
|
||||
+float neftune_noise_alpha
|
||||
+forward(x) Tensor
|
||||
+set_neftune_alpha(alpha)
|
||||
}
|
||||
}
|
||||
|
||||
namespace preprocessing {
|
||||
class BaseMaskBuilder {
|
||||
<<abstract>>
|
||||
+build(item, config, tokenizer) Optional[dict]
|
||||
}
|
||||
|
||||
class SectionedMaskBuilder {
|
||||
+SectionRenderer renderer
|
||||
+build(item, config, tokenizer) Optional[dict]
|
||||
+_build_single(item, config, tokenizer) Optional[dict]
|
||||
+_build_multi(item, sources_spec, config, tokenizer) Optional[dict]
|
||||
}
|
||||
|
||||
class Pipeline {
|
||||
+PipelineConfig config
|
||||
+List[str] paths
|
||||
+str output_dir
|
||||
+str tokenizer_path
|
||||
+AutoTokenizer tokenizer
|
||||
+BaseMaskBuilder mask_builder
|
||||
+PackingStrategy _packer
|
||||
+PositionIdStrategy _position_id
|
||||
+StoreWriter _writer
|
||||
+transform(item) Optional[dict]
|
||||
+run()
|
||||
+_flush(domains, shard_idx)
|
||||
+_inject_doc_reset_position_ids(keys, mode, seqs) Dict
|
||||
+_inject_continuous_position_ids(tensors, mode, seqs) Dict
|
||||
+_to_tensors(keys) Dict
|
||||
}
|
||||
|
||||
class TokenizeTransform {
|
||||
+PipelineConfig config
|
||||
+AutoTokenizer tokenizer
|
||||
+BaseMaskBuilder mask_builder
|
||||
+PositionIdStrategy position_strategy
|
||||
+from_config_file(path) TokenizeTransform
|
||||
+apply(records) Dict[str, list]
|
||||
}
|
||||
}
|
||||
|
||||
namespace tokenize {
|
||||
class AutoTokenizer {
|
||||
+vocab_size int
|
||||
+encode(tokens, out_ids, is_pretokenized, add_special_tokens) List[int]
|
||||
+encode(tokens, out_ids, is_pretokenized, add_special_tokens) List
|
||||
+decode(tokens, skip_special_tokens) str
|
||||
+__getattr__(name) Any (bos_id, eos_id, pad_id, stop_ids)
|
||||
+apply_chat_template(messages, system_prompt, tokenize, add_generation_prompt) Union[str, List[int]]
|
||||
@@ -333,18 +449,19 @@ classDiagram
|
||||
}
|
||||
|
||||
namespace factory {
|
||||
class Registry {
|
||||
class BaseFactory {
|
||||
+Dict _entries
|
||||
+register(name, component_cls, category, priority)
|
||||
+get(name) Type
|
||||
+list_names() List[str]
|
||||
+register(name) decorator
|
||||
+create(name, *args, **kwargs) T
|
||||
+get_component_class(name) Type
|
||||
+list_registered() list
|
||||
+is_registered(name) bool
|
||||
}
|
||||
|
||||
class BaseFactory {
|
||||
+Registry _registry
|
||||
+register(name, category, priority) decorator
|
||||
+create(name, *args, **kwargs) T
|
||||
+list_registered() list
|
||||
class MaskBuilderFactory {
|
||||
+Dict _entries
|
||||
+register(name) decorator
|
||||
+create(name, *args, **kwargs) BaseMaskBuilder
|
||||
}
|
||||
}
|
||||
|
||||
@@ -352,8 +469,8 @@ classDiagram
|
||||
class Trainer {
|
||||
+TrainConfig train_config
|
||||
+List[TrainCallback] callbacks
|
||||
+train(checkpoint)
|
||||
+_get_default_callbacks() List[TrainCallback]
|
||||
+train(resume_dir)
|
||||
-_get_default_callbacks() List[TrainCallback]
|
||||
}
|
||||
|
||||
class TrainContext {
|
||||
@@ -367,13 +484,15 @@ classDiagram
|
||||
+dict model_config
|
||||
+BaseExecutor executor
|
||||
+int epoch
|
||||
+int iteration
|
||||
+int consumed_samples
|
||||
+float loss
|
||||
+float grad_norm
|
||||
+DataLoader val_dataloader
|
||||
+float val_loss
|
||||
+int world_size
|
||||
+int rank
|
||||
+dict kwargs
|
||||
+optimizer_step() int
|
||||
}
|
||||
|
||||
class TrainContextBuilder {
|
||||
@@ -383,13 +502,17 @@ classDiagram
|
||||
}
|
||||
|
||||
class BaseStrategy {
|
||||
+Union[Callable, nn.Module] model
|
||||
+Callable model
|
||||
+Optional[BaseExecutor] executor
|
||||
+Optional[Callable] model_fn
|
||||
+dict extra_kwargs
|
||||
+str device
|
||||
+__call__(batch) Tensor
|
||||
+compute_loss(batch) Tensor
|
||||
}
|
||||
|
||||
class StrategyFactory {
|
||||
+Registry _registry
|
||||
+Dict _entries
|
||||
+register(name) decorator
|
||||
+create(train_type, model, device, **kwargs) BaseStrategy
|
||||
}
|
||||
@@ -412,30 +535,32 @@ classDiagram
|
||||
}
|
||||
|
||||
class GRPOStrategy {
|
||||
+nn.Module old_model
|
||||
+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()
|
||||
+sync_old_model()
|
||||
}
|
||||
|
||||
class BaseScheduler {
|
||||
+get_lr() List[float]
|
||||
+step()
|
||||
+state_dict() dict
|
||||
+load_state_dict(d)
|
||||
}
|
||||
|
||||
class SchedulerFactory {
|
||||
+Registry _registry
|
||||
+Dict _entries
|
||||
+register(name) decorator
|
||||
+create(optimizer, schedule_type, **kwargs) BaseScheduler
|
||||
+create(name, *args, **kwargs) BaseScheduler
|
||||
}
|
||||
|
||||
class CosineScheduler {
|
||||
+int warmup_steps
|
||||
+int lr_decay_steps
|
||||
+int total_steps
|
||||
+float min_rate
|
||||
}
|
||||
|
||||
@@ -446,6 +571,13 @@ classDiagram
|
||||
+int t_mult
|
||||
}
|
||||
|
||||
class WSDScheduler {
|
||||
+int warmup_steps
|
||||
+int stable_steps
|
||||
+int decay_steps
|
||||
+float min_rate
|
||||
}
|
||||
|
||||
class TrainCallback {
|
||||
<<protocol>>
|
||||
+on_train_begin(context)
|
||||
@@ -459,12 +591,12 @@ classDiagram
|
||||
}
|
||||
|
||||
class GradientClippingCallback {
|
||||
+float max_grad_norm
|
||||
+Optional[float] max_grad_norm
|
||||
+on_optimizer_step(context)
|
||||
}
|
||||
|
||||
class GradientCheckpointingCallback {
|
||||
+tuple modules
|
||||
+Optional[List[type]] modules
|
||||
+on_train_begin(context)
|
||||
+on_train_end(context)
|
||||
}
|
||||
@@ -474,55 +606,41 @@ classDiagram
|
||||
+int interval
|
||||
+bool weight_only
|
||||
+Callable save_extra_fn
|
||||
+_save_checkpoint(context)
|
||||
-_save_checkpoint(context)
|
||||
+on_batch_end(context)
|
||||
+on_train_end(context)
|
||||
+on_error(context)
|
||||
+save_extra(context)$
|
||||
+save_extra(context) dict
|
||||
}
|
||||
|
||||
class ProgressBarCallback {
|
||||
+int num_epoch
|
||||
+int log_interval
|
||||
+IO file
|
||||
+tqdm progress_bar
|
||||
+on_epoch_begin(context)
|
||||
+on_batch_end(context)
|
||||
+on_optimizer_step(context)
|
||||
+on_epoch_end(context)
|
||||
}
|
||||
|
||||
class MetricLoggerCallback {
|
||||
+str log_dir
|
||||
class MetricCallback {
|
||||
+Path log_dir
|
||||
+int save_interval
|
||||
+int log_interval
|
||||
+List[str] metrics
|
||||
+on_batch_end(context)
|
||||
+int val_step
|
||||
+on_optimizer_step(context)
|
||||
+on_epoch_end(context)
|
||||
+on_train_end(context)
|
||||
+on_error(context)
|
||||
}
|
||||
|
||||
class ValidationCallback {
|
||||
+_run_validation(context)
|
||||
+on_optimizer_step(context)
|
||||
-_run_validation(context)
|
||||
}
|
||||
|
||||
class CallbackFactory {
|
||||
+Registry _registry
|
||||
+Dict _entries
|
||||
+register(name) decorator
|
||||
+create(name, **kwargs) TrainCallback
|
||||
}
|
||||
|
||||
class Muon {
|
||||
+float lr
|
||||
+float momentum
|
||||
+float weight_decay
|
||||
+bool nesterov
|
||||
+int ns_steps
|
||||
+float adamw_lr
|
||||
+tuple adamw_betas
|
||||
+float adamw_eps
|
||||
+float adamw_wd
|
||||
+step(closure) Optional[float]
|
||||
}
|
||||
}
|
||||
|
||||
namespace inference {
|
||||
@@ -601,20 +719,44 @@ classDiagram
|
||||
}
|
||||
|
||||
class KVCache {
|
||||
-PagePool _pool
|
||||
-Storage _storage
|
||||
-TaskTable _table
|
||||
+int page_size
|
||||
<<abstract>>
|
||||
+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
|
||||
+bind_tasks(task_ids, total_len, device) CacheView
|
||||
}
|
||||
|
||||
class KvcacheView {
|
||||
class PageCache {
|
||||
+int page_size
|
||||
-PagePool _pool
|
||||
-Storage _storage
|
||||
-TaskTable _table
|
||||
+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)
|
||||
+bind_tasks(task_ids, total_len, device) PageCacheView
|
||||
}
|
||||
|
||||
class ContiguousCache {
|
||||
+int max_seq_len
|
||||
+Tensor k, v
|
||||
+task_alloc(task_id, prompt_ids) bool
|
||||
+task_free(task_id)
|
||||
+task_extend(task_id, pos) bool
|
||||
+bind_tasks(task_ids, total_len, device) ContiguousCacheView
|
||||
}
|
||||
|
||||
class CacheView {
|
||||
<<abstract>>
|
||||
+write(layer_id, k, v)
|
||||
+gather(layer_id) Tuple[Tensor, Tensor]
|
||||
}
|
||||
|
||||
class PageCacheView {
|
||||
-Storage _storage
|
||||
+Tensor _page_table
|
||||
+int _total_len
|
||||
@@ -622,6 +764,14 @@ classDiagram
|
||||
+gather(layer_id) Tuple[Tensor, Tensor]
|
||||
}
|
||||
|
||||
class ContiguousCacheView {
|
||||
-ContiguousCache _cache
|
||||
+Tensor _batch_indices
|
||||
+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]
|
||||
@@ -631,23 +781,22 @@ classDiagram
|
||||
+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 Task {
|
||||
+str task_id
|
||||
+List prompt_ids
|
||||
+Optional[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
|
||||
+Optional[float] finish_time
|
||||
+int next_pos
|
||||
+is_finished(stop_ids) bool
|
||||
}
|
||||
|
||||
class TaskStatus {
|
||||
<<enumeration>>
|
||||
@@ -671,6 +820,11 @@ classDiagram
|
||||
+activate(task)
|
||||
+return_to_waiting(tasks)
|
||||
+get_active_tasks() List[Task]
|
||||
+has_work() bool
|
||||
+wait_for_tasks(timeout)
|
||||
+get_waiting_tasks() List[Task]
|
||||
+clear_queues()
|
||||
+wake()
|
||||
+get_stats() Dict
|
||||
}
|
||||
|
||||
@@ -722,7 +876,9 @@ classDiagram
|
||||
|
||||
class ChatMessage {
|
||||
+str role
|
||||
+str content
|
||||
+Optional[str] content
|
||||
+Optional[List[Dict]] tool_calls
|
||||
+Optional[str] tool_call_id
|
||||
}
|
||||
|
||||
class ChatCompletionRequest {
|
||||
@@ -739,6 +895,8 @@ classDiagram
|
||||
+Optional[float] frequency_penalty
|
||||
+Optional[Dict[int, float]] logit_bias
|
||||
+Optional[str] user
|
||||
+Optional[List[ToolDef]] tools
|
||||
+Optional[Union[str, Dict]] tool_choice
|
||||
}
|
||||
|
||||
class AnthropicMessage {
|
||||
@@ -762,7 +920,7 @@ classDiagram
|
||||
<<abstract>>
|
||||
+prepare(request, engine) Tuple[str, GenContext, List[str]]
|
||||
+format_stream_start(ctx) List[str]
|
||||
+format_chunk(token) str
|
||||
+format_chunk(token) List[str]
|
||||
+format_stream_end(ctx, stop) List[str]
|
||||
+format_response(ctx, content, stop) Dict
|
||||
}
|
||||
@@ -770,7 +928,7 @@ classDiagram
|
||||
class OpenAIResponseBuilder {
|
||||
+prepare(request, engine) Tuple
|
||||
+format_stream_start(ctx) List[str]
|
||||
+format_chunk(token) str
|
||||
+format_chunk(token) List[str]
|
||||
+format_stream_end(ctx, stop) List[str]
|
||||
+format_response(ctx, content, stop) Dict
|
||||
}
|
||||
@@ -778,7 +936,7 @@ classDiagram
|
||||
class AnthropicResponseBuilder {
|
||||
+prepare(request, engine) Tuple
|
||||
+format_stream_start(ctx) List[str]
|
||||
+format_chunk(token) str
|
||||
+format_chunk(token) List[str]
|
||||
+format_stream_end(ctx, stop) List[str]
|
||||
+format_response(ctx, content, stop) Dict
|
||||
}
|
||||
@@ -787,12 +945,13 @@ classDiagram
|
||||
+request
|
||||
+engine
|
||||
+builder: ResponseBuilder
|
||||
+handle() Union[StreamingResponse, Dict]
|
||||
-_handle_stream(agen, ctx, stops) StreamingResponse
|
||||
-_handle_non_stream(agen, ctx, stops) Dict
|
||||
+async handle() Union[StreamingResponse, Dict]
|
||||
-_handle_stream(agen, ctx, stop_sequences) StreamingResponse
|
||||
-async _handle_non_stream(agen, ctx, stop_sequences) Dict
|
||||
}
|
||||
|
||||
class StopChecker {
|
||||
+__init__(sequences)
|
||||
+check(text) Optional[str]
|
||||
}
|
||||
|
||||
@@ -804,9 +963,15 @@ classDiagram
|
||||
+int completion_tokens
|
||||
}
|
||||
|
||||
class app {
|
||||
<<singleton>>
|
||||
+FastAPI app
|
||||
class StopInfo {
|
||||
+Optional[str] matched
|
||||
+str body
|
||||
+str yielded
|
||||
}
|
||||
|
||||
class get_app {
|
||||
<<module>>
|
||||
+get_app() FastAPI
|
||||
}
|
||||
}
|
||||
|
||||
@@ -829,14 +994,14 @@ classDiagram
|
||||
}
|
||||
|
||||
namespace parallel {
|
||||
class Functions {
|
||||
class setup {
|
||||
<<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)
|
||||
+setup_parallel(rank, world_size, backend, master_addr, master_port, device_type) contextmanager
|
||||
+get_current_device() str
|
||||
+get_world_size() int
|
||||
+get_rank() int
|
||||
+only_on_rank(rank, sync) decorator
|
||||
+only_on_rank(rank, sync=False) decorator
|
||||
}
|
||||
|
||||
class GradientState {
|
||||
@@ -847,6 +1012,7 @@ classDiagram
|
||||
class AccumOptimizer {
|
||||
+Optimizer optimizer
|
||||
+GradientState gradient_state
|
||||
+param_groups (property)
|
||||
+step(closure)
|
||||
+zero_grad()
|
||||
+state_dict() dict
|
||||
@@ -867,7 +1033,7 @@ classDiagram
|
||||
+prepare(model, optimizer, dataloader, scheduler) tuple
|
||||
+accumulate(model) context manager
|
||||
+backward(loss)
|
||||
+unwrap_model(model) nn.Module
|
||||
+unwrap_model(model) dict
|
||||
+sync_gradients (property) bool
|
||||
+grad_accum_steps (property) int
|
||||
}
|
||||
@@ -876,18 +1042,18 @@ classDiagram
|
||||
}
|
||||
|
||||
class DDPExecutor {
|
||||
+_prepare_model(model) nn.Module
|
||||
+_no_sync(model) context manager
|
||||
+unwrap_model(model) nn.Module
|
||||
-_prepare_model(model) nn.Module
|
||||
-_no_sync(model) context manager
|
||||
+unwrap_model(model) dict
|
||||
}
|
||||
|
||||
class FSDPExecutor {
|
||||
+_prepare_model(model) nn.Module
|
||||
+unwrap_model(model) nn.Module
|
||||
-_prepare_model(model) nn.Module
|
||||
+unwrap_model(model) dict
|
||||
}
|
||||
|
||||
class ExecutorFactory {
|
||||
+Registry _registry
|
||||
+Dict _entries
|
||||
+register(name) decorator
|
||||
+create(parallel_mode, **kwargs) BaseExecutor
|
||||
}
|
||||
@@ -899,11 +1065,25 @@ classDiagram
|
||||
}
|
||||
|
||||
class ColumnParallelLinear {
|
||||
+int in_features
|
||||
+int out_features
|
||||
+int out_features_per_rank
|
||||
+bool gather_results
|
||||
+Parameter weight
|
||||
+Optional[Parameter] bias
|
||||
+forward(x) Tensor
|
||||
+load_state_dict(state_dict)
|
||||
}
|
||||
|
||||
class RowParallelLinear {
|
||||
+int in_features
|
||||
+int out_features
|
||||
+int in_features_per_rank
|
||||
+bool reduce_results
|
||||
+Parameter weight
|
||||
+Optional[Parameter] bias
|
||||
+forward(x) Tensor
|
||||
+load_state_dict(state_dict)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -916,28 +1096,39 @@ classDiagram
|
||||
BaseStrategy <|-- GRPOStrategy
|
||||
BaseScheduler <|-- CosineScheduler
|
||||
BaseScheduler <|-- SGDRScheduler
|
||||
BaseScheduler <|-- WSDScheduler
|
||||
TrainCallback <|-- GradientClippingCallback
|
||||
TrainCallback <|-- GradientCheckpointingCallback
|
||||
TrainCallback <|-- CheckpointCallback
|
||||
TrainCallback <|-- ProgressBarCallback
|
||||
TrainCallback <|-- MetricLoggerCallback
|
||||
TrainCallback <|-- ValidationCallback
|
||||
TrainCallback <|-- MetricCallback
|
||||
BaseDataset <|-- SEQDataset
|
||||
BaseDataset <|-- SFTDataset
|
||||
BaseDataset <|-- DPODataset
|
||||
BaseDataset <|-- GRPODataset
|
||||
BaseDataset <|-- RecordDataset
|
||||
RecordDataset <|-- DPODataset
|
||||
RecordDataset <|-- GRPODataset
|
||||
Store <|-- H5Store
|
||||
Store <|-- MmapStore
|
||||
Store <|-- JsonlStore
|
||||
H5Store --|> Streamable
|
||||
H5Store --|> Recordable
|
||||
MmapStore --|> Streamable
|
||||
MmapStore --|> Recordable
|
||||
JsonlStore --|> Streamable
|
||||
JsonlStore --|> Recordable
|
||||
BaseSamplingStrategy <|-- TemperatureStrategy
|
||||
BaseSamplingStrategy <|-- TopKStrategy
|
||||
BaseSamplingStrategy <|-- TopPStrategy
|
||||
BaseSamplingStrategy <|-- SamplingPipeline
|
||||
ParallelModel <|-- RowParallelLinear
|
||||
ParallelModel <|-- ColumnParallelLinear
|
||||
AutoModel <|-- AutoRegressiveLM
|
||||
AutoModel <|-- EmbeddingEncoder
|
||||
BaseConfig <|-- BaseModelConfig
|
||||
BaseConfig <|-- TrainConfig
|
||||
BaseConfig <|-- InputConfig
|
||||
BaseConfig <|-- ProcessingConfig
|
||||
BaseConfig <|-- OutputConfig
|
||||
BaseConfig <|-- PipelineConfig
|
||||
BaseModelConfig <|-- AutoRegressiveLMConfig
|
||||
BaseModelConfig <|-- EncoderConfig
|
||||
BaseFactory <|-- AutoModel
|
||||
@@ -950,16 +1141,22 @@ classDiagram
|
||||
BaseFactory <|-- StoreFactory
|
||||
BaseFactory <|-- ExecutorFactory
|
||||
BaseFactory <|-- ConfigFactory
|
||||
BaseFactory <|-- MaskBuilderFactory
|
||||
BaseExecutor <|-- NoneExecutor
|
||||
BaseExecutor <|-- DDPExecutor
|
||||
BaseExecutor <|-- FSDPExecutor
|
||||
ResponseBuilder <|-- OpenAIResponseBuilder
|
||||
ResponseBuilder <|-- AnthropicResponseBuilder
|
||||
BaseMaskBuilder <|-- SectionedMaskBuilder
|
||||
KVCache <|-- PageCache
|
||||
KVCache <|-- ContiguousCache
|
||||
CacheView <|-- PageCacheView
|
||||
CacheView <|-- ContiguousCacheView
|
||||
|
||||
%% --- Composition (strong ownership, part destroyed with whole) ---
|
||||
KVCache *-- PagePool
|
||||
KVCache *-- Storage
|
||||
KVCache *-- TaskTable
|
||||
PageCache *-- PagePool
|
||||
PageCache *-- Storage
|
||||
PageCache *-- TaskTable
|
||||
InferenceEngine *-- InferenceScheduler
|
||||
InferenceScheduler *-- KVCache
|
||||
InferenceScheduler *-- Executor
|
||||
@@ -973,7 +1170,6 @@ classDiagram
|
||||
DecoderBlock *-- RMSNorm
|
||||
ChatCompletionRequest *-- ChatMessage
|
||||
MessagesRequest *-- AnthropicMessage
|
||||
BaseFactory *-- Registry
|
||||
BaseExecutor *-- GradientState
|
||||
AccumOptimizer o-- GradientState
|
||||
AccumScheduler o-- GradientState
|
||||
@@ -988,12 +1184,20 @@ classDiagram
|
||||
TrainContext o-- BaseScheduler
|
||||
TrainContext o-- Checkpoint
|
||||
TrainContext o-- BaseExecutor
|
||||
KvcacheView o-- Storage
|
||||
PageCacheView o-- Storage
|
||||
ContiguousCacheView o-- ContiguousCache
|
||||
SamplingPipeline o-- BaseSamplingStrategy
|
||||
BaseDataset o-- Store
|
||||
Pipeline o-- PipelineConfig
|
||||
Pipeline o-- BaseMaskBuilder
|
||||
Pipeline o-- AutoTokenizer
|
||||
TokenizeTransform o-- AutoTokenizer
|
||||
TokenizeTransform o-- BaseMaskBuilder
|
||||
|
||||
%% --- Dependency (uses temporarily) ---
|
||||
TrainConfig ..> BaseStrategy : selects
|
||||
PipelineConfig ..> MaskBuilderFactory : selects
|
||||
MaskBuilderFactory ..> BaseMaskBuilder : creates
|
||||
StrategyFactory ..> BaseStrategy : creates
|
||||
SchedulerFactory ..> BaseScheduler : creates
|
||||
DatasetFactory ..> BaseDataset : creates
|
||||
@@ -1006,6 +1210,7 @@ classDiagram
|
||||
DecoderBlock ..> FFNFactory : uses
|
||||
StoreFactory ..> H5Store : creates
|
||||
StoreFactory ..> MmapStore : creates
|
||||
StoreFactory ..> JsonlStore : creates
|
||||
ConfigFactory ..> AutoRegressiveLMConfig : creates
|
||||
ConfigFactory ..> EncoderConfig : creates
|
||||
ExecutorFactory ..> NoneExecutor : creates
|
||||
@@ -1019,7 +1224,8 @@ classDiagram
|
||||
TrainContextBuilder ..> ResumableDistributedSampler : creates
|
||||
Checkpoint ..> Checkpoint : serializes
|
||||
CheckpointCallback ..> Checkpoint : creates
|
||||
KVCache ..> KvcacheView : binds
|
||||
PageCache ..> PageCacheView : binds
|
||||
ContiguousCache ..> ContiguousCacheView : binds
|
||||
InferenceEngine ..> GenerationRequest : uses
|
||||
InferenceEngine ..> GenerateResult : creates
|
||||
OpenAIResponseBuilder ..> ChatCompletionRequest : receives
|
||||
@@ -1030,7 +1236,7 @@ classDiagram
|
||||
%% --- Association (general usage) ---
|
||||
Trainer --> TrainConfig
|
||||
DPOStrategy --> AutoModel
|
||||
GRPOStrategy --> AutoModel
|
||||
GRPOStrategy --> AutoModel : policy/old/ref
|
||||
InferenceScheduler --> Task
|
||||
InferenceScheduler --> TaskStatus
|
||||
Task --> TaskStatus
|
||||
@@ -1046,23 +1252,24 @@ classDiagram
|
||||
|
||||
| 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, Store–MmapStore, StoreFactory, ResumableDistributedSampler, DatasetFactory | Dataset loading and management |
|
||||
| **astrai.config** | BaseConfig, BaseModelConfig, AutoRegressiveLMConfig, EncoderConfig, ConfigFactory, TrainConfig, PipelineConfig, InputConfig, ProcessingConfig, OutputConfig | Configuration management (to_dict/from_dict, to_file/from_file) |
|
||||
| **astrai.preprocessing** | BaseMaskBuilder, MaskBuilderFactory, SectionedMaskBuilder, SingleOutputMaskBuilder, MultiOutputMaskBuilder, Pipeline, TokenizeTransform, filter_by_length, PackingStrategy, PackingStrategyFactory, plan_bfd, PositionIdStrategy, PositionIdStrategyFactory, StoreWriter, StoreWriterFactory, core (shared helpers) | Declarative JSON-driven data preprocessing |
|
||||
| **astrai.dataset** | BaseDataset–RecordDataset–DPO/GRPODataset, SEQDataset, SFTDataset, Store, Streamable, Recordable, H5Store, MmapStore, JsonlStore, StoreFactory, 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, GenerateResult, BaseSamplingStrategy–SamplingPipeline, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, ChatMessage–MessagesRequest, app | Inference service |
|
||||
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–WSDScheduler, SchedulerFactory, TrainCallback(Protocol)–MetricCallback, CallbackFactory | Training workflow |
|
||||
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, KVCache–ContiguousCache/PageCache, CacheView–ContiguousCacheView/PageCacheView, Allocator–Storage, Task, TaskManager, TaskStatus, GenerationRequest, GenerateResult, BaseSamplingStrategy–SamplingPipeline, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, ChatMessage–MessagesRequest, app | Inference service |
|
||||
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, GradientState, AccumOptimizer, AccumScheduler, ParallelModel, RowParallelLinear, ColumnParallelLinear | Distributed parallel & gradient accumulation |
|
||||
| **astrai.factory** | Registry, BaseFactory[T] | Component registration |
|
||||
| **astrai.factory** | BaseFactory | Component registration |
|
||||
| **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
|
||||
|
||||
## Design Patterns
|
||||
|
||||
| Pattern | Classes | Purpose |
|
||||
|---------|---------|---------|
|
||||
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory` | Decorator-based component creation |
|
||||
| **Registry** | `BaseFactory`, `Registry` | Component registration with category/priority |
|
||||
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory`, `MaskBuilderFactory`, `StoreWriterFactory`, `PackingStrategyFactory`, `PositionIdStrategyFactory` | Decorator-based component creation |
|
||||
| **Registry** | `BaseFactory` | Component registration |
|
||||
| **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching |
|
||||
| **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations |
|
||||
| **Strategy (API)** | `ResponseBuilder`, `OpenAIResponseBuilder`, `AnthropicResponseBuilder` | HTTP API handler with format hooks |
|
||||
@@ -1070,23 +1277,23 @@ classDiagram
|
||||
| **Observer** | `TrainCallback`, callback implementations | Training process monitoring |
|
||||
| **Context** | `TrainContext` | Unified training state bag |
|
||||
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction |
|
||||
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor` | Gradient accumulation & model distribution |
|
||||
| **Storage** | `Store`, `H5Store`, `MmapStore` | Format-agnostic data access with multi-segment support |
|
||||
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution |
|
||||
| **Storage** | `Store`, `H5Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support |
|
||||
| **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, `parallel_mode`, `executor_kwargs`
|
||||
1. **Config → Training**: `TrainConfig` holds `model_fn`, `dataset`, `optimizer_fn`, `scheduler_fn`, `parallel_mode`, `executor_kwargs`
|
||||
2. **Training Flow**: `Trainer` → `TrainContextBuilder` → `TrainContext`, uses `BaseStrategy` for loss, `BaseExecutor` for gradient accumulation + model distribution
|
||||
3. **Strategy Selection**: `StrategyFactory` creates strategy by `train_type`
|
||||
4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)` → `NoneExecutor` / `DDPExecutor` / `FSDPExecutor`
|
||||
5. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `AutoRegressiveLM`, backed by `KVCache` + `SamplingPipeline`
|
||||
6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
|
||||
7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (H5Store/MmapStore) loads data with explicit `_length` and multi-segment `_data`
|
||||
7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (H5Store/MmapStore/JsonlStore) loads data with explicit `_length` and multi-segment `_data`
|
||||
8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only), extra state saved as `{key}.pt`
|
||||
9. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`
|
||||
9. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`/`WSDScheduler`
|
||||
10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
|
||||
11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers
|
||||
|
||||
> Document Update Time: 2026-05-28
|
||||
> Document Update Time: 2026-07-19
|
||||
|
||||
+91
-18
@@ -1,46 +1,119 @@
|
||||
# Data Flow
|
||||
|
||||
This document describes the data pipeline: from raw text to model input tensors.
|
||||
This document describes the data pipeline: from raw text to model input tensors. For creating preprocessing configs, see [Preprocessing Guide](preprocessing.md).
|
||||
|
||||
## Contents
|
||||
|
||||
- [Overview](#overview)
|
||||
- [Data Preparation](#data-preparation) — tokenization, format detection, backends
|
||||
- [Data Keys by Training Type](#data-keys-by-training-type)
|
||||
- [Dataset Architecture](#dataset-architecture)
|
||||
- [Sampler](#sampler)
|
||||
- [DataLoader](#dataloader)
|
||||
|
||||
## Overview
|
||||
|
||||
```
|
||||
Raw Text → AutoTokenizer → Token IDs → .h5/.bin → Dataset → Sampler → DataLoader → Training/Inference
|
||||
JSONL Lines → Pipeline (mask builder) → Tokenized Tensors
|
||||
↓
|
||||
.h5 or .bin storage
|
||||
↓
|
||||
Store.load()
|
||||
↓
|
||||
Store.fetch(begin, end, keys)
|
||||
↓
|
||||
BaseDataset.__getitem__(idx)
|
||||
↓
|
||||
Sampler → DataLoader → Training / Inference
|
||||
```
|
||||
|
||||
## Data Preparation
|
||||
|
||||
Raw text is tokenized via `AutoTokenizer.encode()` and saved as HDF5 (`.h5`) or binary (`.bin` + `meta.json`) files with keyed tensor groups.
|
||||
|
||||
### Tokenization
|
||||
|
||||
The `Pipeline` reads JSONL lines, applies the mask builder (see [Preprocessing](preprocessing.md)), and produces flat token sequences:
|
||||
|
||||
```python
|
||||
# Per JSONL line: messages → chat template → token IDs + loss mask
|
||||
tokens = tokenizer.encode(rendered_text) # List[int]
|
||||
loss_mask = [0, 0, 0, 1, 1, 1, 1, 1, 1] # 0=masked, 1=train
|
||||
# Stored as flat tensors, packed with other lines by packing strategy
|
||||
```
|
||||
|
||||
The output `meta.json` records the storage format, key names, dtype, total token count, and tensor shapes for each shard.
|
||||
|
||||
### Format Detection
|
||||
|
||||
`detect_format(load_path)` inspects the path:
|
||||
|
||||
- If `load_path` is a file: checks suffix — `.h5`/`.hdf5` → `"h5"`, `.jsonl` → `"jsonl"`, unknown suffix raises `ValueError`
|
||||
- If `load_path` is a directory: recursively globs for `*.h5`/`*.hdf5` files → `"h5"`, `*.bin` + `**/meta.json` → `"bin"`, or `*.jsonl` + `dataset_config.json` → `"jsonl"`
|
||||
|
||||
### Store Backends
|
||||
|
||||
Storage format is auto-detected by `detect_format()`; backends are dispatched via registry:
|
||||
|
||||
```
|
||||
StoreFactory.create("h5") → H5Store
|
||||
StoreFactory.create("bin") → MmapStore
|
||||
StoreFactory.create("h5") → H5Store
|
||||
StoreFactory.create("bin") → MmapStore
|
||||
StoreFactory.create("jsonl") → JsonlStore
|
||||
```
|
||||
|
||||
H5 backend supports shared memory via `.share_memory_()`. Bin (mmap) uses OS page-cache sharing natively.
|
||||
All three inherit `Store` (base, owns `_data`/`_cum`/`_offsets`/`_normalize`) plus the `Streamable` and `Recordable` mixins, so every backend supports both `fetch(begin, end, keys)` (stream) and `fetch_record(index, keys)` (record) APIs.
|
||||
|
||||
**H5Store**: Reads HDF5 files. Tensors are loaded into host memory and normalized into segmented storage. `segments_are_records=True` — each `data_i` dataset is one record.
|
||||
|
||||
**MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`. `segments_are_records=False` — bin segments are contiguous streams; record access is driven by `_offsets` (written when `save_bin(..., record_keys=...)` was used at preprocessing time).
|
||||
|
||||
**JsonlStore**: On-the-fly tokenization of raw JSONL files at load time. Requires a `dataset_config.json` alongside the `.jsonl` files following the same `PipelineConfig` schema with an additional `tokenizer_path` field. Two modes: eager (default, applies `TokenizeTransform` to all records at load) and lazy (`processor=fn` given, defers tokenisation to `fetch_record` — used by DPO/GRPO).
|
||||
|
||||
All backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for bisect-based stream indexing) + `Store._offsets[Dict[str, List[int]]]` (per-record offsets for record-mode indexing). Nested keys (GRPO `responses`/`masks` as `List[List[Tensor]]`) are stored as-is and excluded from both bookkeepings — they are only accessed record-by-record.
|
||||
|
||||
## Data Keys by Training Type
|
||||
|
||||
| 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` |
|
||||
| Type | Storage Keys | Access Mode |
|
||||
|------|-------------|-------------|
|
||||
| `seq` | `sequence` (→ input_ids, target_ids via offset-by-1) | stream (`fetch`) |
|
||||
| `sft` | `sequence`, `loss_mask`, `position_ids` | stream (`fetch`) |
|
||||
| `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` | record (`fetch_record`) |
|
||||
| `grpo` | `prompts`, `responses`, `masks`, `rewards` | record (`fetch_record`) |
|
||||
|
||||
## Dataset Architecture
|
||||
|
||||
```
|
||||
DatasetFactory.load(train_type, load_path, window_size, stride, storage_type)
|
||||
→ StoreFactory.create(detect_format(path))
|
||||
→ Store._data[Dict[str, List[Tensor]]] + _cum[Dict[str, List[int]]]
|
||||
→ BaseDataset.__getitem__(idx)
|
||||
→ sliding window [begin, end) via get_index(idx)
|
||||
DatasetFactory.load(train_type, load_path, window_size, stride=None,
|
||||
storage_type=None, tokenizer_path=None,
|
||||
max_len=2048, store=None)
|
||||
→ BaseDataset.load(load_path, storage_type=None)
|
||||
→ detect_format(load_path)
|
||||
→ StoreFactory.create(storage_type)
|
||||
→ Store.load(load_path)
|
||||
→ _normalize(raw) # base Store, shared by both backends
|
||||
→ Store._data[Dict[str, List[Tensor]]]
|
||||
+ _cum[Dict[str, List[int]]] (stream mode)
|
||||
+ _offsets[Dict[str, List[int]]] (record mode)
|
||||
|
||||
Stream datasets (SEQ/SFT):
|
||||
BaseDataset.__getitem__(idx)
|
||||
→ get_index(idx) → [begin, end)
|
||||
→ Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor]
|
||||
|
||||
Record datasets (DPO/GRPO via RecordDataset):
|
||||
RecordDataset.__getitem__(idx)
|
||||
→ Store.fetch_record(idx, keys) → Tensor / Dict[str, Tensor]
|
||||
```
|
||||
|
||||
`window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`).
|
||||
Class hierarchy: `BaseDataset` ← `SEQDataset` / `SFTDataset` (stream); `BaseDataset` ← `RecordDataset` ← `DPODataset` / `GRPODataset` (record).
|
||||
|
||||
`window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`, optional). Only meaningful for stream datasets — record datasets ignore both. `storage_type` defaults to `None` (auto-detect via `detect_format`).
|
||||
|
||||
`tokenizer_path` triggers lazy on-the-fly tokenisation for record datasets on raw JSONL (DPO builds a `dpo_processor`; SEQ/SFT/pre-tokenised backends ignore it). `store` (pre-built `Store`) bypasses `load_path`/`storage_type`/`tokenizer_path` entirely — the caller controls Store construction.
|
||||
|
||||
`Store.fetch(begin, end, keys)` (stream mode, on `Streamable`): accepts a single key (`str`) returning a `Tensor`, or a list of keys returning `Dict[str, Tensor]`. Internally uses `bisect` across multi-segment tensors. Raises `RuntimeError("Store not loaded")` if called before `load()`.
|
||||
|
||||
`Store.fetch_record(index, keys)` (record mode, on `Recordable`): same key API. Uses `_offsets[key]` when present (bin layout with per-record offsets), otherwise indexes `_data[key]` directly (H5/JSONL where each segment is one record).
|
||||
|
||||
## Sampler
|
||||
|
||||
@@ -54,4 +127,4 @@ DatasetFactory.load(train_type, load_path, window_size, stride, storage_type)
|
||||
|
||||
Standard PyTorch `DataLoader` with configurable `batch_size`, `num_workers`, `pin_memory`, `prefetch_factor`. Sampler produces indices; dataloader fetches tensor batches via `__getitem__`.
|
||||
|
||||
> Document Update Time: 2026-05-28
|
||||
> Document Update Time: 2026-07-19
|
||||
|
||||
+122
-18
@@ -1,5 +1,16 @@
|
||||
# Inference
|
||||
|
||||
## Contents
|
||||
|
||||
- [KV Cache](#kv-cache)
|
||||
- [KVCache System](#kvcache-system)
|
||||
- [Continuous Batching](#continuous-batching)
|
||||
- [Sampling](#sampling-strategy-pattern)
|
||||
- [Protocol Handlers](#protocol-handlers-strategy-pattern)
|
||||
- [Engine & GenerateResult](#engine--generateresult)
|
||||
- [HTTP API](#http-api) — endpoints, SSE, errors, stats
|
||||
- [Engine API](#engine-api)
|
||||
|
||||
## KV Cache
|
||||
|
||||
At decode time, only the last query token matters. All previous K/V are cached to avoid recomputation:
|
||||
@@ -12,29 +23,40 @@ RoPE is applied **before** KV cache write, not after — otherwise position enco
|
||||
|
||||
## KVCache System
|
||||
|
||||
Six classes working together:
|
||||
Seven classes working together, with two concrete cache implementations:
|
||||
|
||||
### ContiguousCache (default)
|
||||
|
||||
```
|
||||
KVCache (facade)
|
||||
├── PagePool orchestrates page allocation + prefix matching
|
||||
│ ├── Allocator bitmask-based page allocator + ref-count + LRU eviction (inside PagePool)
|
||||
│ └── PrefixCache hash-based prefix matching (page_hash via polynomial hash) (inside PagePool)
|
||||
├── 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 (returned by bind())
|
||||
ContiguousCache (simple contiguous per-slot cache)
|
||||
├── ContiguousCacheView bundles k/v tensors + slot indices for attention layers
|
||||
```
|
||||
|
||||
`KVCache.bind(page_table, total_len)` returns a `KvcacheView` used by attention layers via `write()` / `gather()`.
|
||||
Created by default when no cache is passed to `InferenceScheduler`. Each task occupies a fixed slot of `[max_seq_len, n_kv_heads, head_dim]`. Simple and efficient for small-to-medium batch sizes.
|
||||
|
||||
### PageCache (paged with prefix sharing)
|
||||
|
||||
```
|
||||
PageCache (paged KV cache with prefix sharing, alternative)
|
||||
├── PagePool orchestrates page allocation + prefix matching
|
||||
│ ├── Allocator bitmask-based page allocator + ref-count + LRU
|
||||
│ └── PrefixCache hash-based prefix matching (page_hash via polynomial hash)
|
||||
├── 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)
|
||||
└── PageCacheView bundles Storage + page_table + total_len for attention layers
|
||||
```
|
||||
|
||||
`isinstance(cache, KVCache)` checks dispatch to the correct view. Both implement the abstract `KVCache` interface used by `Executor` and `InferenceScheduler`.
|
||||
|
||||
## 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
|
||||
1. Cleanup → Remove finished tasks, free KV cache slots/pages
|
||||
2. Refill → Pop from waiting_queue, task_alloc resources, activate
|
||||
3. Prefill → Group by (prompt_len, start_pos), run full forward
|
||||
4. Decode → Pick largest same-position group, single-token forward
|
||||
4. Decode → Run single-token forward for each same-position group
|
||||
```
|
||||
|
||||
## Sampling (Strategy Pattern)
|
||||
@@ -43,7 +65,8 @@ KVCache (facade)
|
||||
BaseSamplingStrategy (ABC)
|
||||
├── TemperatureStrategy
|
||||
├── TopKStrategy
|
||||
└── TopPStrategy
|
||||
├── TopPStrategy
|
||||
└── SamplingPipeline
|
||||
```
|
||||
|
||||
`SamplingPipeline` composes them: Temperature → Top-K → Top-P → softmax → multinomial.
|
||||
@@ -73,7 +96,9 @@ Adding a protocol = one builder file, no handler subclassing needed.
|
||||
InferenceEngine
|
||||
├── generate(prompt, stream, ...) → str | List[str] | Generator
|
||||
├── generate_with_request(req) → same
|
||||
└── generate_async(prompt, ...) → AsyncGenerator
|
||||
├── generate_async(prompt, ...) → AsyncGenerator
|
||||
├── get_stats() → Dict
|
||||
└── shutdown()
|
||||
```
|
||||
|
||||
`GenerateResult` uses `Condition` for non-streaming (`wait_completion()`) and `Event` for streaming (`wait()`). Stream callback is `cb(token)`.
|
||||
@@ -124,12 +149,90 @@ Supports `stop_sequences` and streaming via `event: content_block_delta`.
|
||||
| Param | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `messages` | List[dict] | required | Chat messages (role, content) |
|
||||
| `temperature` | float | 1.0 | Sampling temperature (>= 0.0) |
|
||||
| `top_p` | float | 1.0 | Nucleus threshold |
|
||||
| `top_k` | int | 50 | Top-k count |
|
||||
| `top_p` | float | 1.0 | Nucleus threshold |
|
||||
| `temperature` | float | 1.0 | Sampling temperature (> 0.0) |
|
||||
| `max_tokens` | Optional[int] | None | Max generation length |
|
||||
| `stream` | bool | False | Stream output |
|
||||
|
||||
### SSE Streaming Format
|
||||
|
||||
**OpenAI** (`/v1/chat/completions`, `stream=true`):
|
||||
|
||||
```
|
||||
data: {"id":"chatcmpl-...","object":"chat.completion.chunk","created":...,"model":"astrai",
|
||||
"choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}
|
||||
|
||||
data: {"id":"chatcmpl-...","object":"chat.completion.chunk","created":0,"model":"astrai",
|
||||
"choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]}
|
||||
|
||||
data: {"id":"chatcmpl-...","object":"chat.completion.chunk","created":...,"model":"astrai",
|
||||
"choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}
|
||||
|
||||
data: {"prompt_tokens":5,"completion_tokens":1,"total_tokens":6}
|
||||
|
||||
data: [DONE]
|
||||
```
|
||||
|
||||
**Anthropic** (`/v1/messages`, `stream=true`):
|
||||
|
||||
```
|
||||
event: message_start
|
||||
data: {"type":"message_start","message":{"id":"msg_...","model":"astrai","role":"assistant",
|
||||
"content":[],"usage":{"input_tokens":0}}}
|
||||
|
||||
event: content_block_start
|
||||
data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}
|
||||
|
||||
event: content_block_delta
|
||||
data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello"}}
|
||||
|
||||
event: content_block_stop
|
||||
data: {"type":"content_block_stop","index":0}
|
||||
|
||||
event: message_delta
|
||||
data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{...}}
|
||||
|
||||
event: message_stop
|
||||
data: {"type":"message_stop"}
|
||||
```
|
||||
|
||||
### Error Responses
|
||||
|
||||
The server returns standard HTTP status codes. Pydantic validation errors (e.g. missing required fields)
|
||||
are handled automatically by FastAPI with 422 status. The only application-level error is engine initialization:
|
||||
|
||||
| Status | Meaning |
|
||||
|--------|---------|
|
||||
| 200 | Success |
|
||||
| 422 | Unprocessable entity (Pydantic validation) |
|
||||
| 503 | Service unavailable (model not loaded, engine not ready) |
|
||||
|
||||
Error response body (503):
|
||||
|
||||
```json
|
||||
{
|
||||
"detail": "Engine not initialized"
|
||||
}
|
||||
```
|
||||
|
||||
### Stats Endpoint
|
||||
|
||||
```
|
||||
GET /stats
|
||||
```
|
||||
|
||||
Response:
|
||||
|
||||
```json
|
||||
{
|
||||
"total_tasks": 128,
|
||||
"total_tokens": 10240,
|
||||
"active_tasks": 3,
|
||||
"waiting_queue": 2
|
||||
}
|
||||
```
|
||||
|
||||
## Engine API
|
||||
|
||||
```python
|
||||
@@ -142,7 +245,8 @@ 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]
|
||||
async for token in engine.generate_async("Hello", ...): # -> AsyncGenerator[str]
|
||||
print(token)
|
||||
```
|
||||
|
||||
> Document Update Time: 2026-05-28
|
||||
> Document Update Time: 2026-07-09
|
||||
|
||||
+117
-12
@@ -1,4 +1,11 @@
|
||||
# Parameter Documentation
|
||||
# CLI Parameter Reference
|
||||
|
||||
## Contents
|
||||
|
||||
- [Training Parameters](#training-parameters)
|
||||
- [Inference Server](#inference-server-serverpy)
|
||||
- [Generate](#generate-generatepy)
|
||||
- [Preprocess](#preprocess-preprocesspy)
|
||||
|
||||
## Training Parameters
|
||||
|
||||
@@ -19,15 +26,19 @@
|
||||
|-----------|-------------|---------|
|
||||
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
|
||||
| `--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 (None disables) | None |
|
||||
|
||||
### Optimizer (AdamW)
|
||||
### Optimizer (MuonMix)
|
||||
|
||||
Combined optimizer: matrix parameters via **Muon**, non-matrix via **AdamW** (`fused=True`).
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--adamw_beta1` | AdamW beta1 | 0.9 |
|
||||
| `--adamw_beta2` | AdamW beta2 | 0.95 |
|
||||
| `--adamw_weight_decay` | AdamW weight decay | 0.01 |
|
||||
| `--weight_decay` | Weight decay (applied to Muon matrix params; non-matrix use 0) | 0.1 |
|
||||
| `--muon_momentum` | Muon momentum factor | 0.95 |
|
||||
| `--muon_nesterov` | Enable Nesterov momentum for Muon | True |
|
||||
| `--muon_ns_steps` | Newton-Schulz iteration steps for Muon | 5 |
|
||||
| `--muon_adjust_lr` | Muon LR adjustment strategy (`original`, `match_rms_adamw`) | `match_rms_adamw` |
|
||||
|
||||
### Data Loading
|
||||
|
||||
@@ -46,7 +57,27 @@
|
||||
| `--ckpt_interval` | Iterations between checkpoints | 5000 |
|
||||
| `--ckpt_dir` | Checkpoint save directory | checkpoint |
|
||||
| `--start_epoch` | Resume from epoch (0 = from scratch) | 0 |
|
||||
| `--start_batch` | Resume from batch iteration | 0 |
|
||||
| `--start_samples` | Resume from sample count per rank | 0 |
|
||||
|
||||
### Validation
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--val_split` | Ratio to split from training dataset for validation (e.g. 0.05) | None |
|
||||
| `--val_step` | Number of optimizer steps between validation runs | 1000 |
|
||||
|
||||
### Logging
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--log_dir` | Directory for metric logs | checkpoint/logs |
|
||||
| `--metrics` | Metrics to log (e.g. --metrics loss lr val_loss) | ["loss", "lr", "grad_norm"] |
|
||||
|
||||
### Gradient Checkpointing
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--gradient_checkpointing` | Enable activation checkpointing for DecoderBlock modules | False |
|
||||
|
||||
### Distributed Training
|
||||
|
||||
@@ -56,17 +87,32 @@
|
||||
| `--parallel_mode` | Parallel strategy (`none`, `ddp`, or `fsdp`) | none |
|
||||
| `--device_type` | Device type | cuda |
|
||||
| `--start_method` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) | spawn |
|
||||
| `--backend` | Distributed training backend | nccl |
|
||||
| `--master_addr` | Master node address | localhost |
|
||||
| `--master_port` | Master node port | 29500 |
|
||||
|
||||
### Strategy-specific
|
||||
|
||||
| Parameter | Description | Default | Used by |
|
||||
|-----------|-------------|---------|---------|
|
||||
| `--dpo_beta` | DPO beta value | 0.1 | `dpo` |
|
||||
| `--label_smoothing` | Label smoothing for cross-entropy loss | 0.05 | `seq`, `sft` |
|
||||
| `--label_smoothing` | Label smoothing for cross-entropy loss | 0.0 | `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` |
|
||||
| `--neftune_alpha` | NEFTune noise alpha (0=disabled, typical: 5.0) | 0.0 | `sft` |
|
||||
|
||||
### Scheduler
|
||||
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--schedule_type` | LR scheduler type (`cosine`, `sgdr`, `wsd`) | cosine |
|
||||
| `--min_rate` | Minimum LR as fraction of base LR | None (scheduler default: 0.01) |
|
||||
| `--cycle_length` | SGDR first cycle length in steps | None (total_steps - warmup_steps) |
|
||||
| `--t_mult` | SGDR cycle length multiplier per restart | 2 |
|
||||
| `--stable_steps` | WSD stable plateau steps | None (required for wsd) |
|
||||
| `--decay_steps` | WSD decay steps | None (total_steps - warmup_steps - stable_steps) |
|
||||
|
||||
### Usage Example
|
||||
|
||||
@@ -75,6 +121,7 @@ export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||
|
||||
nohup python scripts/tools/train.py \
|
||||
--nprocs=4 \
|
||||
--parallel_mode=ddp \
|
||||
--train_type=seq \
|
||||
--data_root_path=/path/to/dataset \
|
||||
--param_path=/path/to/model \
|
||||
@@ -83,9 +130,7 @@ nohup python scripts/tools/train.py \
|
||||
--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 \
|
||||
--weight_decay=0.1 \
|
||||
--window_size=2048 \
|
||||
--ckpt_interval=10000 \
|
||||
--ckpt_dir=./checkpoint \
|
||||
@@ -96,4 +141,64 @@ nohup python scripts/tools/train.py \
|
||||
|
||||
---
|
||||
|
||||
> Document Update Time: 2026-05-24
|
||||
## Inference Server (`server.py`)
|
||||
|
||||
| Parameter | Type | Default | Description |
|
||||
|-----------|------|---------|-------------|
|
||||
| `--host` | str | `0.0.0.0` | Host address |
|
||||
| `--port` | int | `8000` | Port number |
|
||||
| `--param_path` | path | `project_root/params` | Path to model parameters |
|
||||
| `--device` | str | `cuda` | Device to load model on |
|
||||
| `--dtype` | str | `bfloat16` | Model weights dtype (`bfloat16`, `float16`, `float32`) |
|
||||
| `--max_batch_size` | int | `16` | Maximum batch size for continuous batching |
|
||||
| `--reload` | flag | `False` | Enable auto-reload for development |
|
||||
|
||||
Usage:
|
||||
```bash
|
||||
python scripts/tools/server.py --param_path ./params --device cuda --dtype bfloat16
|
||||
```
|
||||
|
||||
See [Inference Guide](inference.md) for HTTP API documentation.
|
||||
|
||||
## Generate (`generate.py`)
|
||||
|
||||
| Parameter | Type | Default | Description |
|
||||
|-----------|------|---------|-------------|
|
||||
| `--param_path` | str | required | Path to the model directory |
|
||||
| `--input_json_file` | str | required | Path to the input JSONL file |
|
||||
| `--output_json_file` | str | required | Path to the output JSONL file |
|
||||
| `--question_key` | str | `question` | Key for the question in input JSON |
|
||||
| `--response_key` | str | `response` | Key for the response in output JSON |
|
||||
| `--temperature` | float | `0.60` | Sampling temperature |
|
||||
| `--top_k` | int | `30` | Top-k filtering |
|
||||
| `--top_p` | float | `0.95` | Nucleus sampling threshold |
|
||||
| `--batch_size` | int | `1` | Batch size for generation |
|
||||
| `--max_tokens` | int | model config `max_len` | Maximum tokens to generate |
|
||||
|
||||
Usage:
|
||||
```bash
|
||||
python scripts/tools/generate.py \
|
||||
--param_path ./params \
|
||||
--input_json_file input.jsonl \
|
||||
--output_json_file output.jsonl
|
||||
```
|
||||
|
||||
## Preprocess (`preprocess.py`)
|
||||
|
||||
| Parameter | Type | Default | Description |
|
||||
|-----------|------|---------|-------------|
|
||||
| `input_files` | path(s) | required | Input JSONL file(s), supports glob (`data/*.jsonl`) |
|
||||
| `--output_dir`, `-o` | path | required | Output directory for processed data |
|
||||
| `--config`, `-c` | path | required | Preprocessing pipeline config (JSON) |
|
||||
| `--tokenizer_path` | str | `params` | Path to tokenizer directory |
|
||||
|
||||
Usage:
|
||||
```bash
|
||||
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c sft.json
|
||||
```
|
||||
|
||||
See [Preprocessing Guide](preprocessing.md) for config file format and examples.
|
||||
|
||||
---
|
||||
|
||||
> Document Update Time: 2026-07-19
|
||||
@@ -0,0 +1,364 @@
|
||||
# Preprocessing Pipeline
|
||||
|
||||
Declarative JSON-driven data preprocessing. `MaskBuilderFactory` supports three registered builders: `"single"` (single-output via `input.sections`), `"multi"` (multi-output via `input.sources`), and `"sectioned"` (façade dispatching to `single` or `multi` based on config).
|
||||
|
||||
## Contents
|
||||
|
||||
- [Philosophy](#philosophy)
|
||||
- [Config Structure](#config-structure)
|
||||
- [Quick Start](#quick-start) — SFT Chat, SFT Instruction, Pretrain, DPO, GRPO examples
|
||||
- [Configuration Reference](#configuration-reference) — all fields
|
||||
- [Mask Algorithm](#mask-algorithm)
|
||||
- [Output Layout](#output-layout)
|
||||
- [CLI](#cli)
|
||||
- [Python API](#python-api)
|
||||
|
||||
## Philosophy
|
||||
|
||||
| Component | Responsibility |
|
||||
|-----------|---------------|
|
||||
| `tokenizer_config.json` (`chat_template`) | Formatting -- how roles become tokens |
|
||||
| `pipeline.json` (`mask`) | Masking -- which roles participate in training |
|
||||
|
||||
A single config file captures the entire pipeline, reusable and version-controllable.
|
||||
|
||||
## Config Structure
|
||||
|
||||
```json
|
||||
{
|
||||
"version": 1,
|
||||
"input": {}, // sections (single) or sources (multi)
|
||||
"mask": {}, // role -> "train" | "mask"
|
||||
"mask_default": "mask",
|
||||
"preprocessing": {},
|
||||
"output": {}
|
||||
}
|
||||
```
|
||||
|
||||
### Section Fields
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `field` | str | -- | JSONL key to read |
|
||||
| `action` | str | -- | `"train"` / `"mask"` / `"$role"` |
|
||||
| `template` | bool | `false` | Apply `chat_template` per message |
|
||||
| `add_special_tokens` | bool | `true` for first non-template section | Add special tokens during encode |
|
||||
|
||||
### Source Fields (multi-output mode)
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `sections` | list[dict] | -- | Same as single-output section list |
|
||||
| `list_field` | bool | `false` | JSONL field holds a list; tokenise each element |
|
||||
| `mask_key` | str | `"{key}_mask"` | Explicit output key for loss mask |
|
||||
|
||||
---
|
||||
|
||||
## Quick Start
|
||||
|
||||
### SFT Chat
|
||||
|
||||
Input JSONL:
|
||||
|
||||
```json
|
||||
{"messages": [{"role": "system", "content": "You are helpful."}, {"role": "user", "content": "Hi"}, {"role": "assistant", "content": "Hello!"}]}
|
||||
```
|
||||
|
||||
Config:
|
||||
|
||||
```json
|
||||
{
|
||||
"input": {
|
||||
"sections": [
|
||||
{"field": "messages", "action": "$role", "template": true}
|
||||
]
|
||||
},
|
||||
"mask": {
|
||||
"system": "mask",
|
||||
"user": "mask",
|
||||
"assistant": "train"
|
||||
},
|
||||
"mask_default": "mask",
|
||||
"preprocessing": {
|
||||
"max_seq_len": 2048
|
||||
},
|
||||
"output": {
|
||||
"storage_format": "bin",
|
||||
"dtype": {"loss_mask": "bool"}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Output keys: `sequence` (int32), `loss_mask` (bool)
|
||||
|
||||
### SFT Instruction
|
||||
|
||||
Input JSONL:
|
||||
|
||||
```json
|
||||
{"prompt": "Translate to French: Hello", "response": "Bonjour"}
|
||||
```
|
||||
|
||||
Config:
|
||||
|
||||
```json
|
||||
{
|
||||
"input": {
|
||||
"sections": [
|
||||
{"field": "prompt", "action": "mask", "add_special_tokens": true},
|
||||
{"field": "response", "action": "train"}
|
||||
]
|
||||
},
|
||||
"mask_default": "mask",
|
||||
"preprocessing": {
|
||||
"max_seq_len": 2048
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Output keys: `sequence`, `loss_mask`
|
||||
|
||||
### Pretrain
|
||||
|
||||
Input JSONL:
|
||||
|
||||
```json
|
||||
{"text": "Artificial Intelligence is a field of computer science..."}
|
||||
```
|
||||
|
||||
Config:
|
||||
|
||||
```json
|
||||
{
|
||||
"input": {
|
||||
"sections": [
|
||||
{"field": "text", "action": "train"}
|
||||
]
|
||||
},
|
||||
"preprocessing": {
|
||||
"max_seq_len": 8192,
|
||||
"min_chars": 100
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Output keys: `sequence` (no `loss_mask` — all tokens trained)
|
||||
|
||||
### DPO
|
||||
|
||||
Input JSONL:
|
||||
|
||||
```json
|
||||
{"chosen": [{"role": "user", "content": "What is 2+2?"}, {"role": "assistant", "content": "4"}], "rejected": [{"role": "user", "content": "What is 2+2?"}, {"role": "assistant", "content": "5"}]}
|
||||
```
|
||||
|
||||
Config:
|
||||
|
||||
```json
|
||||
{
|
||||
"input": {
|
||||
"sources": {
|
||||
"chosen": {
|
||||
"sections": [
|
||||
{"field": "chosen", "action": "$role", "template": true}
|
||||
]
|
||||
},
|
||||
"rejected": {
|
||||
"sections": [
|
||||
{"field": "rejected", "action": "$role", "template": true}
|
||||
]
|
||||
}
|
||||
}
|
||||
},
|
||||
"mask": {
|
||||
"user": "mask",
|
||||
"assistant": "train"
|
||||
},
|
||||
"mask_default": "mask"
|
||||
}
|
||||
```
|
||||
|
||||
Output keys: `chosen`, `chosen_mask`, `rejected`, `rejected_mask`
|
||||
|
||||
### GRPO
|
||||
|
||||
Input JSONL:
|
||||
|
||||
```json
|
||||
{"prompt": [{"role": "user", "content": "What is 2+2?"}], "responses": ["4", "Five", "Four"], "rewards": [1.0, 0.3, 0.8]}
|
||||
```
|
||||
|
||||
Config:
|
||||
|
||||
```json
|
||||
{
|
||||
"input": {
|
||||
"sources": {
|
||||
"prompts": {
|
||||
"sections": [
|
||||
{"field": "prompt", "action": "mask", "template": true}
|
||||
]
|
||||
},
|
||||
"responses": {
|
||||
"sections": [
|
||||
{"field": "responses", "action": "train"}
|
||||
],
|
||||
"list_field": true,
|
||||
"mask_key": "masks"
|
||||
},
|
||||
"rewards": {
|
||||
"sections": [
|
||||
{"field": "rewards", "action": "value"}
|
||||
]
|
||||
}
|
||||
}
|
||||
},
|
||||
"mask": {
|
||||
"user": "mask",
|
||||
"assistant": "train"
|
||||
},
|
||||
"mask_default": "mask"
|
||||
}
|
||||
```
|
||||
|
||||
Output keys: `prompts`, `prompts_mask`, `responses`, `masks`, `rewards` (float32)
|
||||
|
||||
- `action: "value"` — extract raw values from JSONL without tokenisation
|
||||
- `list_field: true` — tokenise each list element independently, then concatenate
|
||||
- `mask_key: "masks"` — rename the auto-generated mask key (default: `responses_mask`)
|
||||
- `prompts_mask` is auto-generated (all masked) and unused by GRPOStrategy
|
||||
|
||||
---
|
||||
|
||||
## Configuration Reference
|
||||
|
||||
### `input`
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `sections` | list[dict] or null | `null` | Section specs for single-output mode |
|
||||
| `sources` | dict[str, dict] or null | `null` | Source specs for multi-output mode (DPO/GRPO) |
|
||||
|
||||
When `sources` is set, `sections` is ignored.
|
||||
|
||||
### `mask`
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `mask` | dict | `{}` | `{role: "train" \| "mask"}` |
|
||||
| `mask_default` | str | `"mask"` | Default action for unlisted roles |
|
||||
|
||||
### `preprocessing`
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `max_seq_len` | int | `2048` | Truncate sequences to this length |
|
||||
| `min_chars` | int | `50` | Skip text-mode items shorter than this |
|
||||
| `max_chars` | int | `2000000` | Skip text-mode items longer than this |
|
||||
| `max_items` | int or null | `null` | Stop after N documents |
|
||||
| `packing_strategy` | str | `"simple"` | Packing strategy: `"simple"`, `"bfd"`, `"bfd_split"` |
|
||||
| `max_packed_len` | int | `8192` | Maximum length of a packed bin |
|
||||
| `truncation_mode` | str | `"keep_start"` | How to truncate sequences: `"keep_start"` or `"keep_end"` |
|
||||
|
||||
### `output`
|
||||
|
||||
| Field | Type | Default | Description |
|
||||
|-------|------|---------|-------------|
|
||||
| `domain_key` | str or null | `null` | JSONL key for domain grouping |
|
||||
| `storage_format` | str | `"bin"` | `"bin"` (mmap) or `"h5"` |
|
||||
| `max_tokens_per_shard` | int | `100000000` | Flush threshold in cumulative tokens |
|
||||
| `dtype` | dict[str, str] | `{}` | Per-key tensor dtype override (e.g. `{"loss_mask": "bool"}`) |
|
||||
| `position_ids_mode` | str | `"doc_reset"` | How to compute position_ids: `"none"`, `"doc_reset"`, `"continuous"` |
|
||||
|
||||
---
|
||||
|
||||
## Mask Algorithm
|
||||
|
||||
### Template mode (`template: true`)
|
||||
|
||||
1. Prepend BOS token (masked)
|
||||
2. For each message in the field's array:
|
||||
1. Render through `chat_template` for that single message
|
||||
2. Encode rendered text
|
||||
3. Apply mask rule for the message's role
|
||||
|
||||
### Non-template mode
|
||||
|
||||
Encode the field value as text. Mask value is 1 (train) or 0 (mask) per the section's `action`.
|
||||
|
||||
### Text config detection
|
||||
|
||||
When no section uses `template` and all sections have `action: "train"`, the builder omits `loss_mask` from the output — all tokens are trained.
|
||||
|
||||
---
|
||||
|
||||
## Output Layout
|
||||
|
||||
### Single-Shard (`bin`)
|
||||
|
||||
```
|
||||
output/
|
||||
__default__/
|
||||
shard_0000/
|
||||
meta.json
|
||||
sequence.bin
|
||||
loss_mask.bin
|
||||
wiki/
|
||||
shard_0000/
|
||||
meta.json
|
||||
sequence.bin
|
||||
loss_mask.bin
|
||||
```
|
||||
|
||||
### Multi-Shard (`bin`)
|
||||
|
||||
When `max_tokens_per_shard` is exceeded:
|
||||
|
||||
```
|
||||
output/
|
||||
__default__/
|
||||
shard_0000/
|
||||
meta.json
|
||||
sequence.bin
|
||||
loss_mask.bin
|
||||
shard_0001/
|
||||
meta.json
|
||||
sequence.bin
|
||||
loss_mask.bin
|
||||
```
|
||||
|
||||
For `bin` format, `MmapStore` discovers all shards under the domain directory via `rglob("meta.json")`. For `h5` format, `H5Store` discovers `.h5`/`.hdf5` files via recursive glob.
|
||||
|
||||
---
|
||||
|
||||
## CLI
|
||||
|
||||
```bash
|
||||
# SFT
|
||||
python scripts/tools/preprocess.py data/sft/*.jsonl -o output/sft/ -c configs/sft_chat.json
|
||||
|
||||
# DPO
|
||||
python scripts/tools/preprocess.py data/dpo/*.jsonl -o output/dpo/ -c configs/dpo.json --tokenizer_path params
|
||||
|
||||
# GRPO
|
||||
python scripts/tools/preprocess.py data/grpo/*.jsonl -o output/grpo/ -c configs/grpo.json
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Python API
|
||||
|
||||
```python
|
||||
from astrai.preprocessing.pipeline import Pipeline
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
|
||||
config = PipelineConfig.from_file("sft.json")
|
||||
Pipeline(
|
||||
config,
|
||||
["data_part1.jsonl", "data_part2.jsonl"],
|
||||
output_dir="output/",
|
||||
tokenizer_path="params",
|
||||
).run()
|
||||
```
|
||||
|
||||
> Document Update Time: 2026-07-09
|
||||
+50
-52
@@ -1,37 +1,17 @@
|
||||
# Training
|
||||
|
||||
## Model Architecture
|
||||
## Contents
|
||||
|
||||
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](#autoregression)
|
||||
- [Causal Mask](#causal-mask)
|
||||
- [Rotary Position Embedding (RoPE)](#rotary-position-embedding-rope)
|
||||
- [Training Loop](#training-loop)
|
||||
- [Strategies](#strategies) — SEQ, SFT, DPO, GRPO
|
||||
- [LR Schedulers](#lr-schedulers)
|
||||
- [Gradient Checkpointing](#gradient-checkpointing)
|
||||
- [Checkpoint](#checkpoint)
|
||||
- [TrainContextBuilder](#traincontextbuilder-builder-pattern)
|
||||
- [Training CLI](#training-cli)
|
||||
|
||||
### Autoregression
|
||||
|
||||
@@ -69,14 +49,18 @@ Two-level loop: **epoch** → **batch**. Optimizer step fires every `grad_accum_
|
||||
|
||||
```
|
||||
on_train_begin
|
||||
model.train()
|
||||
on_epoch_begin
|
||||
for batch in dataloader:
|
||||
on_batch_begin
|
||||
with executor.accumulate(model):
|
||||
loss = strategy(batch)
|
||||
loss = strategy.compute_loss(batch)
|
||||
context.loss = loss.item()
|
||||
stand_loss = loss / executor.grad_accum_steps
|
||||
executor.backward(stand_loss)
|
||||
iteration += 1
|
||||
context.consumed_samples += (
|
||||
context.config.batch_per_device * context.world_size
|
||||
)
|
||||
on_batch_end
|
||||
|
||||
if executor.sync_gradients:
|
||||
@@ -94,11 +78,15 @@ on_train_end
|
||||
| Hook | Fires | Default callback |
|
||||
|------|-------|-----------------|
|
||||
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback` |
|
||||
| `on_optimizer_step` | Every accumulation window | `GradientClippingCallback`, `ValidationCallback` |
|
||||
| `on_batch_end` | Every batch | `CheckpointCallback`, `MetricLoggerCallback`, `ProgressBarCallback` |
|
||||
| `on_train_end` | Training ends | `CheckpointCallback`, `MetricLoggerCallback` (final save) |
|
||||
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
|
||||
| `on_batch_begin` | Every batch | — |
|
||||
| `on_optimizer_step` | Every accumulation window | `GradientClippingCallback`, `MetricCallback`, `ProgressBarCallback` |
|
||||
| `on_batch_end` | Every batch | `CheckpointCallback` |
|
||||
| `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` |
|
||||
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
|
||||
| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` |
|
||||
|
||||
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric_logger` (JSONL, rank-0), `progress_bar` (tqdm), `gradient_clipping`, `validation` (periodic validation on val_dataset).
|
||||
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping` (always registered; computes grad norm, clips only when `max_grad_norm` is not `None`).
|
||||
|
||||
## Strategies
|
||||
|
||||
@@ -110,7 +98,7 @@ $$
|
||||
L_{\text{PT}} = -\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta)
|
||||
$$
|
||||
|
||||
Keys: `input_ids`, `target_ids`
|
||||
Keys: `input_ids`, `target_ids`. Optional: `label_smoothing`.
|
||||
|
||||
### SFT (Supervised Fine-Tuning)
|
||||
|
||||
@@ -120,7 +108,7 @@ $$
|
||||
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`
|
||||
Keys: `input_ids`, `target_ids`, `loss_mask`, `position_ids`. Optional: `label_smoothing`.
|
||||
|
||||
### DPO (Direct Preference Optimization)
|
||||
|
||||
@@ -130,21 +118,31 @@ $$
|
||||
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`.
|
||||
Parameters: `beta=0.1`, `reduction="sum"`. Keys: `chosen`, `rejected`, `chosen_mask`, `rejected_mask`.
|
||||
|
||||
### GRPO (Group Relative Policy Optimization)
|
||||
|
||||
On-policy PPO with group-normalized advantages:
|
||||
Token-level PPO with group-normalized advantages. Advantages are derived from
|
||||
scalar per-response rewards, group-normalized, and broadcast across all response
|
||||
tokens. Only response tokens contribute to the loss (prompt tokens are masked
|
||||
out):
|
||||
|
||||
$$
|
||||
\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]
|
||||
L_{\text{GRPO}} = -\mathbb{E}_t\left[\min\left(\rho_t A,\; \text{clip}\left(\rho_t, 1-\epsilon, 1+\epsilon\right)A\right)\right] + \lambda \cdot \mathbb{E}_t\left[\frac{\pi_{\text{ref}}}{\pi_\theta} - \log\frac{\pi_{\text{ref}}}{\pi_\theta} - 1\right]
|
||||
$$
|
||||
|
||||
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`, `sync_interval=200`.
|
||||
where $\rho_t = \pi_\theta(a_t|s_t) / \pi_{\text{old}}(a_t|s_t)$ is the
|
||||
per-token importance sampling ratio against the behaviour policy
|
||||
(`old_model`, synced externally between data-generation rounds) and the
|
||||
expectations are over valid response tokens. The KL term regularises
|
||||
$\pi_\theta$ towards a frozen reference model (`ref_model`, typically
|
||||
the SFT checkpoint).
|
||||
|
||||
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`. External sync of `old_model` weights via `sync_old_model()` between data-generation rounds.
|
||||
|
||||
Keys: `prompts`, `responses`, `masks`, `rewards`.
|
||||
|
||||
@@ -154,8 +152,9 @@ Keys: `prompts`, `responses`, `masks`, `rewards`.
|
||||
|------|-------|-------------|
|
||||
| Cosine | `CosineScheduler` | Linear warmup → cosine decay to `min_rate` |
|
||||
| SGDR | `SGDRScheduler` | Cosine annealing with warm restarts (`t_mult=2`) |
|
||||
| WSD | `WSDScheduler` | Warmup-Stable-Decay with sqrt cooldown |
|
||||
|
||||
Created by `SchedulerFactory.create(optimizer, schedule_type, **kwargs)`.
|
||||
Created by `SchedulerFactory.create(schedule_type, optimizer, **kwargs)`. Valid types: `"cosine"`, `"sgdr"`, `"wsd"`. Omit to use no scheduler.
|
||||
|
||||
## Gradient Checkpointing
|
||||
|
||||
@@ -171,9 +170,9 @@ Callback wraps each `DecoderBlock.forward` with `torch.utils.checkpoint.checkpoi
|
||||
## Checkpoint
|
||||
|
||||
```
|
||||
Checkpoint(state_dict, epoch, iteration, extra, meta, config)
|
||||
├── save(save_dir) rank-0 only: meta.json (epoch/iteration/timestamp) + config.json (model config) + state_dict.safetensors + optional {key}.pt (optimizer.pt, scheduler.pt)
|
||||
└── load(save_dir) broadcasts metadata from rank-0
|
||||
Checkpoint(state_dict, epoch, consumed_samples, extra, meta, config)
|
||||
├── save(save_dir) rank-0 only: meta.json (epoch/consumed_samples/timestamp) + config.json (model config) + model.safetensors + optional {key}.pt (optimizer.pt, scheduler.pt)
|
||||
└── load(save_dir, broadcast=False) loads from local disk; set broadcast=True to broadcast metadata from rank-0
|
||||
```
|
||||
|
||||
Optimizer/scheduler state persisted by default via `Checkpoint.extra`.
|
||||
@@ -194,7 +193,7 @@ context = (
|
||||
- Creates executor via `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)`
|
||||
- Calls `executor.prepare(model, optimizer, dataloader, scheduler)` for model distribution (e.g. DDP) + gradient accumulation wrappers
|
||||
- Creates `ResumableDistributedSampler` for shuffle+resume
|
||||
- Builds strategy via `StrategyFactory.create(train_type, ...)`
|
||||
- Builds strategy via `StrategyFactory.create(train_type, model, device, **kwargs)`
|
||||
|
||||
## Training CLI
|
||||
|
||||
@@ -203,6 +202,7 @@ export CUDA_VISIBLE_DEVICES=0,1,2,3
|
||||
|
||||
nohup python scripts/tools/train.py \
|
||||
--nprocs=4 \
|
||||
--parallel_mode=ddp \
|
||||
--train_type=seq \
|
||||
--data_root_path=/path/to/dataset \
|
||||
--param_path=/path/to/model \
|
||||
@@ -211,9 +211,7 @@ nohup python scripts/tools/train.py \
|
||||
--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 \
|
||||
--weight_decay=0.1 \
|
||||
--window_size=2048 \
|
||||
--ckpt_interval=10000 \
|
||||
--ckpt_dir=./checkpoint \
|
||||
@@ -224,4 +222,4 @@ nohup python scripts/tools/train.py \
|
||||
|
||||
Full parameter reference at [params.md](params.md).
|
||||
|
||||
> Document Update Time: 2026-05-28
|
||||
> Document Update Time: 2026-07-19
|
||||
|
||||
+77
-13
@@ -1,34 +1,98 @@
|
||||
__version__ = "1.3.7"
|
||||
__version__ = "1.3.9"
|
||||
__author__ = "ViperEkura"
|
||||
|
||||
from astrai.config import (
|
||||
AutoRegressiveLMConfig,
|
||||
BaseModelConfig,
|
||||
ConfigFactory,
|
||||
EncoderConfig,
|
||||
PipelineConfig,
|
||||
TrainConfig,
|
||||
)
|
||||
from astrai.dataset import DatasetFactory
|
||||
from astrai.dataset import (
|
||||
BaseDataset,
|
||||
DatasetFactory,
|
||||
RDSampler,
|
||||
Store,
|
||||
StoreFactory,
|
||||
)
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.inference import (
|
||||
GenerationRequest,
|
||||
InferenceEngine,
|
||||
ProtocolHandler,
|
||||
SamplingPipeline,
|
||||
get_app,
|
||||
run_server,
|
||||
sample,
|
||||
)
|
||||
from astrai.model import (
|
||||
AutoModel,
|
||||
AutoRegressiveLM,
|
||||
EmbeddingEncoder,
|
||||
LoRAConfig,
|
||||
inject_lora,
|
||||
)
|
||||
from astrai.parallel import (
|
||||
ExecutorFactory,
|
||||
get_rank,
|
||||
get_world_size,
|
||||
only_on_rank,
|
||||
spawn_parallel_fn,
|
||||
)
|
||||
from astrai.preprocessing import Pipeline, filter_by_length
|
||||
from astrai.serialization import Checkpoint
|
||||
from astrai.tokenize import AutoTokenizer, ChatTemplate
|
||||
from astrai.trainer import (
|
||||
BaseScheduler,
|
||||
BaseStrategy,
|
||||
CallbackFactory,
|
||||
SchedulerFactory,
|
||||
StrategyFactory,
|
||||
TrainCallback,
|
||||
Trainer,
|
||||
)
|
||||
from astrai.model import AutoModel, AutoRegressiveLM
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
from astrai.trainer import CallbackFactory, SchedulerFactory, StrategyFactory, Trainer
|
||||
|
||||
__all__ = [
|
||||
"AutoRegressiveLM",
|
||||
"AutoRegressiveLMConfig",
|
||||
"EncoderConfig",
|
||||
"TrainConfig",
|
||||
"DatasetFactory",
|
||||
"AutoModel",
|
||||
"AutoTokenizer",
|
||||
"BaseDataset",
|
||||
"BaseFactory",
|
||||
"BaseModelConfig",
|
||||
"BaseScheduler",
|
||||
"BaseStrategy",
|
||||
"CallbackFactory",
|
||||
"ChatTemplate",
|
||||
"Checkpoint",
|
||||
"ConfigFactory",
|
||||
"DatasetFactory",
|
||||
"EmbeddingEncoder",
|
||||
"EncoderConfig",
|
||||
"ExecutorFactory",
|
||||
"GenerationRequest",
|
||||
"InferenceEngine",
|
||||
"Trainer",
|
||||
"CallbackFactory",
|
||||
"StrategyFactory",
|
||||
"LoRAConfig",
|
||||
"Pipeline",
|
||||
"PipelineConfig",
|
||||
"ProtocolHandler",
|
||||
"RDSampler",
|
||||
"SamplingPipeline",
|
||||
"SchedulerFactory",
|
||||
"BaseFactory",
|
||||
"AutoModel",
|
||||
"Store",
|
||||
"StoreFactory",
|
||||
"StrategyFactory",
|
||||
"TrainCallback",
|
||||
"TrainConfig",
|
||||
"Trainer",
|
||||
"filter_by_length",
|
||||
"get_app",
|
||||
"get_rank",
|
||||
"get_world_size",
|
||||
"inject_lora",
|
||||
"only_on_rank",
|
||||
"run_server",
|
||||
"sample",
|
||||
"spawn_parallel_fn",
|
||||
]
|
||||
|
||||
@@ -4,13 +4,22 @@ from astrai.config.model_config import (
|
||||
ConfigFactory,
|
||||
EncoderConfig,
|
||||
)
|
||||
from astrai.config.preprocess_config import (
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
PipelineConfig,
|
||||
ProcessingConfig,
|
||||
)
|
||||
from astrai.config.train_config import TrainConfig
|
||||
|
||||
__all__ = [
|
||||
# Model configuration
|
||||
"BaseModelConfig",
|
||||
"AutoRegressiveLMConfig",
|
||||
"EncoderConfig",
|
||||
"ConfigFactory",
|
||||
"TrainConfig",
|
||||
"InputConfig",
|
||||
"OutputConfig",
|
||||
"PipelineConfig",
|
||||
"ProcessingConfig",
|
||||
]
|
||||
|
||||
+13
-1
@@ -1,6 +1,7 @@
|
||||
import json
|
||||
from dataclasses import MISSING, dataclass, fields
|
||||
from typing import Any, Dict, Optional, Self, get_type_hints
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Optional, Self, Union, get_type_hints
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -83,4 +84,15 @@ class BaseConfig:
|
||||
return value
|
||||
if isinstance(value, target_type):
|
||||
return value
|
||||
if isinstance(value, dict) and issubclass(target_type, BaseConfig):
|
||||
return target_type.from_dict(value)
|
||||
raise TypeError
|
||||
|
||||
@classmethod
|
||||
def from_file(cls, path: Union[str, Path]) -> Self:
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
return cls.from_dict(json.load(f))
|
||||
|
||||
def to_file(self, path: Union[str, Path]):
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
json.dump(self.to_dict(), f, indent=2, ensure_ascii=False)
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional, Self
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from astrai.config.base import BaseConfig
|
||||
from astrai.factory import BaseFactory
|
||||
@@ -21,18 +20,7 @@ class BaseModelConfig(BaseConfig):
|
||||
"""Base config with ``model_type`` dispatch and file I/O."""
|
||||
|
||||
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)
|
||||
neftune_alpha: float = 0.0
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -83,10 +71,12 @@ class EncoderConfig(BaseModelConfig):
|
||||
rope_theta: Optional[float] = None
|
||||
rope_scaling: Optional[dict] = None
|
||||
|
||||
attn_type: str = "gqa"
|
||||
n_heads: Optional[int] = None
|
||||
n_kv_heads: Optional[int] = None
|
||||
use_qk_norm: Optional[bool] = None
|
||||
use_gated_attention: Optional[bool] = None
|
||||
|
||||
ffn_type: str = "mlp"
|
||||
pooling_type: Optional[str] = None
|
||||
normalize_embeddings: Optional[bool] = None
|
||||
|
||||
@@ -0,0 +1,109 @@
|
||||
"""Pipeline configuration for JSONL preprocessing.
|
||||
|
||||
Supports single-sequence (SFT/pretrain) and multi-output (DPO/GRPO)
|
||||
modes, both driven declaratively through ``input.sections`` or
|
||||
``input.sources``.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from astrai.config.base import BaseConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class InputConfig(BaseConfig):
|
||||
"""Declarative input mapping.
|
||||
|
||||
Single-output mode (backward-compatible)::
|
||||
|
||||
{"input": {"sections": [{"field": "messages", ...}]}}
|
||||
|
||||
Multi-output mode (DPO / GRPO)::
|
||||
|
||||
{"input": {"sources": {
|
||||
"chosen": {"sections": [{"field": "chosen", ...}]},
|
||||
"rejected": {"sections": [{"field": "rejected", ...}]},
|
||||
}}}
|
||||
"""
|
||||
|
||||
sections: Optional[List[Dict]] = None
|
||||
sources: Optional[Dict[str, Dict]] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProcessingConfig(BaseConfig):
|
||||
"""Processing configuration.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
max_seq_len : int
|
||||
Maximum sequence length (default: 2048).
|
||||
min_chars : int
|
||||
Minimum number of characters to keep (default: 50).
|
||||
max_chars : int
|
||||
Maximum number of characters to keep (default: 2_000_000).
|
||||
max_items : Optional[int]
|
||||
Maximum number of items to process (default: None, unlimited).
|
||||
packing_strategy : str
|
||||
How to pack sequences into a contiguous stream.
|
||||
|
||||
- ``"simple"``: sequential concatenation (default, backward compatible).
|
||||
- ``"bfd"``: best-fit decreasing bin packing, minimises wasted tokens.
|
||||
- ``"bfd_split"``: BFD with over-length sequences split into chunks.
|
||||
max_packed_len : int
|
||||
Maximum length of a packed bin. Sequences longer than this are
|
||||
truncated or split depending on ``packing_strategy`` (default: 8192).
|
||||
truncation_mode : str
|
||||
How to truncate sequences longer than ``max_packed_len``.
|
||||
|
||||
- ``"keep_start"``: keep the first ``max_packed_len`` tokens (default).
|
||||
- ``"keep_end"``: keep the last ``max_packed_len`` tokens.
|
||||
"""
|
||||
|
||||
max_seq_len: int = 2048
|
||||
min_chars: int = 50
|
||||
max_chars: int = 2_000_000
|
||||
max_items: Optional[int] = None
|
||||
packing_strategy: str = "simple"
|
||||
max_packed_len: int = 8192
|
||||
truncation_mode: str = "keep_start"
|
||||
|
||||
|
||||
@dataclass
|
||||
class OutputConfig(BaseConfig):
|
||||
"""Output configuration.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
domain_key : Optional[str]
|
||||
Domain key for the output store (default: None).
|
||||
storage_format : str
|
||||
Storage format, one of ``"bin"``, ``"jsonl"`` (default: ``"bin"``).
|
||||
max_tokens_per_shard : int
|
||||
Maximum tokens per shard before splitting (default: 100_000_000).
|
||||
dtype : Dict[str, str]
|
||||
Per-key dtype overrides, e.g. ``{"input_ids": "int32"}`` (default: {}).
|
||||
position_ids_mode : Optional[str]
|
||||
How to compute position_ids in packed sequences.
|
||||
|
||||
- ``"none"``: do not generate (default).
|
||||
- ``"doc_reset"``: reset to 0 at each document boundary.
|
||||
- ``"continuous"``: sequential 0, 1, 2, ... (pretrain, single doc).
|
||||
"""
|
||||
|
||||
domain_key: Optional[str] = None
|
||||
storage_format: str = "bin"
|
||||
max_tokens_per_shard: int = 100_000_000
|
||||
dtype: Dict[str, str] = field(default_factory=dict)
|
||||
position_ids_mode: str = "doc_reset"
|
||||
|
||||
|
||||
@dataclass
|
||||
class PipelineConfig(BaseConfig):
|
||||
version: int = 1
|
||||
input: InputConfig = field(default_factory=InputConfig)
|
||||
mask: Dict[str, str] = field(default_factory=dict)
|
||||
mask_default: str = "mask"
|
||||
preprocessing: ProcessingConfig = field(default_factory=ProcessingConfig)
|
||||
output: OutputConfig = field(default_factory=OutputConfig)
|
||||
@@ -1,5 +1,5 @@
|
||||
from dataclasses import dataclass, field, fields
|
||||
from typing import Callable, List, Optional
|
||||
from typing import Any, Callable, Dict, List, Optional
|
||||
|
||||
import torch.nn as nn
|
||||
from torch.optim import Optimizer
|
||||
@@ -37,24 +37,29 @@ class TrainConfig(BaseConfig):
|
||||
grad_accum_steps: int = field(
|
||||
default=1, metadata={"help": "Number of iterations between steps."}
|
||||
)
|
||||
max_grad_norm: float = field(
|
||||
default=1.0, metadata={"help": "Maximum gradient norm."}
|
||||
max_grad_norm: Optional[float] = field(
|
||||
default=None,
|
||||
metadata={"help": "Maximum gradient norm. None disables clipping."},
|
||||
)
|
||||
gradient_checkpointing_modules: list = field(
|
||||
gradient_checkpointing_modules: List[str] = field(
|
||||
default_factory=list,
|
||||
metadata={"help": "Module types to enable activation checkpointing for."},
|
||||
)
|
||||
|
||||
# checkpoint setting
|
||||
start_epoch: int = field(default=0, metadata={"help": "Start epoch for training."})
|
||||
start_batch: int = field(
|
||||
default=0, metadata={"help": "Start batch iteration for training."}
|
||||
start_samples: int = field(
|
||||
default=0,
|
||||
metadata={
|
||||
"help": "Start samples count (per rank). Superseded by checkpoint consumed_samples."
|
||||
},
|
||||
)
|
||||
ckpt_dir: str = field(
|
||||
default="./checkpoint", metadata={"help": "Checkpoint directory."}
|
||||
)
|
||||
ckpt_interval: int = field(
|
||||
default=5000, metadata={"help": "Number of iterations between checkpoints."}
|
||||
default=5000,
|
||||
metadata={"help": "Number of optimizer steps between checkpoints."},
|
||||
)
|
||||
|
||||
# lora setting
|
||||
@@ -67,12 +72,8 @@ class TrainConfig(BaseConfig):
|
||||
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"],
|
||||
default_factory=lambda: ["loss", "lr", "grad_norm"],
|
||||
metadata={"help": "Metrics to record during training."},
|
||||
)
|
||||
|
||||
@@ -87,6 +88,10 @@ class TrainConfig(BaseConfig):
|
||||
pin_memory: bool = field(
|
||||
default=False, metadata={"help": "Pin memory for dataloader."}
|
||||
)
|
||||
collate_fn: Optional[Callable[[List[Any]], Any]] = field(
|
||||
default=None,
|
||||
metadata={"help": "Collate function for dataloader (e.g. dpo_collate_fn)."},
|
||||
)
|
||||
|
||||
# distributed training
|
||||
nprocs: int = field(
|
||||
@@ -118,16 +123,26 @@ class TrainConfig(BaseConfig):
|
||||
val_dataset: Optional[Dataset] = field(
|
||||
default=None, metadata={"help": "Dataset for validation."}
|
||||
)
|
||||
val_split: Optional[float] = field(
|
||||
default=None,
|
||||
metadata={
|
||||
"help": "Ratio to split from training dataset for validation (e.g. 0.05). Ignored if val_dataset is set."
|
||||
},
|
||||
)
|
||||
val_step: int = field(
|
||||
default=1000,
|
||||
metadata={"help": "Number of optimizer steps between validation runs."},
|
||||
)
|
||||
neftune_alpha: float = field(
|
||||
default=0.0,
|
||||
metadata={"help": "NEFTune noise alpha (0=disabled, typical: 5.0)."},
|
||||
)
|
||||
|
||||
executor_kwargs: dict = field(
|
||||
executor_kwargs: Dict[str, Any] = field(
|
||||
default_factory=dict,
|
||||
metadata={"help": "Extra kwargs passed to ExecutorFactory.create()."},
|
||||
)
|
||||
extra_kwargs: dict = field(
|
||||
extra_kwargs: Dict[str, Any] = field(
|
||||
default_factory=dict, metadata={"help": "Other arguments."}
|
||||
)
|
||||
|
||||
|
||||
@@ -1,14 +1,21 @@
|
||||
from astrai.dataset.dataset import (
|
||||
BaseDataset,
|
||||
DatasetFactory,
|
||||
dpo_collate_fn,
|
||||
grpo_collate_fn,
|
||||
)
|
||||
from astrai.dataset.sampler import ResumableDistributedSampler
|
||||
from astrai.dataset.sampler import RDSampler
|
||||
from astrai.dataset.storage import (
|
||||
H5Store,
|
||||
JsonlStore,
|
||||
MmapStore,
|
||||
Recordable,
|
||||
Store,
|
||||
StoreFactory,
|
||||
Streamable,
|
||||
detect_format,
|
||||
)
|
||||
from astrai.serialization import (
|
||||
load_bin,
|
||||
load_h5,
|
||||
save_bin,
|
||||
@@ -18,14 +25,19 @@ from astrai.dataset.storage import (
|
||||
__all__ = [
|
||||
"BaseDataset",
|
||||
"DatasetFactory",
|
||||
"dpo_collate_fn",
|
||||
"grpo_collate_fn",
|
||||
"Store",
|
||||
"Streamable",
|
||||
"Recordable",
|
||||
"StoreFactory",
|
||||
"H5Store",
|
||||
"MmapStore",
|
||||
"JsonlStore",
|
||||
"detect_format",
|
||||
"save_h5",
|
||||
"load_h5",
|
||||
"save_bin",
|
||||
"load_bin",
|
||||
"ResumableDistributedSampler",
|
||||
"RDSampler",
|
||||
]
|
||||
|
||||
+404
-218
@@ -1,7 +1,31 @@
|
||||
"""Dataset implementations with factory pattern for training."""
|
||||
"""Dataset implementations for training.
|
||||
|
||||
Composition over inheritance — every dataset is a thin wrapper that
|
||||
binds a :class:`Store` to a particular train-type's key mapping. All
|
||||
sample-id → token/record indexing lives on the Store; datasets never
|
||||
know about window/stride math or segment layouts.
|
||||
|
||||
Class hierarchy:
|
||||
|
||||
BaseDataset (ABC) — holds a Store, exposes __len__/keys,
|
||||
overrides __getitem__
|
||||
├── SEQDataset — next-token prediction (stream)
|
||||
├── SFTDataset — loss-mask + position_ids (stream)
|
||||
├── DPODataset — chosen/rejected pairs (record)
|
||||
└── GRPODataset — prompt + response group (record)
|
||||
|
||||
``DatasetFactory.load(train_type, load_path, window_size, stride, …)``
|
||||
builds the Store (auto-detecting format) before constructing the
|
||||
matching dataset. Passing ``store=`` skips Store construction.
|
||||
|
||||
When a record dataset (DPO) reads from raw JSONL, a *processor*
|
||||
function (pure ``record -> Dict[str, Tensor]``) is forwarded to
|
||||
:class:`JsonlStore` so tokenisation happens on the fly.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Dict, List, Optional
|
||||
from functools import partial
|
||||
from typing import Callable, Dict, List, Optional
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
@@ -13,296 +37,458 @@ from astrai.dataset.storage import (
|
||||
detect_format,
|
||||
)
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
|
||||
def dpo_tokenize(
|
||||
record: dict,
|
||||
tokenizer,
|
||||
max_len: int = 2048,
|
||||
) -> Optional[dict]:
|
||||
"""Tokenize one DPO record into chosen/rejected + masks.
|
||||
|
||||
Applies the tokenizer's chat template so token sequences match the
|
||||
SFT checkpoint's format. Prompt is rendered with
|
||||
``add_generation_prompt=True``; chosen/rejected are appended as a
|
||||
single assistant turn.
|
||||
|
||||
Accepts:
|
||||
|
||||
- Flat: ``{"prompt": str, "chosen": str, "rejected": str}``
|
||||
- Conv: ``{"prompt": [{role, content}, ...], "chosen": [...], ...}``
|
||||
- Legacy: ``{"input": str, "chosen": str, "rejected": str}``
|
||||
|
||||
No packing, no ``position_ids`` — DPO sequences are independent.
|
||||
"""
|
||||
prompt = record.get("prompt") or record.get("input")
|
||||
chosen = record.get("chosen")
|
||||
rejected = record.get("rejected")
|
||||
if prompt is None or chosen is None or rejected is None:
|
||||
return None
|
||||
|
||||
prompt_messages = _to_messages(prompt)
|
||||
chosen_text = _extract_text(chosen)
|
||||
rejected_text = _extract_text(rejected)
|
||||
if chosen_text is None or rejected_text is None:
|
||||
return None
|
||||
chosen_messages = prompt_messages + [{"role": "assistant", "content": chosen_text}]
|
||||
rejected_messages = prompt_messages + [
|
||||
{"role": "assistant", "content": rejected_text}
|
||||
]
|
||||
|
||||
prompt_ids = tokenizer.apply_chat_template(
|
||||
prompt_messages, tokenize=True, add_generation_prompt=True
|
||||
)
|
||||
ch_ids = tokenizer.apply_chat_template(
|
||||
chosen_messages, tokenize=True, add_generation_prompt=False
|
||||
)
|
||||
re_ids = tokenizer.apply_chat_template(
|
||||
rejected_messages, tokenize=True, add_generation_prompt=False
|
||||
)
|
||||
|
||||
full_ch = ch_ids[:max_len]
|
||||
full_re = re_ids[:max_len]
|
||||
|
||||
prompt_len = min(len(prompt_ids), max_len)
|
||||
ch_mask = [0] * prompt_len + [1] * max(0, len(full_ch) - prompt_len)
|
||||
ch_mask = ch_mask[:max_len]
|
||||
re_mask = [0] * prompt_len + [1] * max(0, len(full_re) - prompt_len)
|
||||
re_mask = re_mask[:max_len]
|
||||
|
||||
return {
|
||||
"chosen": full_ch,
|
||||
"rejected": full_re,
|
||||
"chosen_mask": ch_mask,
|
||||
"rejected_mask": re_mask,
|
||||
}
|
||||
|
||||
|
||||
def _to_messages(value) -> list:
|
||||
"""Accept str or conversation list; return message list."""
|
||||
if isinstance(value, str):
|
||||
return [{"role": "user", "content": value}]
|
||||
if isinstance(value, list):
|
||||
return value
|
||||
return [{"role": "user", "content": str(value)}]
|
||||
|
||||
|
||||
def _extract_text(value) -> Optional[str]:
|
||||
"""Accept str or conversation list; return plain text."""
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
if isinstance(value, list):
|
||||
return "".join(m.get("content", "") for m in value if isinstance(m, dict))
|
||||
return None
|
||||
|
||||
|
||||
def dpo_processor(
|
||||
record: dict,
|
||||
tokenizer,
|
||||
max_len: int = 2048,
|
||||
) -> Dict[str, Tensor]:
|
||||
"""DPO processor: wraps :func:`dpo_tokenize` and returns tensors."""
|
||||
result = dpo_tokenize(record, tokenizer, max_len=max_len)
|
||||
if result is None:
|
||||
raise ValueError(f"Malformed DPO record: {list(record.keys())}")
|
||||
return {
|
||||
"chosen": torch.tensor(result["chosen"], dtype=torch.int32),
|
||||
"rejected": torch.tensor(result["rejected"], dtype=torch.int32),
|
||||
"chosen_mask": torch.tensor(result["chosen_mask"], dtype=torch.bool),
|
||||
"rejected_mask": torch.tensor(result["rejected_mask"], dtype=torch.bool),
|
||||
}
|
||||
|
||||
|
||||
def dpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
||||
"""Collate variable-length DPO samples into padded 2-D tensors.
|
||||
|
||||
Input: list of dicts, each with:
|
||||
- chosen: [C_i]
|
||||
- rejected: [R_i]
|
||||
- chosen_mask: [C_i]
|
||||
- rejected_mask: [R_i]
|
||||
|
||||
Output (padded to the max length across chosen/rejected within the batch):
|
||||
- chosen: [B, S_max]
|
||||
- rejected: [B, S_max]
|
||||
- chosen_mask: [B, S_max]
|
||||
- rejected_mask: [B, S_max]
|
||||
"""
|
||||
B = len(batch)
|
||||
S_max = max(b["chosen"].size(0) for b in batch)
|
||||
S_max = max(S_max, max(b["rejected"].size(0) for b in batch))
|
||||
|
||||
chosen = torch.zeros(B, S_max, dtype=torch.long)
|
||||
rejected = torch.zeros(B, S_max, dtype=torch.long)
|
||||
chosen_mask = torch.zeros(B, S_max, dtype=torch.bool)
|
||||
rejected_mask = torch.zeros(B, S_max, dtype=torch.bool)
|
||||
|
||||
for i, b in enumerate(batch):
|
||||
c_len = b["chosen"].size(0)
|
||||
r_len = b["rejected"].size(0)
|
||||
chosen[i, :c_len] = b["chosen"]
|
||||
rejected[i, :r_len] = b["rejected"]
|
||||
chosen_mask[i, :c_len] = b["chosen_mask"]
|
||||
rejected_mask[i, :r_len] = b["rejected_mask"]
|
||||
|
||||
return {
|
||||
"chosen": chosen,
|
||||
"rejected": rejected,
|
||||
"chosen_mask": chosen_mask,
|
||||
"rejected_mask": rejected_mask,
|
||||
}
|
||||
|
||||
|
||||
def grpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
|
||||
"""Collate variable-length GRPO samples into padded 3-D tensors.
|
||||
|
||||
Input: list of dicts, each with:
|
||||
- prompts: [P_i]
|
||||
- responses: list of G tensors, each [R_ij]
|
||||
- masks: list of G tensors, each [R_ij]
|
||||
- rewards: [G]
|
||||
|
||||
Output:
|
||||
- prompts: [B, P_max]
|
||||
- responses: [B, G, R_max]
|
||||
- masks: [B, G, R_max]
|
||||
- rewards: [B, G]
|
||||
"""
|
||||
B = len(batch)
|
||||
G = len(batch[0]["responses"])
|
||||
P_max = max(b["prompts"].size(0) for b in batch)
|
||||
R_max = max(r.size(0) for b in batch for r in b["responses"])
|
||||
|
||||
prompts = torch.zeros(B, P_max, dtype=torch.long)
|
||||
responses = torch.zeros(B, G, R_max, dtype=torch.long)
|
||||
masks = torch.zeros(B, G, R_max, dtype=torch.bool)
|
||||
rewards = torch.zeros(B, G, dtype=torch.float32)
|
||||
|
||||
for i, b in enumerate(batch):
|
||||
p_len = b["prompts"].size(0)
|
||||
prompts[i, :p_len] = b["prompts"]
|
||||
rewards[i, : b["rewards"].size(0)] = b["rewards"]
|
||||
for g in range(min(G, len(b["responses"]))):
|
||||
r_len = b["responses"][g].size(0)
|
||||
responses[i, g, :r_len] = b["responses"][g]
|
||||
if g < len(b["masks"]):
|
||||
masks[i, g, :r_len] = b["masks"][g]
|
||||
|
||||
return {
|
||||
"prompts": prompts,
|
||||
"responses": responses,
|
||||
"masks": masks,
|
||||
"rewards": rewards,
|
||||
}
|
||||
|
||||
|
||||
def validate_keys(store: Store, required: List[str]) -> None:
|
||||
"""Raise ``KeyError`` if *store* is missing any *required* key."""
|
||||
if not required:
|
||||
return
|
||||
actual = set(store.keys)
|
||||
missing = [k for k in required if k not in actual]
|
||||
if missing:
|
||||
raise KeyError(
|
||||
f"Store at {getattr(store, '_load_path', '?')} is missing required "
|
||||
f"keys {missing}; available keys are {sorted(actual)}."
|
||||
)
|
||||
|
||||
|
||||
class BaseDataset(Dataset, ABC):
|
||||
"""Abstract base class for all dataset types.
|
||||
"""Abstract base class for dataset types.
|
||||
|
||||
Implements common functionality for window-based data fetching.
|
||||
Uses a storage abstraction for format-agnostic data loading.
|
||||
Holds a :class:`Store`. All sample-id indexing is delegated to the
|
||||
store — this class exposes ``__len__`` as ``len(store)`` and the
|
||||
``keys`` property as ``store.keys``. Subclasses implement
|
||||
``__getitem__`` with the train-type-specific key mapping and any
|
||||
training-only index arithmetic (e.g. the next-token ``+1`` shift).
|
||||
"""
|
||||
|
||||
def __init__(self, window_size: int, stride: int):
|
||||
required_keys: List[str] = []
|
||||
|
||||
def __init__(self, store: Store):
|
||||
super().__init__()
|
||||
self.window_size = window_size
|
||||
self.stride = stride
|
||||
self.storage: Optional[Store] = None
|
||||
self.store: Store = store
|
||||
validate_keys(store, self.required_keys)
|
||||
|
||||
@property
|
||||
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):
|
||||
"""Load dataset from the given path.
|
||||
|
||||
Auto-detects the storage format if not specified.
|
||||
|
||||
Args:
|
||||
load_path: Path to the data directory or file
|
||||
storage_type: Force a specific storage type ("h5", "bin"),
|
||||
or None for auto-detection
|
||||
|
||||
Raises:
|
||||
KeyError: If the loaded storage is missing required keys.
|
||||
"""
|
||||
if storage_type is None:
|
||||
storage_type = detect_format(load_path)
|
||||
self.storage = StoreFactory.create(storage_type)
|
||||
self._load_path = load_path
|
||||
self.storage.load(load_path)
|
||||
self._validate_keys()
|
||||
|
||||
@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)
|
||||
def __len__(self) -> int:
|
||||
return len(self.store)
|
||||
|
||||
@property
|
||||
def keys(self) -> List[str]:
|
||||
"""Return the available data keys."""
|
||||
if self.storage is None:
|
||||
return []
|
||||
return self.storage.keys
|
||||
return self.store.keys
|
||||
|
||||
def get_index(self, index: int) -> tuple:
|
||||
"""Calculate begin and end indices for a sample.
|
||||
|
||||
Args:
|
||||
index: Sample index
|
||||
|
||||
Returns:
|
||||
Tuple of (begin_idx, end_idx)
|
||||
"""
|
||||
if self.storage is None:
|
||||
raise RuntimeError("Dataset not loaded, call load() first")
|
||||
total = len(self.storage)
|
||||
if total <= self.window_size:
|
||||
raise ValueError(
|
||||
f"Data too short: {total} tokens <= window_size {self.window_size}"
|
||||
)
|
||||
|
||||
begin_idx = min(index * self.stride, total - 1 - self.window_size)
|
||||
end_idx = min(begin_idx + self.window_size, total - 1)
|
||||
|
||||
return begin_idx, end_idx
|
||||
@property
|
||||
def token_count(self) -> int:
|
||||
return self.store.token_count
|
||||
|
||||
@abstractmethod
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
"""Get a single sample by index.
|
||||
|
||||
Must be implemented by subclasses.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def __len__(self) -> int:
|
||||
if self.storage is None:
|
||||
return 0
|
||||
total = len(self.storage)
|
||||
if total <= self.window_size:
|
||||
return 0
|
||||
return (total - 1 - self.window_size) // self.stride + 1
|
||||
|
||||
|
||||
class DatasetFactory(BaseFactory["BaseDataset"]):
|
||||
"""Factory class for creating dataset instances.
|
||||
"""Factory for creating dataset instances by train-type.
|
||||
|
||||
Supports decorator-based registration for extensible dataset types.
|
||||
All default dataset types (seq, sft, dpo, grpo) are registered automatically
|
||||
when their classes are defined with the decorator.
|
||||
|
||||
Example usage:
|
||||
@DatasetFactory.register("custom")
|
||||
class CustomDataset(BaseDataset):
|
||||
...
|
||||
|
||||
dataset = DatasetFactory.create("custom", window_size, stride)
|
||||
Use :meth:`DatasetFactory.register("custom")` to register new
|
||||
dataset classes; they must inherit from :class:`BaseDataset`.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def _validate_component(cls, dataset_cls: type):
|
||||
"""Validate that the dataset class inherits from BaseDataset."""
|
||||
if not issubclass(dataset_cls, BaseDataset):
|
||||
raise TypeError(f"{dataset_cls.__name__} must inherit from BaseDataset")
|
||||
|
||||
@classmethod
|
||||
def create(cls, train_type: str, window_size: int, stride: int) -> "BaseDataset":
|
||||
"""Create a dataset instance.
|
||||
|
||||
Args:
|
||||
train_type: Type of training ("seq", "sft", "dpo", "grpo")
|
||||
window_size: Window size for data sampling
|
||||
stride: Stride between consecutive samples
|
||||
|
||||
Returns:
|
||||
Dataset instance
|
||||
"""
|
||||
return super().create(train_type, window_size, stride)
|
||||
|
||||
@classmethod
|
||||
def load(
|
||||
cls,
|
||||
train_type: str,
|
||||
load_path: str,
|
||||
window_size: int,
|
||||
load_path: Optional[str] = None,
|
||||
window_size: int = 0,
|
||||
stride: Optional[int] = None,
|
||||
storage_type: Optional[str] = None,
|
||||
tokenizer_path: Optional[str] = None,
|
||||
max_len: int = 2048,
|
||||
store: Optional[Store] = None,
|
||||
**kwargs,
|
||||
) -> "BaseDataset":
|
||||
"""Create and load a dataset in one step.
|
||||
|
||||
Two entry points:
|
||||
|
||||
- **store given**: bind it directly — the caller fully controls
|
||||
Store construction and processor setup. *load_path*,
|
||||
*storage_type*, *tokenizer_path*, *window_size*, *stride* are
|
||||
ignored.
|
||||
- **store is None**: build a Store from *load_path*, auto-detecting
|
||||
format and constructing a processor when *tokenizer_path* is
|
||||
given for a record dataset on JSONL.
|
||||
|
||||
Args:
|
||||
train_type: Type of training dataset
|
||||
load_path: Path to the data file
|
||||
window_size: Window size for data sampling
|
||||
stride: Stride between consecutive samples (default: same as window_size)
|
||||
storage_type: Storage type ("h5", "bin") or None for auto-detection
|
||||
train_type: Registered dataset name ("seq", "sft", "dpo",
|
||||
"grpo", …).
|
||||
load_path: Path to the data file or directory (ignored if
|
||||
*store* is given).
|
||||
window_size: Stream window length — only meaningful for
|
||||
stream datasets (SEQ/SFT). Record datasets ignore it.
|
||||
stride: Stride between consecutive stream samples
|
||||
(default: same as *window_size*).
|
||||
storage_type: Storage backend ("h5", "bin", "jsonl") or
|
||||
None for auto-detection.
|
||||
tokenizer_path: Path to tokenizer for lazy JSONL
|
||||
tokenisation (record datasets only).
|
||||
max_len: Max sequence length forwarded to processors.
|
||||
store: Pre-built, already-loaded Store instance.
|
||||
**kwargs: Extra arguments forwarded to ``store.load()``.
|
||||
|
||||
Returns:
|
||||
Loaded dataset instance
|
||||
Loaded dataset instance.
|
||||
"""
|
||||
if store is not None:
|
||||
return cls.create(train_type, store=store)
|
||||
|
||||
if load_path is None:
|
||||
raise ValueError("Either load_path or store must be provided")
|
||||
|
||||
if storage_type is None:
|
||||
storage_type = detect_format(load_path)
|
||||
|
||||
if stride is None:
|
||||
stride = window_size
|
||||
|
||||
dataset = cls.create(train_type, window_size, stride)
|
||||
dataset.load(load_path, storage_type=storage_type)
|
||||
processor = cls._maybe_build_processor(
|
||||
train_type, storage_type, tokenizer_path, max_len
|
||||
)
|
||||
|
||||
return dataset
|
||||
store_window = cls._store_window_for(train_type, window_size)
|
||||
store = StoreFactory.create(
|
||||
storage_type,
|
||||
window_size=store_window,
|
||||
stride=stride if stride else store_window,
|
||||
)
|
||||
if processor is not None:
|
||||
store.load(load_path, processor=processor, **kwargs)
|
||||
else:
|
||||
store.load(load_path, **kwargs)
|
||||
|
||||
@classmethod
|
||||
def available_types(cls) -> list:
|
||||
"""Return list of registered dataset type names."""
|
||||
return cls.list_registered()
|
||||
return cls.create(train_type, store=store)
|
||||
|
||||
@staticmethod
|
||||
def _store_window_for(train_type: str, window_size: int) -> int:
|
||||
"""Stream datasets consume ``window_size``; record datasets ignore it.
|
||||
|
||||
Record datasets (dpo/grpo) treat each record as an independent
|
||||
training unit and never window, so the store is built with
|
||||
``window_size=0`` and ``len(store)`` returns the record count.
|
||||
"""
|
||||
if train_type in ("seq", "sft"):
|
||||
return window_size
|
||||
return 0
|
||||
|
||||
@staticmethod
|
||||
def _maybe_build_processor(
|
||||
train_type: str,
|
||||
storage_type: str,
|
||||
tokenizer_path: Optional[str],
|
||||
max_len: int,
|
||||
) -> Optional[Callable[[dict], Dict[str, Tensor]]]:
|
||||
"""Build an on-the-fly tokenisation processor if applicable.
|
||||
|
||||
Only raw JSONL + record datasets (DPO/GRPO) need a processor;
|
||||
pre-tokenised backends (H5/bin) and stream datasets (SEQ/SFT)
|
||||
return ``None`` so no tokenizer is loaded.
|
||||
"""
|
||||
if tokenizer_path is None or storage_type != "jsonl":
|
||||
return None
|
||||
if train_type == "dpo":
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
||||
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
|
||||
return None
|
||||
|
||||
|
||||
@DatasetFactory.register("seq")
|
||||
class SEQDataset(BaseDataset):
|
||||
"""Dataset for sequential next-token prediction training."""
|
||||
"""Dataset for sequential next-token prediction training.
|
||||
|
||||
def __init__(self, window_size: int, stride: int):
|
||||
super().__init__(window_size, stride)
|
||||
Stream mode: ``store.fetch(begin, end, "sequence")`` returns the
|
||||
input window; the +1 shifted call returns the next-token target.
|
||||
"""
|
||||
|
||||
@property
|
||||
def required_keys(self) -> List[str]:
|
||||
return ["sequence"]
|
||||
required_keys = ["sequence"]
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int) -> Tensor:
|
||||
return self.storage.fetch(begin_idx, end_idx, "sequence")
|
||||
|
||||
def __getitem__(self, index):
|
||||
begin_idx, end_idx = self.get_index(index)
|
||||
|
||||
x = self._fetch_data(begin_idx, end_idx).to(dtype=torch.long)
|
||||
y = self._fetch_data(begin_idx + 1, end_idx + 1).to(dtype=torch.long)
|
||||
|
||||
return {"input_ids": x, "target_ids": y}
|
||||
def __getitem__(self, index: int):
|
||||
begin, end = self.store.sample_window(index)
|
||||
x = self.store.fetch(begin, end, "sequence")
|
||||
y = self.store.fetch(begin + 1, end + 1, "sequence")
|
||||
return {
|
||||
"input_ids": x.to(dtype=torch.long),
|
||||
"target_ids": y.to(dtype=torch.long),
|
||||
}
|
||||
|
||||
|
||||
@DatasetFactory.register("sft")
|
||||
class SFTDataset(BaseDataset):
|
||||
"""Dataset for supervised fine-tuning with loss masking."""
|
||||
"""Dataset for supervised fine-tuning with loss masking.
|
||||
|
||||
def __init__(self, window_size: int, stride: int):
|
||||
super().__init__(window_size, stride)
|
||||
Stream mode: ``sequence``/``loss_mask``/``position_ids`` are sliced
|
||||
to the window. ``loss_mask`` and ``target_ids`` use the +1 shifted
|
||||
slice so they align with the predicted positions.
|
||||
"""
|
||||
|
||||
@property
|
||||
def required_keys(self) -> List[str]:
|
||||
return ["sequence", "loss_mask"]
|
||||
required_keys = ["sequence", "loss_mask", "position_ids"]
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||
return self.storage.fetch(begin_idx, end_idx, key)
|
||||
|
||||
def __getitem__(self, index):
|
||||
begin_idx, end_idx = self.get_index(index)
|
||||
|
||||
x = self._fetch_data(begin_idx, end_idx, "sequence").to(dtype=torch.long)
|
||||
y = self._fetch_data(begin_idx + 1, end_idx + 1, "sequence").to(
|
||||
dtype=torch.long
|
||||
)
|
||||
loss_mask = self._fetch_data(begin_idx + 1, end_idx + 1, "loss_mask").to(
|
||||
dtype=torch.bool
|
||||
)
|
||||
|
||||
return {"input_ids": x, "target_ids": y, "loss_mask": loss_mask}
|
||||
def __getitem__(self, index: int):
|
||||
begin, end = self.store.sample_window(index)
|
||||
x = self.store.fetch(begin, end, "sequence")
|
||||
y = self.store.fetch(begin + 1, end + 1, "sequence")
|
||||
position_ids = self.store.fetch(begin, end, "position_ids")
|
||||
loss_mask = self.store.fetch(begin + 1, end + 1, "loss_mask")
|
||||
return {
|
||||
"input_ids": x.to(dtype=torch.long),
|
||||
"target_ids": y.to(dtype=torch.long),
|
||||
"position_ids": position_ids.to(dtype=torch.long),
|
||||
"loss_mask": loss_mask.to(dtype=torch.bool),
|
||||
}
|
||||
|
||||
|
||||
@DatasetFactory.register("dpo")
|
||||
class DPODataset(BaseDataset):
|
||||
"""Dataset for Direct Preference Optimization training."""
|
||||
"""Record-structured dataset for Direct Preference Optimization.
|
||||
|
||||
def __init__(self, window_size: int, stride: int):
|
||||
super().__init__(window_size, stride)
|
||||
Each sample is one preference pair (chosen + rejected) and is an
|
||||
independent training unit — no windowing, stride, or cross-record
|
||||
concatenation. This keeps each sequence self-contained so attention
|
||||
never leaks across preference pairs.
|
||||
|
||||
@property
|
||||
def required_keys(self) -> List[str]:
|
||||
return ["chosen", "rejected", "chosen_mask", "rejected_mask"]
|
||||
Two loading paths (handled by :class:`DatasetFactory`):
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||
return self.storage.fetch(begin_idx, end_idx, key)
|
||||
- **Pre-tokenized** (H5/bin): ``store.load(path)`` reads per-record
|
||||
tensors; ``__getitem__`` returns them directly.
|
||||
- **Raw JSONL** (``tokenizer_path=...``): builds a lazy processor
|
||||
via :func:`dpo_processor` that tokenises on the fly — no packing,
|
||||
no ``position_ids``.
|
||||
"""
|
||||
|
||||
def __getitem__(self, index: int):
|
||||
begin_idx, end_idx = self.get_index(index)
|
||||
required_keys = ["chosen", "rejected", "chosen_mask", "rejected_mask"]
|
||||
|
||||
chosen = self._fetch_data(begin_idx, end_idx, "chosen").to(dtype=torch.long)
|
||||
rejected = self._fetch_data(begin_idx, end_idx, "rejected").to(dtype=torch.long)
|
||||
chosen_mask = self._fetch_data(begin_idx, end_idx, "chosen_mask").to(
|
||||
dtype=torch.bool
|
||||
)
|
||||
rejected_mask = self._fetch_data(begin_idx, end_idx, "rejected_mask").to(
|
||||
dtype=torch.bool
|
||||
)
|
||||
def make_processor(self, tokenizer, max_len: int):
|
||||
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
|
||||
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
return {
|
||||
"chosen": chosen,
|
||||
"rejected": rejected,
|
||||
"chosen_mask": chosen_mask,
|
||||
"rejected_mask": rejected_mask,
|
||||
"chosen": self.store.fetch_record(index, "chosen").to(dtype=torch.long),
|
||||
"rejected": self.store.fetch_record(index, "rejected").to(dtype=torch.long),
|
||||
"chosen_mask": self.store.fetch_record(index, "chosen_mask").to(
|
||||
dtype=torch.bool
|
||||
),
|
||||
"rejected_mask": self.store.fetch_record(index, "rejected_mask").to(
|
||||
dtype=torch.bool
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
@DatasetFactory.register("grpo")
|
||||
class GRPODataset(BaseDataset):
|
||||
"""Dataset for Group Relative Policy Optimization training."""
|
||||
"""Dataset for offline Group Relative Policy Optimization.
|
||||
|
||||
def __init__(self, window_size: int, stride: int):
|
||||
super().__init__(window_size, stride)
|
||||
Each sample is one prompt with its group of responses and scalar
|
||||
rewards — an independent training unit with no windowing or stride.
|
||||
|
||||
@property
|
||||
def required_keys(self) -> List[str]:
|
||||
return ["prompts", "responses", "masks", "rewards"]
|
||||
Expected storage layout (produced by JsonlStore or pre-tokenized):
|
||||
|
||||
def _fetch_data(self, begin_idx: int, end_idx: int, key: str) -> Tensor:
|
||||
return self.storage.fetch(begin_idx, end_idx, key)
|
||||
- ``prompts``: List[Tensor] — one 1-D token tensor per record
|
||||
- ``responses``: List[List[Tensor]] — G response tensors per record
|
||||
- ``masks``: List[List[Tensor]] — G mask tensors per record
|
||||
- ``rewards``: List[Tensor] — one 1-D float tensor (len G) per record
|
||||
"""
|
||||
|
||||
required_keys = ["prompts", "responses", "masks", "rewards"]
|
||||
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
begin_idx, end_idx = self.get_index(index)
|
||||
|
||||
prompts = self._fetch_data(begin_idx, end_idx, "prompts").to(dtype=torch.long)
|
||||
responses = self._fetch_data(begin_idx, end_idx, "responses").to(
|
||||
dtype=torch.long
|
||||
)
|
||||
masks = self._fetch_data(begin_idx, end_idx, "masks").to(dtype=torch.bool)
|
||||
rewards = self._fetch_data(begin_idx, end_idx, "rewards")
|
||||
|
||||
prompts = self.store.fetch_record(index, "prompts")
|
||||
responses = self.store.fetch_record(index, "responses")
|
||||
masks = self.store.fetch_record(index, "masks")
|
||||
rewards = self.store.fetch_record(index, "rewards")
|
||||
return {
|
||||
"prompts": prompts,
|
||||
"responses": responses,
|
||||
"masks": masks,
|
||||
"rewards": rewards,
|
||||
"prompts": prompts.to(dtype=torch.long),
|
||||
"responses": [r.to(dtype=torch.long) for r in responses],
|
||||
"masks": [m.to(dtype=torch.bool) for m in masks],
|
||||
"rewards": rewards.to(dtype=torch.float32),
|
||||
}
|
||||
|
||||
@@ -5,7 +5,15 @@ import torch.distributed as dist
|
||||
from torch.utils.data import Dataset, Sampler
|
||||
|
||||
|
||||
class ResumableDistributedSampler(Sampler[int]):
|
||||
class RDSampler(Sampler[int]):
|
||||
"""Resumable Distributed Sampler.
|
||||
|
||||
A distributed sampler that supports checkpoint-based resume: iteration
|
||||
state (epoch, position) is tracked so training can continue from the
|
||||
exact sample after a restart. Shards the dataset across
|
||||
``dist.world_size`` replicas with optional shuffling.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
data_source: Dataset,
|
||||
@@ -74,6 +82,7 @@ class ResumableDistributedSampler(Sampler[int]):
|
||||
|
||||
self.epoch += 1
|
||||
self._indices = None
|
||||
self.iter = self.iter % self.num_samples_per_replica
|
||||
|
||||
@property
|
||||
def _remaining(self):
|
||||
|
||||
+546
-160
@@ -1,97 +1,70 @@
|
||||
"""Storage backends for different data formats.
|
||||
|
||||
Layers:
|
||||
- I/O layer: save_* / load_* functions, read/write raw files (HDF5/bin)
|
||||
return Dict[str, List[Tensor]] — format-specific, no state
|
||||
- Store (ABC): central abstraction, normalizes multi-segment into
|
||||
Dict[str, List[Tensor]] per key via _normalize(),
|
||||
fetch() uses bisect across segments — no forced concat
|
||||
- Dataset layer: BaseDataset owns a Store, only calls store.fetch(begin, end, key)
|
||||
Architecture (composition over inheritance):
|
||||
|
||||
Key properties:
|
||||
- Multi-segment: segments kept as-is, no forced concatenation — safe for
|
||||
datasets larger than RAM
|
||||
- Explicit length: _length = min(total elements across keys), set at load,
|
||||
__len__ returns O(1)
|
||||
- Zero-copy mmap: MmapStore wraps np.memmap(mode="r"), all DataLoader
|
||||
workers share OS page-cache pages
|
||||
Store (ABC) — owns _data/_cum/_offsets bookkeeping
|
||||
+ window_size/stride for sample-id
|
||||
indexing. __getitem__/__len__ produce
|
||||
the smallest iterable unit so Dataset
|
||||
classes are pure delegators.
|
||||
Streamable (mixin) — raw token slice fetch(begin, end, keys)
|
||||
Recordable (mixin) — raw record slice fetch_record(idx, keys)
|
||||
|
||||
H5Store(Store, Streamable, Recordable)
|
||||
MmapStore(Store, Streamable, Recordable)
|
||||
JsonlStore(Store, Streamable, Recordable)
|
||||
|
||||
Each mixin is a stateless trait that relies on ``self._data`` etc.
|
||||
provided by :class:`Store`. Concrete stores mix in whichever access
|
||||
primitives they support — ``Store`` is the sole base class, so there is
|
||||
no diamond inheritance or MRO ambiguity.
|
||||
|
||||
Sample-id indexing lives on :class:`Store`, not on the dataset:
|
||||
|
||||
- **Stream mode** (``window_size > 0``): ``len(store)`` returns the number
|
||||
of ``(window_size, stride)`` windows that fit in the token river;
|
||||
``store[i]`` returns the *i*-th window as a dict of per-key tensors;
|
||||
``store.sample_window(i)`` exposes the underlying ``(begin, end)``
|
||||
token slice for callers (e.g. next-token trainers) that need a +1
|
||||
shifted companion window.
|
||||
- **Record mode** (``num_records > 0``): ``len(store)`` returns the
|
||||
record count; ``store[i]`` returns the *i*-th record dict.
|
||||
|
||||
Raw token/record access via :meth:`fetch` / :meth:`fetch_record`
|
||||
remains available for low-level callers that want explicit index
|
||||
control. ``store.token_count`` is the total stream token count (what
|
||||
``len(store)`` used to mean in the legacy stream-only API).
|
||||
|
||||
``segments_are_records`` (class attribute on each Store subclass)
|
||||
tells ``_normalize`` whether segments are inherently per-record (H5/
|
||||
JSONL) or opaque shards (bin). Record access for bin relies on
|
||||
``_offsets`` instead.
|
||||
|
||||
:class:`JsonlStore` supports a lazy mode (``processor=fn``) that keeps
|
||||
raw records and defers tokenisation to ``fetch_record`` — used by DPO
|
||||
to train directly from a ``.jsonl`` file without a pre-tokenised copy.
|
||||
"""
|
||||
|
||||
import bisect
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Union
|
||||
from typing import Callable, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import h5py
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.preprocessing.transform import TokenizeTransform
|
||||
from astrai.serialization import (
|
||||
load_bin,
|
||||
load_bin_offsets,
|
||||
load_h5,
|
||||
)
|
||||
|
||||
|
||||
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_bin(file_path: str, tensor_group: Dict[str, List[Tensor]]):
|
||||
os.makedirs(file_path, exist_ok=True)
|
||||
meta = {}
|
||||
for key, tensors in tensor_group.items():
|
||||
cat = torch.cat(tensors, dim=0)
|
||||
meta[key] = {"shape": list(cat.shape), "dtype": str(cat.dtype).split(".")[-1]}
|
||||
np.asarray(cat.cpu().numpy()).tofile(os.path.join(file_path, f"{key}.bin"))
|
||||
with open(os.path.join(file_path, "meta.json"), "w") as f:
|
||||
json.dump(meta, f)
|
||||
|
||||
|
||||
def load_bin(file_path: str) -> Dict[str, List[Tensor]]:
|
||||
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
||||
meta = json.load(f)
|
||||
segments: Dict[str, List[Tensor]] = {}
|
||||
for key, info in meta.items():
|
||||
arr = np.memmap(
|
||||
os.path.join(file_path, f"{key}.bin"),
|
||||
dtype=info["dtype"],
|
||||
mode="r+",
|
||||
shape=tuple(info["shape"]),
|
||||
)
|
||||
segments[key] = [torch.from_numpy(arr)]
|
||||
return segments
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def detect_format(load_path: str) -> str:
|
||||
@@ -101,7 +74,7 @@ def detect_format(load_path: str) -> str:
|
||||
load_path: Directory or file path
|
||||
|
||||
Returns:
|
||||
Format string ("h5" or "bin")
|
||||
Format string ("h5", "bin", "jsonl", or "processed")
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If no supported data files are found
|
||||
@@ -111,140 +84,553 @@ def detect_format(load_path: str) -> str:
|
||||
suffix = root.suffix.lower()
|
||||
if suffix in (".h5", ".hdf5"):
|
||||
return "h5"
|
||||
if suffix == ".jsonl":
|
||||
return "jsonl"
|
||||
raise ValueError(f"Unsupported file format: {suffix}")
|
||||
|
||||
h5_files = list(root.rglob("*.h5")) + list(root.rglob("*.hdf5"))
|
||||
h5_files = [
|
||||
Path(p)
|
||||
for pattern in ("*.h5", "*.hdf5")
|
||||
for p in glob.glob(str(root / "**" / pattern), recursive=True)
|
||||
]
|
||||
if h5_files:
|
||||
return "h5"
|
||||
bin_files = list(root.rglob("*.bin"))
|
||||
if bin_files and (root / "meta.json").exists():
|
||||
return "bin"
|
||||
bin_files = [Path(p) for p in glob.glob(str(root / "**" / "*.bin"), recursive=True)]
|
||||
if bin_files:
|
||||
has_meta = (root / "meta.json").exists() or len(
|
||||
[Path(p) for p in glob.glob(str(root / "**" / "meta.json"), recursive=True)]
|
||||
) > 0
|
||||
if has_meta:
|
||||
return "bin"
|
||||
jsonl_files = [
|
||||
Path(p) for p in glob.glob(str(root / "**" / "*.jsonl"), recursive=True)
|
||||
]
|
||||
if jsonl_files:
|
||||
return "jsonl"
|
||||
raise FileNotFoundError(f"No supported data files found at {load_path}")
|
||||
|
||||
|
||||
class Store(ABC):
|
||||
"""String keys -> segmented tensors with ``fetch(begin, end, keys)``.
|
||||
"""Common base for all storage backends.
|
||||
|
||||
Each key maps to one or more tensor segments (no forced concatenation).
|
||||
``len(store)`` returns ``self._length`` (explicit, O(1)), the minimum
|
||||
total element count across all keys.
|
||||
A Store owns both its data layout AND its sample-id → token/record
|
||||
index translation. Datasets are thin wrappers that bind a Store
|
||||
to a particular train-type's key mapping; they never know about
|
||||
window/stride math.
|
||||
|
||||
Subclasses fill ``self._data`` and ``self._cum`` during ``load()``
|
||||
via ``_normalize()``.
|
||||
Two iteration modes:
|
||||
|
||||
- **Stream** (``window_size > 0``): data is treated as one long
|
||||
token river. ``len(store)`` returns the number of windows;
|
||||
``store[i]`` slices every stream-compatible key to window ``i``;
|
||||
``store.sample_window(i)`` returns the ``(begin, end)`` token
|
||||
slice for callers needing a +1 shifted companion window.
|
||||
- **Record** (``num_records > 0``): data is per-record.
|
||||
``len(store)`` returns ``num_records``; ``store[i]`` returns
|
||||
the *i*-th record as a dict.
|
||||
|
||||
Raw token slicing is still available via :meth:`fetch` (mixed in
|
||||
by :class:`Streamable`) when a store has stream support configured.
|
||||
Raw record slicing via :meth:`fetch_record` (mixed in by
|
||||
:class:`Recordable`) when a store has record support.
|
||||
|
||||
``token_count`` exposes the raw total stream length — this is what
|
||||
``len(store)`` returned in the legacy stream-only API and what
|
||||
stream-bound ``fetch`` uses for its bounds check.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
segments_are_records: bool = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window_size: int = 0,
|
||||
stride: Optional[int] = None,
|
||||
):
|
||||
self._data: Dict[str, List[Tensor]] = {}
|
||||
self._cum: Dict[str, List[int]] = {}
|
||||
self._offsets: Dict[str, List[int]] = {}
|
||||
self._length: int = 0
|
||||
self._num_records: int = 0
|
||||
self._window_size: int = int(window_size)
|
||||
self._stride: int = int(stride) if stride is not None else int(window_size)
|
||||
|
||||
@abstractmethod
|
||||
def load(self, path: str) -> None:
|
||||
def load(self, path: str, **kwargs) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
@property
|
||||
def keys(self) -> List[str]:
|
||||
return list(self._data.keys())
|
||||
|
||||
def __len__(self) -> int:
|
||||
@property
|
||||
def window_size(self) -> int:
|
||||
return self._window_size
|
||||
|
||||
@property
|
||||
def stride(self) -> int:
|
||||
return self._stride
|
||||
|
||||
@property
|
||||
def token_count(self) -> int:
|
||||
"""Total tokens across all stream segments.
|
||||
|
||||
Useful for the bounds-checked raw :meth:`fetch` and as the
|
||||
legacy ``len(store)`` value.
|
||||
"""
|
||||
return self._length
|
||||
|
||||
@property
|
||||
def num_records(self) -> int:
|
||||
"""Number of records available via :meth:`fetch_record`.
|
||||
|
||||
Non-zero only when the backing layout provides per-record
|
||||
indexing (H5/JSONL segments or bin ``_offsets``).
|
||||
"""
|
||||
return self._num_records
|
||||
|
||||
@property
|
||||
def num_samples(self) -> int:
|
||||
"""Number of items produced by ``__getitem__``.
|
||||
|
||||
Stream-mode wins when ``window_size > 0`` and there are tokens
|
||||
to slice; otherwise falls back to ``num_records``.
|
||||
"""
|
||||
if self._window_size > 0 and self._length > 0:
|
||||
total = self._length
|
||||
w = self._window_size
|
||||
if total <= w:
|
||||
return 0
|
||||
return (total - 1 - w) // self._stride + 1
|
||||
return self._num_records
|
||||
|
||||
def __len__(self) -> int:
|
||||
return self.num_samples
|
||||
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
if index < 0:
|
||||
index += self.num_samples
|
||||
if not 0 <= index < self.num_samples:
|
||||
raise IndexError(
|
||||
f"Store index out of range: {index}, num_samples={self.num_samples}"
|
||||
)
|
||||
if self._window_size > 0 and self._length > 0:
|
||||
begin, end = self.sample_window(index)
|
||||
keys = self._stream_keys()
|
||||
return {k: self.fetch(begin, end, k) for k in keys}
|
||||
return self.fetch_record(index, self._record_keys())
|
||||
|
||||
def sample_window(self, index: int) -> Tuple[int, int]:
|
||||
"""Return ``(begin, end)`` token positions for stream sample *index*.
|
||||
|
||||
The clipped tail keeps the last reachable window inside the
|
||||
token river instead of overshooting. Caller is responsible
|
||||
for staying within :attr:`num_samples`: an out-of-range index
|
||||
raises ``IndexError``.
|
||||
"""
|
||||
if self._window_size <= 0:
|
||||
raise RuntimeError("sample_window() requires window_size > 0 (stream mode)")
|
||||
if self._window_size <= 0 or self._length <= self._window_size:
|
||||
raise IndexError(
|
||||
f"Data too short for window: token_count={self._length}, "
|
||||
f"window_size={self._window_size}"
|
||||
)
|
||||
if not 0 <= index < self.num_samples:
|
||||
raise IndexError(
|
||||
f"Sample index out of range: {index}, num_samples={self.num_samples}"
|
||||
)
|
||||
total = self._length
|
||||
begin = min(index * self._stride, total - 1 - self._window_size)
|
||||
end = min(begin + self._window_size, total - 1)
|
||||
return begin, end
|
||||
|
||||
def _stream_keys(self) -> List[str]:
|
||||
out: List[str] = []
|
||||
for k, tensors in self._data.items():
|
||||
if tensors and isinstance(tensors[0], list):
|
||||
continue
|
||||
out.append(k)
|
||||
return out
|
||||
|
||||
def _record_keys(self) -> List[str]:
|
||||
return list(self._data.keys())
|
||||
|
||||
def _normalize(
|
||||
self,
|
||||
raw: Dict[str, list],
|
||||
offsets: Optional[Dict[str, List[int]]] = None,
|
||||
):
|
||||
"""Register segments and pre-compute indices for both access modes.
|
||||
|
||||
Stream mode: ``_cum[key]`` accumulates per-segment lengths so
|
||||
``Streamable._fetch_stream_key`` can bisect across segments
|
||||
without concatenation.
|
||||
|
||||
Record mode: if *offsets* is provided (bin layout),
|
||||
``_offsets[key]`` stores cumulative per-record offsets into the
|
||||
single concatenated segment. Otherwise, when
|
||||
``segments_are_records`` is True (H5/JSONL), ``_data[key]`` is
|
||||
a per-record list and ``fetch_record`` indexes it directly.
|
||||
|
||||
Nested keys (GRPO ``responses``/``masks`` as
|
||||
``List[List[Tensor]]``) are stored as-is and excluded from both
|
||||
cumulative bookkeepings — they are only accessed record-by-record.
|
||||
"""
|
||||
flat_lengths = []
|
||||
for key, tensors in raw.items():
|
||||
self._data[key] = tensors
|
||||
if not tensors:
|
||||
self._cum[key] = []
|
||||
flat_lengths.append(0)
|
||||
continue
|
||||
if isinstance(tensors[0], list):
|
||||
self._cum[key] = []
|
||||
continue
|
||||
cum = []
|
||||
total = 0
|
||||
for t in tensors:
|
||||
total += t.shape[0]
|
||||
cum.append(total)
|
||||
self._cum[key] = cum
|
||||
flat_lengths.append(cum[-1] if cum else 0)
|
||||
self._length = min(flat_lengths) if flat_lengths else 0
|
||||
|
||||
valid_offsets: Dict[str, List[int]] = {}
|
||||
if offsets:
|
||||
for key, off in offsets.items():
|
||||
segs = self._data.get(key, [])
|
||||
if len(segs) == 1 and len(off) > 1:
|
||||
valid_offsets[key] = off
|
||||
elif len(segs) > 1:
|
||||
logger.warning(
|
||||
"Key '%s' has %d segments with offsets — record mode "
|
||||
"disabled for this key (multi-shard bin+offsets not "
|
||||
"supported). Merge shards or use H5/JSONL.",
|
||||
key,
|
||||
len(segs),
|
||||
)
|
||||
self._offsets = valid_offsets
|
||||
if valid_offsets:
|
||||
record_counts = [len(v) - 1 for v in valid_offsets.values()]
|
||||
self._num_records = min(record_counts) if record_counts else 0
|
||||
elif self.segments_are_records:
|
||||
per_record_counts = []
|
||||
for key, tensors in self._data.items():
|
||||
if tensors and isinstance(tensors[0], list):
|
||||
continue
|
||||
per_record_counts.append(len(tensors))
|
||||
self._num_records = min(per_record_counts) if per_record_counts else 0
|
||||
else:
|
||||
self._num_records = 0
|
||||
|
||||
|
||||
class Streamable:
|
||||
"""Mixin granting raw token-stream access via :meth:`fetch`.
|
||||
|
||||
Stateless trait relying on ``self._data``, ``self._cum``,
|
||||
``self._length`` maintained by :class:`Store`. Stream mode is
|
||||
active when the owning store has ``window_size > 0``; for stores
|
||||
that can also serve record access (H5/JSONL/bin+offsets), the
|
||||
``fetch_record`` API from :class:`Recordable` is used instead.
|
||||
"""
|
||||
|
||||
def fetch(
|
||||
self,
|
||||
begin: int,
|
||||
end: int,
|
||||
keys: Union[str, List[str]],
|
||||
):
|
||||
if not self._data:
|
||||
raise RuntimeError("Store not loaded")
|
||||
if not (0 <= begin < self._length and 0 <= end <= self._length):
|
||||
raise ValueError(
|
||||
f"Index out of bounds: begin={begin}, end={end}, length={self._length}"
|
||||
)
|
||||
if isinstance(keys, str):
|
||||
return self._fetch_key(keys, begin, end)
|
||||
return {k: self._fetch_key(k, begin, end) for k in keys}
|
||||
return _stream_fetch(self, begin, end, keys)
|
||||
|
||||
def _fetch_key(self, key: str, begin: int, end: int) -> Tensor:
|
||||
"""Fetch slice [begin, end) across potentially multiple segments."""
|
||||
segments = self._data[key]
|
||||
cum = self._cum[key]
|
||||
seg_start = bisect.bisect_right(cum, begin)
|
||||
seg_end = bisect.bisect_left(cum, end)
|
||||
|
||||
results = []
|
||||
for i in range(seg_start, seg_end + 1):
|
||||
prev = cum[i - 1] if i > 0 else 0
|
||||
s = max(begin - prev, 0)
|
||||
e = min(end - prev, segments[i].shape[0])
|
||||
results.append(segments[i][s:e])
|
||||
|
||||
return results[0] if len(results) == 1 else torch.cat(results, dim=0)
|
||||
|
||||
def _normalize(self, raw: Dict[str, List[Tensor]]):
|
||||
"""Register segments and pre-compute cumulative lengths.
|
||||
|
||||
Does NOT concatenate — segments are kept as-is to avoid OOM on
|
||||
large datasets. Sets ``self._length`` to the minimum total
|
||||
element count across all keys.
|
||||
"""
|
||||
for key, tensors in raw.items():
|
||||
self._data[key] = tensors
|
||||
cum = []
|
||||
total = 0
|
||||
for t in tensors:
|
||||
total += t.shape[0]
|
||||
cum.append(total)
|
||||
self._cum[key] = cum
|
||||
self._length = (
|
||||
min((cum[-1] if cum else 0) for cum in self._cum.values())
|
||||
if self._cum
|
||||
else 0
|
||||
def _stream_fetch(self, begin: int, end: int, keys: Union[str, List[str]]):
|
||||
if not getattr(self, "_data", None):
|
||||
raise RuntimeError("Store not loaded")
|
||||
if not (0 <= begin < self._length and 0 <= end <= self._length):
|
||||
raise ValueError(
|
||||
f"Index out of bounds: begin={begin}, end={end}, length={self._length}"
|
||||
)
|
||||
if isinstance(keys, str):
|
||||
return _fetch_stream_key(self, keys, begin, end)
|
||||
return {k: _fetch_stream_key(self, k, begin, end) for k in keys}
|
||||
|
||||
|
||||
def _fetch_stream_key(self, key: str, begin: int, end: int) -> Tensor:
|
||||
segments = self._data[key]
|
||||
cum = self._cum[key]
|
||||
seg_start = bisect.bisect_right(cum, begin)
|
||||
seg_end = bisect.bisect_left(cum, end)
|
||||
|
||||
results = []
|
||||
for i in range(seg_start, seg_end + 1):
|
||||
prev = cum[i - 1] if i > 0 else 0
|
||||
s = max(begin - prev, 0)
|
||||
e = min(end - prev, segments[i].shape[0])
|
||||
results.append(segments[i][s:e])
|
||||
|
||||
return results[0] if len(results) == 1 else torch.cat(results, dim=0)
|
||||
|
||||
|
||||
class Recordable:
|
||||
"""Mixin granting raw record access via :meth:`fetch_record`.
|
||||
|
||||
Stateless trait relying on ``self._data``, ``self._offsets``,
|
||||
``self._num_records`` maintained by :class:`Store`.
|
||||
"""
|
||||
|
||||
def fetch_record(
|
||||
self,
|
||||
index: int,
|
||||
keys: Union[str, List[str]],
|
||||
):
|
||||
return _record_fetch(self, index, keys)
|
||||
|
||||
|
||||
def _record_fetch(self, index: int, keys: Union[str, List[str]]):
|
||||
if not getattr(self, "_data", None) and self._num_records == 0:
|
||||
raise RuntimeError("Store not loaded")
|
||||
if not 0 <= index < self._num_records:
|
||||
raise ValueError(
|
||||
f"Record index out of bounds: {index}, num_records={self._num_records}"
|
||||
)
|
||||
if isinstance(keys, str):
|
||||
return _fetch_record_key(self, keys, index)
|
||||
return {k: _fetch_record_key(self, k, index) for k in keys}
|
||||
|
||||
|
||||
def _fetch_record_key(self, key: str, index: int):
|
||||
offsets = self._offsets.get(key)
|
||||
if offsets:
|
||||
start = offsets[index]
|
||||
end = (
|
||||
offsets[index + 1]
|
||||
if index + 1 < len(offsets)
|
||||
else self._data[key][0].shape[0]
|
||||
)
|
||||
return self._data[key][0][start:end]
|
||||
return self._data[key][index]
|
||||
|
||||
|
||||
class StoreFactory(BaseFactory["Store"]):
|
||||
"""Factory for creating Store instances by type name.
|
||||
|
||||
Example::
|
||||
|
||||
@StoreFactory.register("custom")
|
||||
class CustomStore(Store):
|
||||
...
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def _validate_component(cls, store_cls: type):
|
||||
if not issubclass(store_cls, Store):
|
||||
raise TypeError(f"{store_cls.__name__} must inherit from Store")
|
||||
"""Factory for creating Store instances by type name."""
|
||||
|
||||
|
||||
@StoreFactory.register("h5")
|
||||
class H5Store(Store):
|
||||
"""HDF5-based storage backend (pre-tokenized data)."""
|
||||
class H5Store(Store, Streamable, Recordable):
|
||||
"""HDF5-based storage backend (pre-tokenized data).
|
||||
|
||||
def load(self, path: str):
|
||||
Each key is stored as a group of per-record datasets (``data_0``,
|
||||
``data_1``, …). Supports both access modes:
|
||||
|
||||
- **Stream**: ``fetch(begin, end, key)`` and ``store[i]`` slice
|
||||
across concatenated records via ``_cum`` — used by SEQ/SFT.
|
||||
- **Record**: ``fetch_record(i, key)`` and ``store[i]`` (when
|
||||
``window_size == 0``) index ``_data[key]`` directly — used by
|
||||
DPO/GRPO.
|
||||
"""
|
||||
|
||||
segments_are_records = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window_size: int = 0,
|
||||
stride: Optional[int] = None,
|
||||
):
|
||||
super().__init__(window_size=window_size, stride=stride)
|
||||
|
||||
def load(self, path: str, **kwargs):
|
||||
self._normalize(load_h5(path))
|
||||
|
||||
|
||||
@StoreFactory.register("bin")
|
||||
class MmapStore(Store):
|
||||
class MmapStore(Store, Streamable, Recordable):
|
||||
"""Memory-mapped binary storage backend.
|
||||
|
||||
Each key is a single .bin file backed by ``np.memmap(mode="r")``.
|
||||
No per-process memory duplication — all DataLoader workers share the
|
||||
same OS page-cache pages.
|
||||
|
||||
Format on disk::
|
||||
Supports both access modes:
|
||||
|
||||
data_root/
|
||||
meta.json # {key: {shape, dtype}, ...}
|
||||
<key>.bin # raw numpy array, one per key
|
||||
- **Stream**: always available via :meth:`fetch`.
|
||||
- **Record** (``fetch_record(i, key)``): only when ``meta.json``
|
||||
contains per-record ``offsets`` (written via
|
||||
``save_bin(..., record_keys=...)``). Legacy bin files without
|
||||
offsets have ``num_records == 0`` and ``len(store)`` reflects the
|
||||
windowed sample count when ``window_size > 0``.
|
||||
|
||||
``segments_are_records`` is ``False`` here (bin segments are
|
||||
contiguous streams, not per-record) — record access is driven
|
||||
purely by ``_offsets``.
|
||||
"""
|
||||
|
||||
def load(self, path: str):
|
||||
segments_are_records = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window_size: int = 0,
|
||||
stride: Optional[int] = None,
|
||||
):
|
||||
super().__init__(window_size=window_size, stride=stride)
|
||||
self._mmap_refs: List[Tensor] = []
|
||||
|
||||
def load(self, path: str, **kwargs):
|
||||
self._mmap_refs = []
|
||||
raw = load_bin(path)
|
||||
self._normalize(raw)
|
||||
root = Path(path)
|
||||
all_raw: Dict[str, List[Tensor]] = {}
|
||||
all_offsets: Dict[str, List[int]] = {}
|
||||
meta_paths = [
|
||||
Path(p) for p in glob.glob(str(root / "**" / "meta.json"), recursive=True)
|
||||
]
|
||||
for meta_path in meta_paths:
|
||||
raw = load_bin(str(meta_path.parent))
|
||||
off = load_bin_offsets(str(meta_path.parent))
|
||||
for key, tensors in raw.items():
|
||||
if key not in all_raw:
|
||||
all_raw[key] = []
|
||||
all_raw[key].extend(tensors)
|
||||
for key, o in off.items():
|
||||
if key not in all_offsets:
|
||||
all_offsets[key] = []
|
||||
all_offsets[key].extend(o)
|
||||
if not meta_paths:
|
||||
raise FileNotFoundError(f"No meta.json found under {path}")
|
||||
self._normalize(all_raw, offsets=all_offsets or None)
|
||||
for tensors in self._data.values():
|
||||
self._mmap_refs.extend(tensors)
|
||||
|
||||
|
||||
class JsonlSource:
|
||||
"""Read raw JSON records from a ``.jsonl`` file or directory.
|
||||
|
||||
A thin reader used by :class:`JsonlStore` in processor mode — holds
|
||||
no tokenizer, performs no tokenisation, just yields dicts.
|
||||
"""
|
||||
|
||||
def __init__(self, path: str):
|
||||
self.path = Path(path)
|
||||
self._records: Optional[List[dict]] = None
|
||||
|
||||
def load(self) -> List[dict]:
|
||||
if self._records is None:
|
||||
self._records = self._read(self.path)
|
||||
return self._records
|
||||
|
||||
@staticmethod
|
||||
def _read(root: Path) -> List[dict]:
|
||||
if root.is_file():
|
||||
return JsonlSource._read_file(root)
|
||||
return JsonlSource._read_dir(root)
|
||||
|
||||
@staticmethod
|
||||
def _read_file(path: Path) -> List[dict]:
|
||||
records: List[dict] = []
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
records.append(json.loads(line))
|
||||
except json.JSONDecodeError:
|
||||
logger.warning("Failed to parse JSON line in %s, skipping", path)
|
||||
return records
|
||||
|
||||
@staticmethod
|
||||
def _read_dir(root: Path) -> List[dict]:
|
||||
records: List[dict] = []
|
||||
for jsonl_path in sorted(root.glob("*.jsonl")):
|
||||
records.extend(JsonlSource._read_file(jsonl_path))
|
||||
return records
|
||||
|
||||
|
||||
@StoreFactory.register("jsonl")
|
||||
class JsonlStore(Store, Streamable, Recordable):
|
||||
"""JSONL reader with two tokenisation modes.
|
||||
|
||||
A JSONL dataset is a ``.jsonl`` file or a directory of ``*.jsonl``
|
||||
files plus (optionally) a ``dataset_config.json`` describing the
|
||||
tokenization pipeline.
|
||||
|
||||
Two modes, selected at :meth:`load` time:
|
||||
|
||||
- **Eager** (default): applies a :class:`TokenizeTransform` to every
|
||||
record at load time and registers per-key tensors via
|
||||
``_normalize``. Both ``fetch`` (stream) and ``fetch_record``
|
||||
(record) work.
|
||||
- **Lazy** (``processor=fn`` passed): keeps raw records and defers
|
||||
tokenisation to ``fetch_record``. Only record access works —
|
||||
``len(store)`` returns ``num_records``; stream primitives raise.
|
||||
"""
|
||||
|
||||
CONFIG_NAME = "dataset_config.json"
|
||||
segments_are_records = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
window_size: int = 0,
|
||||
stride: Optional[int] = None,
|
||||
):
|
||||
super().__init__(window_size=window_size, stride=stride)
|
||||
self._source: Optional[JsonlSource] = None
|
||||
self._processor: Optional[Callable[[dict], Dict[str, Tensor]]] = None
|
||||
self._keys_cache: Optional[List[str]] = None
|
||||
|
||||
def load(self, path: str, transform=None, processor=None, **kwargs):
|
||||
self._source = JsonlSource(path)
|
||||
records = self._source.load()
|
||||
|
||||
if processor is not None:
|
||||
self._processor = processor
|
||||
self._num_records = len(records)
|
||||
return
|
||||
|
||||
if transform is None:
|
||||
root = Path(path)
|
||||
config_path = root / self.CONFIG_NAME if root.is_dir() else None
|
||||
if config_path is None or not config_path.exists():
|
||||
raise FileNotFoundError(
|
||||
f"JSONL dataset config not found. Expected "
|
||||
f"{self.CONFIG_NAME} alongside *.jsonl files, pass an "
|
||||
f"explicit transform, or pass processor= for lazy "
|
||||
f"on-the-fly tokenisation."
|
||||
)
|
||||
transform = TokenizeTransform.from_config_file(str(config_path))
|
||||
|
||||
transformed = transform.apply(records)
|
||||
self._normalize(transformed)
|
||||
|
||||
@property
|
||||
def keys(self) -> List[str]:
|
||||
if self._processor is not None:
|
||||
if self._keys_cache is None and self._num_records > 0:
|
||||
sample = self._processor(self._source.load()[0])
|
||||
self._keys_cache = list(sample.keys())
|
||||
return self._keys_cache or []
|
||||
return list(self._data.keys())
|
||||
|
||||
def fetch_record(self, index: int, keys: Union[str, List[str]]):
|
||||
if self._processor is not None:
|
||||
if not 0 <= index < self._num_records:
|
||||
raise ValueError(
|
||||
f"Record index out of bounds: {index}, "
|
||||
f"num_records={self._num_records}"
|
||||
)
|
||||
record = self._source.load()[index]
|
||||
data = self._processor(record)
|
||||
if isinstance(keys, str):
|
||||
return data[keys]
|
||||
return {k: data[k] for k in keys}
|
||||
return _record_fetch(self, index, keys)
|
||||
|
||||
def fetch(self, begin: int, end: int, keys: Union[str, List[str]]):
|
||||
if self._processor is not None:
|
||||
raise RuntimeError(
|
||||
"JsonlStore in lazy (processor) mode does not support "
|
||||
"stream fetch(); use fetch_record() instead."
|
||||
)
|
||||
return _stream_fetch(self, begin, end, keys)
|
||||
|
||||
def __getitem__(self, index: int) -> Dict[str, Tensor]:
|
||||
if self._processor is not None:
|
||||
return self.fetch_record(index, self._record_keys())
|
||||
return super().__getitem__(index)
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
"""CUDA attention kernel wrappers with torch fallback.
|
||||
|
||||
Public API:
|
||||
- ``attn_decode`` — single-query decode attention
|
||||
- ``attn_prefill`` — multi-query prefill attention
|
||||
- ``attn_paged_decode`` — paged decode attention (direct page-table access)
|
||||
|
||||
Interface (shared by all wrappers):
|
||||
causal_offset: -1 = non-causal; >=0 = absolute position of first Q token
|
||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True = keep)
|
||||
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
|
||||
layout: "bhld" (default) or "blhd"
|
||||
|
||||
Causal and mask can coexist — both are applied simultaneously.
|
||||
|
||||
Each wrapper dispatches to its compiled CUDA kernel (``astrai.extension.attn_*``)
|
||||
when available, otherwise falls back to ``torch.nn.functional.scaled_dot_product_attention``.
|
||||
"""
|
||||
|
||||
from astrai.extension.loader import KERNEL_NAMES, is_available
|
||||
from astrai.extension.ops import attn_decode, attn_paged_decode, attn_prefill
|
||||
|
||||
__all__ = [
|
||||
"attn_decode",
|
||||
"attn_paged_decode",
|
||||
"attn_prefill",
|
||||
"is_available",
|
||||
"KERNEL_NAMES",
|
||||
]
|
||||
@@ -0,0 +1,36 @@
|
||||
"""Dynamic discovery and loading of compiled CUDA kernel modules.
|
||||
|
||||
Each kernel is registered in ``csrc/build.py`` and built into a ``.so`` placed
|
||||
in this package directory. On import we try to load each one; kernels that
|
||||
failed to build (or are running on a CPU-only machine) are marked unavailable
|
||||
so the wrapper functions can fall back to ``torch`` SDPA.
|
||||
"""
|
||||
|
||||
import importlib
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
KERNEL_NAMES = ["attn_decode", "attn_prefill", "attn_paged_decode"]
|
||||
|
||||
_available: dict[str, bool] = {}
|
||||
_modules: dict[str, object] = {}
|
||||
|
||||
for _name in KERNEL_NAMES:
|
||||
try:
|
||||
_mod = importlib.import_module(f".{_name}", package=__package__)
|
||||
_available[_name] = True
|
||||
_modules[_name] = _mod
|
||||
except ImportError:
|
||||
_available[_name] = False
|
||||
_modules[_name] = None
|
||||
|
||||
|
||||
def is_available(name: str) -> bool:
|
||||
"""Return ``True`` if the compiled kernel ``name`` was loaded."""
|
||||
return _available.get(name, False)
|
||||
|
||||
|
||||
def get_module(name: str) -> object:
|
||||
"""Return the loaded kernel module for ``name``, or ``None`` if unavailable."""
|
||||
return _modules.get(name)
|
||||
@@ -0,0 +1,246 @@
|
||||
"""GQA attention wrapper functions — one entry point per compiled kernel.
|
||||
|
||||
Each wrapper dispatches to its CUDA kernel (loaded in ``loader.py``) when
|
||||
available, otherwise falls back to ``torch`` SDPA.
|
||||
|
||||
Interface (all functions):
|
||||
causal_offset: -1 = non-causal; >=0 = absolute position of first Q token
|
||||
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool)
|
||||
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
|
||||
layout: "bhld" (default) or "blhd"
|
||||
|
||||
Add new kernel wrappers here; split into per-variant files only if this file
|
||||
grows large.
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from astrai.extension.loader import _available, _modules
|
||||
|
||||
_LAYOUT_CODES: dict[str, int] = {"bhld": 0, "blhd": 1}
|
||||
|
||||
|
||||
def _parse_layout(layout: str | int) -> int:
|
||||
if isinstance(layout, int):
|
||||
return layout
|
||||
code = _LAYOUT_CODES.get(layout.lower())
|
||||
if code is None:
|
||||
raise ValueError(
|
||||
f"unknown layout '{layout}', expected one of {list(_LAYOUT_CODES)}"
|
||||
)
|
||||
return code
|
||||
|
||||
|
||||
def _to_bhld(t: torch.Tensor, layout: int) -> torch.Tensor:
|
||||
"""Normalize to b h l d view. Zero-copy transpose if layout==1 (b l h d)."""
|
||||
if layout == 1:
|
||||
return t.transpose(1, 2)
|
||||
return t
|
||||
|
||||
|
||||
def _expand_kv_heads(
|
||||
k: torch.Tensor, v: torch.Tensor, q_head: int
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Expand K/V heads to match Q heads for GQA fallback."""
|
||||
kv_head = k.size(1)
|
||||
if kv_head == q_head:
|
||||
return k, v
|
||||
group = q_head // kv_head
|
||||
k = k.repeat_interleave(group, dim=1)
|
||||
v = v.repeat_interleave(group, dim=1)
|
||||
return k, v
|
||||
|
||||
|
||||
def _build_attn_mask(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
mask: torch.Tensor | None,
|
||||
causal_offset: int,
|
||||
scale: float,
|
||||
) -> tuple[torch.Tensor | None, float]:
|
||||
"""Build SDPA-compatible attn_mask + resolved scale.
|
||||
|
||||
q and k must already be in b h l d layout.
|
||||
Causal and mask can coexist: causal sets -inf above the diagonal, mask
|
||||
sets -inf for padded positions. Both are OR'd into a single bool mask.
|
||||
"""
|
||||
q_len = q.size(2)
|
||||
kv_len = k.size(2)
|
||||
head_dim = q.size(3)
|
||||
resolved_scale = scale if scale and scale > 0 else 1.0 / math.sqrt(head_dim)
|
||||
|
||||
attn_mask = None
|
||||
|
||||
if mask is not None:
|
||||
if mask.dim() == 2:
|
||||
# [batch, kv_len] → [batch, 1, 1, kv_len]
|
||||
attn_mask = mask[:, None, None, :]
|
||||
elif mask.dim() == 3:
|
||||
# [batch, q_len, kv_len] → [batch, 1, q_len, kv_len]
|
||||
attn_mask = mask[:, None, :, :]
|
||||
else:
|
||||
raise ValueError(f"mask must be 2D or 3D, got {mask.dim()}D")
|
||||
|
||||
if causal_offset >= 0:
|
||||
batch = q.size(0)
|
||||
# q row i attends to kv cols 0..(causal_offset + i)
|
||||
q_idx = torch.arange(q_len, device=q.device).unsqueeze(1) # [q_len, 1]
|
||||
kv_idx = torch.arange(kv_len, device=q.device).unsqueeze(0) # [1, kv_len]
|
||||
causal_bool = kv_idx > (causal_offset + q_idx) # True = masked out
|
||||
causal_mask = causal_bool.unsqueeze(0).expand(
|
||||
batch, -1, -1
|
||||
) # [batch, q_len, kv_len]
|
||||
causal_mask = causal_mask[:, None, :, :] # [batch, 1, q_len, kv_len]
|
||||
|
||||
if attn_mask is not None:
|
||||
attn_mask = attn_mask | causal_mask
|
||||
else:
|
||||
attn_mask = causal_mask
|
||||
|
||||
return attn_mask, resolved_scale
|
||||
|
||||
|
||||
def _torch_fallback(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: torch.Tensor | None,
|
||||
causal_offset: int,
|
||||
scale: float,
|
||||
q_layout: int,
|
||||
kv_layout: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Reference attention via ``scaled_dot_product_attention``.
|
||||
|
||||
q_layout / kv_layout: 0 = b h l d, 1 = b l h d.
|
||||
If kv_layout is None, uses q_layout (Q and K/V share the same layout).
|
||||
"""
|
||||
if kv_layout is None:
|
||||
kv_layout = q_layout
|
||||
q = _to_bhld(q, q_layout)
|
||||
k = _to_bhld(k, kv_layout)
|
||||
v = _to_bhld(v, kv_layout)
|
||||
k, v = _expand_kv_heads(k, v, q.size(1))
|
||||
attn_mask, resolved_scale = _build_attn_mask(q, k, mask, causal_offset, scale)
|
||||
out = F.scaled_dot_product_attention(
|
||||
q, k, v, attn_mask=attn_mask, is_causal=False, scale=resolved_scale
|
||||
)
|
||||
# Restore Q's original layout
|
||||
if q_layout == 1:
|
||||
out = out.transpose(1, 2)
|
||||
return out
|
||||
|
||||
|
||||
def _gather_kv_from_pages(
|
||||
page_table: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
v_cache: torch.Tensor,
|
||||
page_size: int,
|
||||
kv_len: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Gather contiguous K/V from paged cache for torch SDPA fallback.
|
||||
|
||||
Shapes:
|
||||
page_table : [batch, max_pages] (int64)
|
||||
k_cache : [n_pages, page_size, n_kv_heads, head_dim]
|
||||
v_cache : same as k_cache
|
||||
Returns:
|
||||
k, v : [batch, kv_len, n_kv_heads, head_dim] (b l h d)
|
||||
"""
|
||||
batch, max_pages = page_table.shape
|
||||
_, ps, n_kv_heads, head_dim = k_cache.shape
|
||||
if ps != page_size:
|
||||
raise ValueError(f"k_cache page_size mismatch: {ps} vs {page_size}")
|
||||
|
||||
# Vectorized gather: build physical page + offset indices, then advanced-index
|
||||
positions = torch.arange(kv_len, device=page_table.device)
|
||||
logical_pages = positions // page_size # [kv_len]
|
||||
page_offsets = positions % page_size # [kv_len]
|
||||
|
||||
phys_pages = page_table[:, logical_pages] # [batch, kv_len]
|
||||
# k_cache[phys_pages, page_offsets] → [batch, kv_len, n_kv_heads, head_dim] (b l h d)
|
||||
k = k_cache[phys_pages, page_offsets]
|
||||
v = v_cache[phys_pages, page_offsets]
|
||||
return k, v
|
||||
|
||||
|
||||
def attn_decode(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: torch.Tensor | None = None,
|
||||
causal_offset: int = -1,
|
||||
scale: float = 0.0,
|
||||
layout: str = "bhld",
|
||||
) -> torch.Tensor:
|
||||
li = _parse_layout(layout)
|
||||
if _available["attn_decode"]:
|
||||
return _modules["attn_decode"].attn_decode(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
mask=mask,
|
||||
causal_offset=causal_offset,
|
||||
scale=scale,
|
||||
layout=li,
|
||||
)
|
||||
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
|
||||
|
||||
|
||||
def attn_prefill(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
mask: torch.Tensor | None = None,
|
||||
causal_offset: int = -1,
|
||||
scale: float = 0.0,
|
||||
layout: str = "bhld",
|
||||
) -> torch.Tensor:
|
||||
li = _parse_layout(layout)
|
||||
if _available["attn_prefill"]:
|
||||
return _modules["attn_prefill"].attn_prefill(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
mask=mask,
|
||||
causal_offset=causal_offset,
|
||||
scale=scale,
|
||||
layout=li,
|
||||
)
|
||||
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
|
||||
|
||||
|
||||
def attn_paged_decode(
|
||||
q: torch.Tensor,
|
||||
page_table: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
v_cache: torch.Tensor,
|
||||
page_size: int,
|
||||
kv_len: int,
|
||||
mask: torch.Tensor | None = None,
|
||||
causal_offset: int = -1,
|
||||
scale: float = 0.0,
|
||||
layout: str = "bhld",
|
||||
) -> torch.Tensor:
|
||||
li = _parse_layout(layout)
|
||||
if _available["attn_paged_decode"]:
|
||||
return _modules["attn_paged_decode"].attn_paged_decode(
|
||||
q,
|
||||
page_table,
|
||||
k_cache,
|
||||
v_cache,
|
||||
page_size,
|
||||
kv_len,
|
||||
mask=mask,
|
||||
causal_offset=causal_offset,
|
||||
scale=scale,
|
||||
layout=li,
|
||||
)
|
||||
# Gathered K/V are always b l h d
|
||||
k, v = _gather_kv_from_pages(page_table, k_cache, v_cache, page_size, kv_len)
|
||||
return _torch_fallback(
|
||||
q, k, v, mask, causal_offset, scale, q_layout=li, kv_layout=1
|
||||
)
|
||||
+76
-158
@@ -1,149 +1,103 @@
|
||||
"""Base factory class for extensible component registration."""
|
||||
"""Base factory with decorator-based registration and kwarg-filtered instantiation."""
|
||||
|
||||
import inspect
|
||||
import sys
|
||||
from abc import ABC
|
||||
from typing import Callable, Dict, Generic, List, Optional, Tuple, Type, TypeVar
|
||||
from typing import (
|
||||
Callable,
|
||||
Dict,
|
||||
ForwardRef,
|
||||
Generic,
|
||||
List,
|
||||
Optional,
|
||||
Type,
|
||||
TypeVar,
|
||||
Union,
|
||||
)
|
||||
from typing import get_args as _get_args
|
||||
from typing import get_origin as _get_origin
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
class Registry:
|
||||
"""Flexible registry for component classes with category and priority support.
|
||||
def _resolve_type(
|
||||
arg: Union[Type, str, ForwardRef], factory_cls: type
|
||||
) -> Optional[Type]:
|
||||
"""Resolve a generic type-arg (str forward-ref, ForwardRef, or class)."""
|
||||
if not isinstance(arg, (str, ForwardRef)):
|
||||
return arg
|
||||
|
||||
This registry stores component classes with optional metadata (category, priority).
|
||||
It provides methods for registration, retrieval, and listing with filtering.
|
||||
"""
|
||||
name = arg if isinstance(arg, str) else arg.__forward_arg__
|
||||
if name == factory_cls.__name__:
|
||||
return factory_cls
|
||||
|
||||
def __init__(self):
|
||||
self._entries = {} # name -> (component_cls, category, priority)
|
||||
mod = sys.modules.get(factory_cls.__module__)
|
||||
if mod is None:
|
||||
return None
|
||||
ns = vars(mod)
|
||||
|
||||
def register(
|
||||
self,
|
||||
name: str,
|
||||
component_cls: Type,
|
||||
category: Optional[str] = None,
|
||||
priority: int = 0,
|
||||
):
|
||||
"""Register a component class with optional category and priority."""
|
||||
if name in self._entries:
|
||||
raise ValueError(f"Component '{name}' is already registered")
|
||||
self._entries[name] = (component_cls, category, priority)
|
||||
if isinstance(arg, ForwardRef):
|
||||
return arg._evaluate(ns, None, recursive_guard=frozenset())
|
||||
|
||||
def get(self, name: str) -> Type:
|
||||
"""Get component class by name."""
|
||||
if name not in self._entries:
|
||||
raise KeyError(f"Component '{name}' not found in registry")
|
||||
return self._entries[name][0]
|
||||
|
||||
def get_with_metadata(self, name: str) -> Tuple[Type, Optional[str], int]:
|
||||
"""Get component class with its metadata."""
|
||||
entry = self._entries.get(name)
|
||||
if entry is None:
|
||||
raise KeyError(f"Component '{name}' not found in registry")
|
||||
return entry
|
||||
|
||||
def contains(self, name: str) -> bool:
|
||||
"""Check if a name is registered."""
|
||||
return name in self._entries
|
||||
|
||||
def list_names(self) -> List[str]:
|
||||
"""Return list of registered component names."""
|
||||
return sorted(self._entries.keys())
|
||||
|
||||
def list_by_category(self, category: str) -> List[str]:
|
||||
"""Return names of components belonging to a specific category."""
|
||||
return sorted(
|
||||
name for name, (_, cat, _) in self._entries.items() if cat == category
|
||||
)
|
||||
|
||||
def list_by_priority(self, reverse: bool = False) -> List[str]:
|
||||
"""Return names sorted by priority (default ascending)."""
|
||||
return sorted(
|
||||
self._entries.keys(),
|
||||
key=lambda name: self._entries[name][2],
|
||||
reverse=reverse,
|
||||
)
|
||||
|
||||
def entries(self) -> Dict[str, Tuple[Type, Optional[str], int]]:
|
||||
"""Return raw entries dictionary."""
|
||||
return self._entries.copy()
|
||||
return ns.get(name)
|
||||
|
||||
|
||||
class BaseFactory(ABC, Generic[T]):
|
||||
"""Generic factory class for component registration and creation.
|
||||
"""Generic factory with decorator-based component registration.
|
||||
|
||||
This base class provides a decorator-based registration pattern
|
||||
for creating extensible component factories.
|
||||
|
||||
Example usage:
|
||||
class MyFactory(BaseFactory[MyBaseClass]):
|
||||
class MyFactory(BaseFactory[MyBase]):
|
||||
pass
|
||||
|
||||
@MyFactory.register("custom")
|
||||
class CustomComponent(MyBaseClass):
|
||||
class CustomComponent(MyBase):
|
||||
...
|
||||
|
||||
component = MyFactory.create("custom", *args, **kwargs)
|
||||
obj = MyFactory.create("custom", *args, **kwargs)
|
||||
|
||||
``create()`` filters kwargs to match the component's ``__init__``
|
||||
signature so components don't need ``**kwargs`` just to absorb
|
||||
unrelated parameters.
|
||||
"""
|
||||
|
||||
_registry: Registry
|
||||
_entries: Dict[str, Type[T]]
|
||||
|
||||
def __init_subclass__(cls, **kwargs):
|
||||
super().__init_subclass__(**kwargs)
|
||||
cls._registry = Registry()
|
||||
for orig_base in getattr(cls, "__orig_bases__", ()):
|
||||
if _get_origin(orig_base) is BaseFactory:
|
||||
(arg,) = _get_args(orig_base)
|
||||
cls._entries = {}
|
||||
cls._component_base = _resolve_type(arg, cls)
|
||||
return
|
||||
|
||||
@classmethod
|
||||
def register(
|
||||
cls, name: str, category: Optional[str] = None, priority: int = 0
|
||||
) -> Callable[[Type[T]], Type[T]]:
|
||||
"""Decorator to register a component class with optional category and priority.
|
||||
def register(cls, name: str) -> Callable[[Type[T]], Type[T]]:
|
||||
"""Decorator to register a component class.
|
||||
|
||||
Args:
|
||||
name: Registration name for the component
|
||||
category: Optional category for grouping components
|
||||
priority: Priority for ordering (default 0)
|
||||
|
||||
Returns:
|
||||
Decorator function that registers the component class
|
||||
|
||||
Raises:
|
||||
TypeError: If the decorated class doesn't inherit from the base type
|
||||
Validates that the decorated class inherits from the generic
|
||||
type parameter ``T`` declared on the factory.
|
||||
"""
|
||||
|
||||
def decorator(component_cls: Type[T]) -> Type[T]:
|
||||
cls._validate_component(component_cls)
|
||||
cls._registry.register(
|
||||
name, component_cls, category=category, priority=priority
|
||||
)
|
||||
if name in cls._entries:
|
||||
raise ValueError(f"Component '{name}' is already registered")
|
||||
cls._entries[name] = component_cls
|
||||
return component_cls
|
||||
|
||||
return decorator
|
||||
|
||||
@classmethod
|
||||
def create(cls, name: str, *args, **kwargs) -> T:
|
||||
"""Create a component instance by name.
|
||||
|
||||
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:
|
||||
name: Registered name of the component
|
||||
*args: Positional arguments passed to component constructor
|
||||
**kwargs: Keyword arguments passed to component constructor
|
||||
|
||||
Returns:
|
||||
Component instance
|
||||
|
||||
Raises:
|
||||
ValueError: If the component name is not registered
|
||||
"""Create a component instance by name, filtering kwargs to match
|
||||
the component's ``__init__`` signature.
|
||||
"""
|
||||
if not cls._registry.contains(name):
|
||||
entry = cls._entries.get(name)
|
||||
if entry is None:
|
||||
raise ValueError(
|
||||
f"Unknown component: '{name}'. "
|
||||
f"Supported types: {sorted(cls._registry.list_names())}"
|
||||
f"Unknown component: '{name}'. Supported types: {sorted(cls._entries)}"
|
||||
)
|
||||
component_cls = cls._registry.get(name)
|
||||
component_cls = entry
|
||||
sig = inspect.signature(component_cls.__init__)
|
||||
has_var_kwargs = any(
|
||||
p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()
|
||||
@@ -159,68 +113,32 @@ class BaseFactory(ABC, Generic[T]):
|
||||
|
||||
@classmethod
|
||||
def _validate_component(cls, component_cls: Type[T]):
|
||||
"""Validate that the component class is valid for this factory.
|
||||
"""Validate the decorated class inherits from the factory's base type.
|
||||
|
||||
Override this method in subclasses to add custom validation.
|
||||
|
||||
Args:
|
||||
component_cls: Component class to validate
|
||||
|
||||
Raises:
|
||||
TypeError: If the component class is invalid
|
||||
Override for custom validation beyond ``issubclass``.
|
||||
"""
|
||||
pass
|
||||
base = cls._component_base
|
||||
if base is not None and not issubclass(component_cls, base):
|
||||
raise TypeError(
|
||||
f"{component_cls.__name__} must inherit from {base.__name__}"
|
||||
)
|
||||
|
||||
@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):
|
||||
"""Get the registered component class without instantiating it."""
|
||||
entry = cls._entries.get(name)
|
||||
if entry is None:
|
||||
raise ValueError(
|
||||
f"Unknown component: '{name}'. "
|
||||
f"Supported types: {sorted(cls._registry.list_names())}"
|
||||
f"Unknown component: '{name}'. Supported types: {sorted(cls._entries)}"
|
||||
)
|
||||
return cls._registry.get(name)
|
||||
return entry
|
||||
|
||||
@classmethod
|
||||
def list_registered(cls) -> list:
|
||||
"""List all registered component names.
|
||||
|
||||
Returns:
|
||||
List of registered component names
|
||||
"""
|
||||
return cls._registry.list_names()
|
||||
def list_registered(cls) -> List[str]:
|
||||
"""List all registered component names."""
|
||||
return sorted(cls._entries)
|
||||
|
||||
@classmethod
|
||||
def is_registered(cls, name: str) -> bool:
|
||||
"""Check if a component name is registered.
|
||||
|
||||
Args:
|
||||
name: Component name to check
|
||||
|
||||
Returns:
|
||||
True if registered, False otherwise
|
||||
"""
|
||||
return cls._registry.contains(name)
|
||||
|
||||
@classmethod
|
||||
def list_by_category(cls, category: str) -> List[str]:
|
||||
"""List registered component names in a category."""
|
||||
return cls._registry.list_by_category(category)
|
||||
|
||||
@classmethod
|
||||
def list_by_priority(cls, reverse: bool = False) -> List[str]:
|
||||
"""List registered component names sorted by priority."""
|
||||
return cls._registry.list_by_priority(reverse)
|
||||
|
||||
|
||||
__all__ = ["Registry", "BaseFactory"]
|
||||
"""Check if a component name is registered."""
|
||||
return name in cls._entries
|
||||
|
||||
@@ -6,18 +6,23 @@ Layers:
|
||||
- protocols/: Response builders (OpenAI, Anthropic)
|
||||
- transport/: SSE transport utilities
|
||||
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
|
||||
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy)
|
||||
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy, FrequencyPenaltyStrategy)
|
||||
"""
|
||||
|
||||
from astrai.inference.api import (
|
||||
AnthropicMessage,
|
||||
BaseToolParser,
|
||||
ChatCompletionRequest,
|
||||
ChatMessage,
|
||||
FunctionDef,
|
||||
GenContext,
|
||||
MessagesRequest,
|
||||
ProtocolHandler,
|
||||
SimpleJsonToolParser,
|
||||
StopChecker,
|
||||
app,
|
||||
ToolDef,
|
||||
ToolParserFactory,
|
||||
get_app,
|
||||
run_server,
|
||||
)
|
||||
from astrai.inference.api.anthropic import AnthropicResponseBuilder
|
||||
@@ -25,10 +30,14 @@ from astrai.inference.api.openai import OpenAIResponseBuilder
|
||||
from astrai.inference.core import (
|
||||
STOP,
|
||||
Allocator,
|
||||
CacheView,
|
||||
ContiguousCache,
|
||||
ContiguousCacheView,
|
||||
Executor,
|
||||
InferenceScheduler,
|
||||
KVCache,
|
||||
KvcacheView,
|
||||
PageCache,
|
||||
PageCacheView,
|
||||
PagePool,
|
||||
PrefixCache,
|
||||
Storage,
|
||||
@@ -41,6 +50,7 @@ from astrai.inference.core import (
|
||||
from astrai.inference.engine import GenerationRequest, InferenceEngine
|
||||
from astrai.inference.sample import (
|
||||
BaseSamplingStrategy,
|
||||
FrequencyPenaltyStrategy,
|
||||
SamplingPipeline,
|
||||
TemperatureStrategy,
|
||||
TopKStrategy,
|
||||
@@ -58,8 +68,12 @@ __all__ = [
|
||||
"TaskManager",
|
||||
"TaskStatus",
|
||||
"Allocator",
|
||||
"CacheView",
|
||||
"KVCache",
|
||||
"KvcacheView",
|
||||
"ContiguousCache",
|
||||
"ContiguousCacheView",
|
||||
"PageCache",
|
||||
"PageCacheView",
|
||||
"PagePool",
|
||||
"PrefixCache",
|
||||
"Storage",
|
||||
@@ -70,16 +84,22 @@ __all__ = [
|
||||
"TemperatureStrategy",
|
||||
"TopKStrategy",
|
||||
"TopPStrategy",
|
||||
"FrequencyPenaltyStrategy",
|
||||
"SamplingPipeline",
|
||||
"ProtocolHandler",
|
||||
"StopChecker",
|
||||
"GenContext",
|
||||
"BaseToolParser",
|
||||
"SimpleJsonToolParser",
|
||||
"ToolParserFactory",
|
||||
"OpenAIResponseBuilder",
|
||||
"AnthropicResponseBuilder",
|
||||
"ChatMessage",
|
||||
"ChatCompletionRequest",
|
||||
"FunctionDef",
|
||||
"ToolDef",
|
||||
"AnthropicMessage",
|
||||
"MessagesRequest",
|
||||
"app",
|
||||
"get_app",
|
||||
"run_server",
|
||||
]
|
||||
|
||||
@@ -1,23 +1,39 @@
|
||||
"""Inference API: protocol handler, stop checker, and FastAPI server."""
|
||||
"""Inference API: protocol handler, stop checker, tool parsers, and FastAPI server.
|
||||
|
||||
``app`` is no longer a module-level global. Use :func:`get_app` to access the
|
||||
lazy singleton FastAPI instance.
|
||||
"""
|
||||
|
||||
from astrai.inference.api.protocol import GenContext, ProtocolHandler, StopChecker
|
||||
from astrai.inference.api.server import (
|
||||
AnthropicMessage,
|
||||
ChatCompletionRequest,
|
||||
ChatMessage,
|
||||
FunctionDef,
|
||||
MessagesRequest,
|
||||
app,
|
||||
ToolDef,
|
||||
get_app,
|
||||
run_server,
|
||||
)
|
||||
from astrai.inference.api.tool_parser import (
|
||||
BaseToolParser,
|
||||
SimpleJsonToolParser,
|
||||
ToolParserFactory,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ProtocolHandler",
|
||||
"StopChecker",
|
||||
"GenContext",
|
||||
"BaseToolParser",
|
||||
"SimpleJsonToolParser",
|
||||
"ToolParserFactory",
|
||||
"AnthropicMessage",
|
||||
"ChatCompletionRequest",
|
||||
"ChatMessage",
|
||||
"FunctionDef",
|
||||
"ToolDef",
|
||||
"MessagesRequest",
|
||||
"app",
|
||||
"get_app",
|
||||
"run_server",
|
||||
]
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
"""Anthropic message completion response builder."""
|
||||
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Tuple, Union
|
||||
|
||||
@@ -39,9 +40,8 @@ class AnthropicResponseBuilder(ResponseBuilder):
|
||||
prompt = engine.tokenizer.apply_chat_template(messages, tokenize=False)
|
||||
ctx = GenContext(
|
||||
resp_id=f"msg_{uuid.uuid4().hex[:24]}",
|
||||
created=0,
|
||||
created=int(time.time()),
|
||||
model=request.model,
|
||||
prompt_tokens=0,
|
||||
)
|
||||
stop_sequences = getattr(request, "stop_sequences", None) or []
|
||||
return prompt, ctx, stop_sequences
|
||||
@@ -72,15 +72,17 @@ class AnthropicResponseBuilder(ResponseBuilder):
|
||||
),
|
||||
]
|
||||
|
||||
def format_chunk(self, token: str) -> str:
|
||||
return sse_event(
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": token},
|
||||
},
|
||||
event="content_block_delta",
|
||||
)
|
||||
def format_chunk(self, token: str, **kwargs) -> List[str]:
|
||||
return [
|
||||
sse_event(
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": token},
|
||||
},
|
||||
event="content_block_delta",
|
||||
)
|
||||
]
|
||||
|
||||
def format_stream_end(self, ctx: GenContext, stop: StopInfo) -> List[str]:
|
||||
events: List[str] = []
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
"""OpenAI chat completion response builder."""
|
||||
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Tuple
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
@@ -11,24 +13,78 @@ from astrai.inference.api.protocol import (
|
||||
StopInfo,
|
||||
sse_event,
|
||||
)
|
||||
from astrai.inference.api.tool_parser import BaseToolParser, ToolParserFactory
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_UNSUPPORTED_PARAMS = (
|
||||
"n",
|
||||
"presence_penalty",
|
||||
"logit_bias",
|
||||
"user",
|
||||
)
|
||||
|
||||
|
||||
def _resolve_tool_choice(
|
||||
request: BaseModel,
|
||||
) -> Union[str, Dict[str, Any]]:
|
||||
tc = getattr(request, "tool_choice", None)
|
||||
if tc is None:
|
||||
return "auto"
|
||||
if isinstance(tc, str):
|
||||
return tc
|
||||
if isinstance(tc, dict):
|
||||
return tc
|
||||
return "auto"
|
||||
|
||||
|
||||
def _resolve_tools(request: BaseModel) -> Optional[List[Dict[str, Any]]]:
|
||||
raw = getattr(request, "tools", None)
|
||||
if not raw:
|
||||
return None
|
||||
if isinstance(raw, list):
|
||||
return [t.model_dump() if hasattr(t, "model_dump") else t for t in raw]
|
||||
return None
|
||||
|
||||
|
||||
class OpenAIResponseBuilder(ResponseBuilder):
|
||||
def prepare(
|
||||
self, request: BaseModel, engine: InferenceEngine
|
||||
) -> Tuple[str, GenContext, List[str]]:
|
||||
messages = [{"role": m.role, "content": m.content} for m in request.messages]
|
||||
prompt = engine.tokenizer.apply_chat_template(messages, tokenize=False)
|
||||
tools = _resolve_tools(request)
|
||||
prompt = engine.tokenizer.apply_chat_template(
|
||||
messages, tokenize=False, tools=tools or []
|
||||
)
|
||||
|
||||
self._resp_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
||||
self._model = request.model
|
||||
|
||||
for param in _UNSUPPORTED_PARAMS:
|
||||
value = getattr(request, param, None)
|
||||
fields = getattr(type(request), "model_fields", {})
|
||||
default = fields[param].default if param in fields else None
|
||||
if value is not None and value != default:
|
||||
logger.warning(
|
||||
"ChatCompletionRequest param '%s'=%r is not supported"
|
||||
" and will be ignored",
|
||||
param,
|
||||
value,
|
||||
)
|
||||
|
||||
self._parser: Optional[BaseToolParser] = None
|
||||
if tools:
|
||||
tool_choice = _resolve_tool_choice(request)
|
||||
self._parser = ToolParserFactory.create(
|
||||
"simple_json", tools=tools, tool_choice=tool_choice
|
||||
)
|
||||
self._content_started = False
|
||||
|
||||
ctx = GenContext(
|
||||
resp_id=self._resp_id,
|
||||
created=0,
|
||||
created=int(time.time()),
|
||||
model=self._model,
|
||||
prompt_tokens=0,
|
||||
)
|
||||
stop = request.stop
|
||||
stop_sequences = (
|
||||
@@ -55,7 +111,82 @@ class OpenAIResponseBuilder(ResponseBuilder):
|
||||
)
|
||||
]
|
||||
|
||||
def format_chunk(self, token: str) -> str:
|
||||
def format_chunk(self, token: str, **kwargs) -> List[str]:
|
||||
body = kwargs.get("body", "")
|
||||
if self._parser is not None:
|
||||
return self._format_tool_chunk(body, **kwargs)
|
||||
|
||||
return [
|
||||
sse_event(
|
||||
{
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 0,
|
||||
"model": self._model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"content": token},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
]
|
||||
|
||||
def _format_tool_chunk(self, body: str, **kwargs) -> List[str]:
|
||||
deltas = self._parser.feed(
|
||||
body,
|
||||
current_token_ids=kwargs.get("current_token_ids"),
|
||||
delta_token_ids=kwargs.get("delta_token_ids"),
|
||||
)
|
||||
events: List[str] = []
|
||||
for d in deltas:
|
||||
if "content" in d:
|
||||
if not self._content_started:
|
||||
events.append(self._role_chunk())
|
||||
self._content_started = True
|
||||
events.append(
|
||||
sse_event(
|
||||
{
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 0,
|
||||
"model": self._model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"content": d["content"]},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
)
|
||||
elif "tool_calls" in d:
|
||||
if not self._content_started:
|
||||
events.append(self._role_chunk())
|
||||
self._content_started = True
|
||||
events.append(
|
||||
sse_event(
|
||||
{
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 0,
|
||||
"model": self._model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"tool_calls": d["tool_calls"]},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
)
|
||||
return events
|
||||
|
||||
def _role_chunk(self) -> str:
|
||||
return sse_event(
|
||||
{
|
||||
"id": self._resp_id,
|
||||
@@ -63,12 +194,19 @@ class OpenAIResponseBuilder(ResponseBuilder):
|
||||
"created": 0,
|
||||
"model": self._model,
|
||||
"choices": [
|
||||
{"index": 0, "delta": {"content": token}, "finish_reason": None}
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"role": "assistant"},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
def format_stream_end(self, ctx: GenContext, stop: StopInfo) -> List[str]:
|
||||
finish_reason = "stop"
|
||||
if self._parser is not None and self._parser.has_tool_calls:
|
||||
finish_reason = "tool_calls"
|
||||
return [
|
||||
sse_event(
|
||||
{
|
||||
@@ -76,7 +214,9 @@ class OpenAIResponseBuilder(ResponseBuilder):
|
||||
"object": "chat.completion.chunk",
|
||||
"created": ctx.created,
|
||||
"model": self._model,
|
||||
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
|
||||
"choices": [
|
||||
{"index": 0, "delta": {}, "finish_reason": finish_reason}
|
||||
],
|
||||
}
|
||||
),
|
||||
sse_event(
|
||||
@@ -91,6 +231,32 @@ class OpenAIResponseBuilder(ResponseBuilder):
|
||||
def format_response(
|
||||
self, ctx: GenContext, content: str, stop: StopInfo
|
||||
) -> Dict[str, Any]:
|
||||
if self._parser is not None:
|
||||
parsed = self._parser.parse_complete(content)
|
||||
if parsed and parsed.get("tool_calls"):
|
||||
return {
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion",
|
||||
"created": ctx.created,
|
||||
"model": self._model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": parsed.get("content"),
|
||||
"tool_calls": parsed["tool_calls"],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": ctx.prompt_tokens,
|
||||
"completion_tokens": ctx.completion_tokens,
|
||||
"total_tokens": ctx.prompt_tokens + ctx.completion_tokens,
|
||||
},
|
||||
}
|
||||
|
||||
return {
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion",
|
||||
|
||||
@@ -35,7 +35,7 @@ class GenContext:
|
||||
resp_id: str
|
||||
created: int
|
||||
model: str
|
||||
prompt_tokens: int
|
||||
prompt_tokens: int = 0
|
||||
completion_tokens: int = 0
|
||||
|
||||
|
||||
@@ -64,7 +64,7 @@ class StopChecker:
|
||||
class ResponseBuilder(ABC):
|
||||
"""Interface for protocol-specific response formatting.
|
||||
|
||||
A new protocol requires one concrete builder implementing 6 methods.
|
||||
A new protocol requires one concrete builder implementing 5 methods.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
@@ -78,8 +78,15 @@ class ResponseBuilder(ABC):
|
||||
"""SSE events that open the stream."""
|
||||
|
||||
@abstractmethod
|
||||
def format_chunk(self, token: str) -> str:
|
||||
"""SSE event for a single generated token."""
|
||||
def format_chunk(self, token: str, **kwargs) -> List[str]:
|
||||
"""SSE events for a single generated token.
|
||||
|
||||
``body`` (the full accumulated text so far) is always provided
|
||||
as a keyword argument. Additional keyword arguments such as
|
||||
``current_token_ids`` and ``delta_token_ids`` may be included
|
||||
for tool parsers that need token-level information.
|
||||
Returns a list of SSE event strings (may be empty).
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def format_stream_end(self, ctx: GenContext, stop: StopInfo) -> List[str]:
|
||||
@@ -118,6 +125,7 @@ class ProtocolHandler:
|
||||
temperature=self.request.temperature,
|
||||
top_p=self.request.top_p,
|
||||
top_k=self.request.top_k,
|
||||
frequency_penalty=getattr(self.request, "frequency_penalty", 0.0),
|
||||
)
|
||||
|
||||
if self.request.stream:
|
||||
@@ -137,15 +145,25 @@ class ProtocolHandler:
|
||||
body = ""
|
||||
yielded = ""
|
||||
matched = None
|
||||
token_ids: List[int] = []
|
||||
async for token in agen:
|
||||
ctx.completion_tokens += 1
|
||||
body += token
|
||||
|
||||
new_ids = self.engine.tokenizer.encode(token)
|
||||
token_ids.extend(new_ids)
|
||||
|
||||
matched = checker.check(body)
|
||||
if matched:
|
||||
break
|
||||
|
||||
yield self.builder.format_chunk(token)
|
||||
ctx.completion_tokens += 1
|
||||
for event in self.builder.format_chunk(
|
||||
token,
|
||||
body=body,
|
||||
current_token_ids=token_ids,
|
||||
delta_token_ids=new_ids,
|
||||
):
|
||||
yield event
|
||||
yielded += token
|
||||
|
||||
stop = StopInfo(matched=matched, body=body, yielded=yielded)
|
||||
@@ -168,7 +186,6 @@ class ProtocolHandler:
|
||||
matched = None
|
||||
|
||||
async for token in agen:
|
||||
ctx.completion_tokens += 1
|
||||
chunks.append(token)
|
||||
body += token
|
||||
|
||||
@@ -176,6 +193,8 @@ class ProtocolHandler:
|
||||
if matched:
|
||||
break
|
||||
|
||||
ctx.completion_tokens += 1
|
||||
|
||||
content = "".join(chunks)
|
||||
stop = StopInfo(matched=matched, body=body)
|
||||
return self.builder.format_response(ctx, content, stop)
|
||||
|
||||
@@ -3,6 +3,9 @@ OpenAI / Anthropic-compatible chat completion server backed by continuous-batchi
|
||||
|
||||
Protocol-specific formatting is delegated to ``astrai.inference.protocol``.
|
||||
This module owns the FastAPI app, request/response schemas, and dependency wiring.
|
||||
|
||||
``app`` is lazily constructed — importing this module does NOT create a FastAPI instance.
|
||||
Use :func:`get_app` to access the singleton.
|
||||
"""
|
||||
|
||||
import logging
|
||||
@@ -12,7 +15,7 @@ from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import torch
|
||||
import uvicorn
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from fastapi import APIRouter, FastAPI, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from astrai.inference.api.anthropic import AnthropicResponseBuilder
|
||||
@@ -24,12 +27,25 @@ from astrai.tokenize import AutoTokenizer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_project_root = Path(__file__).parent.parent.parent
|
||||
_app_instance: Optional[FastAPI] = None
|
||||
|
||||
|
||||
class ChatMessage(BaseModel):
|
||||
role: str
|
||||
content: str
|
||||
content: Optional[str] = None
|
||||
tool_calls: Optional[List[Dict[str, Any]]] = None
|
||||
tool_call_id: Optional[str] = None
|
||||
|
||||
|
||||
class FunctionDef(BaseModel):
|
||||
name: str
|
||||
description: Optional[str] = None
|
||||
parameters: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
class ToolDef(BaseModel):
|
||||
type: str = "function"
|
||||
function: FunctionDef
|
||||
|
||||
|
||||
class ChatCompletionRequest(BaseModel):
|
||||
@@ -48,6 +64,8 @@ class ChatCompletionRequest(BaseModel):
|
||||
frequency_penalty: Optional[float] = Field(default=0.0, ge=-2.0, le=2.0)
|
||||
logit_bias: Optional[Dict[int, float]] = None
|
||||
user: Optional[str] = None
|
||||
tools: Optional[List[ToolDef]] = None
|
||||
tool_choice: Optional[Union[str, Dict[str, Any]]] = "auto"
|
||||
|
||||
|
||||
class AnthropicMessage(BaseModel):
|
||||
@@ -84,17 +102,15 @@ async def lifespan(app: FastAPI):
|
||||
logger.info("Inference engine shutdown complete")
|
||||
|
||||
|
||||
app = FastAPI(title="AstrAI Inference Server", version="0.2.0", lifespan=lifespan)
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _create_engine(
|
||||
param_path: Optional[Path] = None,
|
||||
param_path: Path,
|
||||
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}")
|
||||
|
||||
@@ -112,34 +128,50 @@ def _create_engine(
|
||||
return engine
|
||||
|
||||
|
||||
def get_app() -> FastAPI:
|
||||
"""Return the singleton FastAPI instance (lazily created on first call)."""
|
||||
global _app_instance
|
||||
if _app_instance is None:
|
||||
_app_instance = FastAPI(
|
||||
title="AstrAI Inference Server",
|
||||
version="0.2.0",
|
||||
lifespan=lifespan,
|
||||
)
|
||||
_app_instance.include_router(router)
|
||||
_app_instance.state.server_config = {}
|
||||
_app_instance.state.engine = None
|
||||
return _app_instance
|
||||
|
||||
|
||||
def _get_engine() -> InferenceEngine:
|
||||
engine = app.state.engine
|
||||
engine = get_app().state.engine
|
||||
if engine is None:
|
||||
raise HTTPException(status_code=503, detail="Engine not initialized")
|
||||
return engine
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
@router.get("/health")
|
||||
async def health():
|
||||
app = get_app()
|
||||
return {
|
||||
"status": "ok",
|
||||
"model_loaded": app.state.engine is not None,
|
||||
}
|
||||
|
||||
|
||||
@app.get("/stats")
|
||||
@router.get("/stats")
|
||||
async def get_stats():
|
||||
return _get_engine().get_stats()
|
||||
|
||||
|
||||
@app.post("/v1/chat/completions")
|
||||
@router.post("/v1/chat/completions")
|
||||
async def chat_completion(request: ChatCompletionRequest):
|
||||
engine = _get_engine()
|
||||
handler = ProtocolHandler(request, engine, OpenAIResponseBuilder())
|
||||
return await handler.handle()
|
||||
|
||||
|
||||
@app.post("/v1/messages")
|
||||
@router.post("/v1/messages")
|
||||
async def create_message(request: MessagesRequest):
|
||||
engine = _get_engine()
|
||||
handler = ProtocolHandler(request, engine, AnthropicResponseBuilder())
|
||||
@@ -147,14 +179,15 @@ async def create_message(request: MessagesRequest):
|
||||
|
||||
|
||||
def run_server(
|
||||
param_path: Path,
|
||||
host: str = "0.0.0.0",
|
||||
port: int = 8000,
|
||||
reload: bool = False,
|
||||
device: str = "cuda",
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
param_path: Optional[Path] = None,
|
||||
max_batch_size: int = 16,
|
||||
):
|
||||
app = get_app()
|
||||
app.state.server_config = {
|
||||
"device": device,
|
||||
"dtype": dtype,
|
||||
|
||||
@@ -0,0 +1,325 @@
|
||||
"""Tool call parsers for extracting structured tool calls from model output.
|
||||
|
||||
Patterned after vLLM's ToolParser abstraction. Each parser knows how to
|
||||
detect and incrementally extract tool calls from raw generated text.
|
||||
|
||||
Subclasses may optionally consume ``token_ids`` for token-level parsing
|
||||
(e.g. Harmony / VLM-style parsers).
|
||||
"""
|
||||
|
||||
import re
|
||||
import uuid
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
|
||||
class BaseToolParser(ABC):
|
||||
"""Abstract tool call parser — one instance per request.
|
||||
|
||||
Maintains streaming state internally so that each call to :meth:`feed`
|
||||
can diff against previously emitted content.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
tools : list of dict, optional
|
||||
Tool definitions from the request.
|
||||
tool_choice : str
|
||||
``"auto"`` / ``"required"`` / ``"none"`` or a named tool choice
|
||||
dict.
|
||||
"""
|
||||
|
||||
def __init__(self, tools: Optional[List[Dict]] = None, tool_choice: str = "auto"):
|
||||
self.tools = tools or []
|
||||
self.tool_choice = tool_choice
|
||||
|
||||
@abstractmethod
|
||||
def feed(
|
||||
self,
|
||||
body: str,
|
||||
current_token_ids: Optional[List[int]] = None,
|
||||
delta_token_ids: Optional[List[int]] = None,
|
||||
) -> List[Dict]:
|
||||
"""Feed the *full* accumulated text each step.
|
||||
|
||||
Returns a list of delta dicts to emit. Each delta is one of:
|
||||
|
||||
- ``{"content": "text"}`` — plain text delta
|
||||
- ``{"tool_calls": [...]}`` — tool-call delta (OpenAI format)
|
||||
|
||||
Returns an empty list when nothing new should be emitted.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
body : str
|
||||
The complete accumulated generated text so far.
|
||||
current_token_ids : list of int, optional
|
||||
All token IDs decoded into *body* (cumulative).
|
||||
delta_token_ids : list of int, optional
|
||||
Only the token IDs for this chunk.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def parse_complete(self, body: str) -> Optional[Dict]:
|
||||
"""Parse the *complete* generated text after generation ends.
|
||||
|
||||
Returns ``None`` when no tool calls were found, otherwise a dict
|
||||
with ``content`` (str or None) and ``tool_calls`` (list of dicts).
|
||||
"""
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def has_tool_calls(self) -> bool:
|
||||
"""True if the parser detected at least one tool call in the stream."""
|
||||
|
||||
|
||||
class ToolParserFactory(BaseFactory["BaseToolParser"]):
|
||||
pass
|
||||
|
||||
|
||||
_TOOL_CALL_HEAD_RE = re.compile(r'\{\s*"name"\s*:')
|
||||
|
||||
|
||||
def _scan_json(text: str, start: int = 0):
|
||||
"""Scan for a complete JSON object starting at *start*.
|
||||
|
||||
Returns ``(end, complete)`` where *end* is one-past the closing
|
||||
brace (or ``len(text)`` if unclosed), and *complete* is a bool.
|
||||
"""
|
||||
depth = 0
|
||||
in_string = False
|
||||
escape = False
|
||||
for i in range(start, len(text)):
|
||||
c = text[i]
|
||||
if escape:
|
||||
escape = False
|
||||
continue
|
||||
if c == "\\":
|
||||
escape = True
|
||||
continue
|
||||
if c == '"':
|
||||
in_string = not in_string
|
||||
continue
|
||||
if in_string:
|
||||
continue
|
||||
if c == "{":
|
||||
depth += 1
|
||||
elif c == "}":
|
||||
depth -= 1
|
||||
if depth == 0:
|
||||
return i + 1, True
|
||||
return len(text), False
|
||||
|
||||
|
||||
def _parse_tool_call_json(json_str: str, complete: bool):
|
||||
"""Extract *name* and *arguments* from a tool-call JSON string.
|
||||
|
||||
Returns ``(name, args, valid)``.
|
||||
"""
|
||||
name_match = re.search(r'"name"\s*:\s*"([^"]*)"', json_str)
|
||||
if not name_match:
|
||||
return None, "", False
|
||||
name = name_match.group(1)
|
||||
|
||||
args_match = re.search(r'"arguments"\s*:\s*(.*)', json_str, re.DOTALL)
|
||||
if not args_match:
|
||||
return name, "", True
|
||||
|
||||
raw = args_match.group(1).rstrip()
|
||||
if complete and raw.endswith("}"):
|
||||
raw = raw[:-1].rstrip()
|
||||
if raw.startswith("{"):
|
||||
inner = raw[1:].rstrip()
|
||||
if inner.endswith("}"):
|
||||
inner = inner[:-1].rstrip()
|
||||
raw = inner
|
||||
return name, raw, True
|
||||
|
||||
|
||||
def _find_tool_calls(text: str, start_pos: int = 0):
|
||||
"""Find all complete ``{...}`` tool-call objects in *text*.
|
||||
|
||||
Returns a list of dicts with keys *start*, *end*, *name*, *args*,
|
||||
*complete*.
|
||||
"""
|
||||
results = []
|
||||
pos = start_pos
|
||||
|
||||
while True:
|
||||
brace = text.find("{", pos)
|
||||
if brace == -1:
|
||||
break
|
||||
|
||||
end, complete = _scan_json(text, brace)
|
||||
if not complete:
|
||||
break
|
||||
|
||||
json_str = text[brace:end]
|
||||
if not _TOOL_CALL_HEAD_RE.search(json_str):
|
||||
pos = end
|
||||
continue
|
||||
|
||||
name, args, valid = _parse_tool_call_json(json_str, complete=True)
|
||||
if not valid or name is None:
|
||||
pos = end
|
||||
continue
|
||||
|
||||
results.append(
|
||||
{
|
||||
"start": brace,
|
||||
"end": end,
|
||||
"name": name,
|
||||
"args": args,
|
||||
"complete": True,
|
||||
}
|
||||
)
|
||||
pos = end
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def _find_partial_tool_call(text: str, start_pos: int = 0):
|
||||
"""Find one incomplete (still-generating) tool-call JSON object."""
|
||||
brace = text.find("{", start_pos)
|
||||
if brace == -1:
|
||||
return None
|
||||
|
||||
json_str = text[brace:]
|
||||
if not _TOOL_CALL_HEAD_RE.search(json_str):
|
||||
return None
|
||||
|
||||
name, args, valid = _parse_tool_call_json(json_str, complete=False)
|
||||
if not valid or name is None:
|
||||
return None
|
||||
|
||||
return {
|
||||
"start": brace,
|
||||
"name": name,
|
||||
"args": args,
|
||||
"complete": False,
|
||||
}
|
||||
|
||||
|
||||
@ToolParserFactory.register("simple_json")
|
||||
class SimpleJsonToolParser(BaseToolParser):
|
||||
"""Parser for models that output tool calls as plain JSON objects.
|
||||
|
||||
Detects ``{"name": "<func>", "arguments": {...}}`` anywhere in the
|
||||
generated text. Handles single and (non-overlapping) multiple tool
|
||||
calls. Text preceding the first tool call is emitted as plain
|
||||
``content`` deltas.
|
||||
"""
|
||||
|
||||
def __init__(self, tools=None, tool_choice="auto"):
|
||||
super().__init__(tools, tool_choice)
|
||||
self._emitted_content_len = 0
|
||||
self._tc_state: List[Dict] = []
|
||||
self._has_tool_calls = False
|
||||
|
||||
# -------------------------------------------------------------- feed
|
||||
|
||||
def feed(
|
||||
self,
|
||||
body: str,
|
||||
current_token_ids: Optional[List[int]] = None,
|
||||
delta_token_ids: Optional[List[int]] = None,
|
||||
) -> List[Dict]:
|
||||
deltas: List[Dict] = []
|
||||
|
||||
completed = _find_tool_calls(body)
|
||||
|
||||
if not completed:
|
||||
partial = _find_partial_tool_call(body)
|
||||
if not partial:
|
||||
return self._emit_plain_content(body, deltas)
|
||||
all_tcs = [partial]
|
||||
else:
|
||||
all_tcs = completed
|
||||
partial = _find_partial_tool_call(body, completed[-1]["end"])
|
||||
if partial:
|
||||
all_tcs = completed + [partial]
|
||||
|
||||
first_start = all_tcs[0]["start"]
|
||||
if first_start > self._emitted_content_len:
|
||||
content = body[self._emitted_content_len : first_start]
|
||||
self._emitted_content_len = first_start
|
||||
if content:
|
||||
deltas.append({"content": content})
|
||||
|
||||
for i, tc in enumerate(all_tcs):
|
||||
if i >= len(self._tc_state):
|
||||
self._tc_state.append(
|
||||
{
|
||||
"id": f"call_{uuid.uuid4().hex[:12]}",
|
||||
"name_emitted": False,
|
||||
"args_emitted_len": 0,
|
||||
}
|
||||
)
|
||||
self._has_tool_calls = True
|
||||
st = self._tc_state[i]
|
||||
|
||||
if not st["name_emitted"]:
|
||||
st["name_emitted"] = True
|
||||
deltas.append(
|
||||
{
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": i,
|
||||
"id": st["id"],
|
||||
"type": "function",
|
||||
"function": {"name": tc["name"], "arguments": ""},
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
|
||||
new_args = tc["args"]
|
||||
if len(new_args) > st["args_emitted_len"]:
|
||||
diff = new_args[st["args_emitted_len"] :]
|
||||
st["args_emitted_len"] = len(new_args)
|
||||
deltas.append(
|
||||
{
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": i,
|
||||
"function": {"arguments": diff},
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
|
||||
return deltas
|
||||
|
||||
def _emit_plain_content(self, body: str, deltas: List[Dict]) -> List[Dict]:
|
||||
new_content = body[self._emitted_content_len :]
|
||||
if new_content:
|
||||
self._emitted_content_len = len(body)
|
||||
deltas.append({"content": new_content})
|
||||
return deltas
|
||||
|
||||
# -------------------------------------------------------- complete
|
||||
|
||||
def parse_complete(self, body: str) -> Optional[Dict]:
|
||||
completed = _find_tool_calls(body)
|
||||
if not completed:
|
||||
return None
|
||||
|
||||
content = body[: completed[0]["start"]].strip() or None
|
||||
tool_calls = []
|
||||
for i, tc in enumerate(completed):
|
||||
tool_calls.append(
|
||||
{
|
||||
"id": f"call_{uuid.uuid4().hex[:12]}",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tc["name"],
|
||||
"arguments": tc["args"],
|
||||
},
|
||||
}
|
||||
)
|
||||
return {"content": content, "tool_calls": tool_calls}
|
||||
|
||||
@property
|
||||
def has_tool_calls(self) -> bool:
|
||||
return self._has_tool_calls
|
||||
@@ -2,8 +2,12 @@
|
||||
|
||||
from astrai.inference.core.cache import (
|
||||
Allocator,
|
||||
CacheView,
|
||||
ContiguousCache,
|
||||
ContiguousCacheView,
|
||||
KVCache,
|
||||
KvcacheView,
|
||||
PageCache,
|
||||
PageCacheView,
|
||||
PagePool,
|
||||
PrefixCache,
|
||||
Storage,
|
||||
@@ -16,8 +20,12 @@ from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
||||
|
||||
__all__ = [
|
||||
"Allocator",
|
||||
"CacheView",
|
||||
"KVCache",
|
||||
"KvcacheView",
|
||||
"ContiguousCache",
|
||||
"ContiguousCacheView",
|
||||
"PageCache",
|
||||
"PageCacheView",
|
||||
"PagePool",
|
||||
"PrefixCache",
|
||||
"Storage",
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import threading
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import OrderedDict
|
||||
from typing import Callable, Dict, List, Optional, Tuple
|
||||
|
||||
@@ -62,7 +63,8 @@ class Allocator:
|
||||
|
||||
def touch(self, idx: int):
|
||||
with self._lock:
|
||||
self._lru.move_to_end(idx)
|
||||
if idx in self._lru:
|
||||
self._lru.move_to_end(idx)
|
||||
|
||||
|
||||
class PrefixCache:
|
||||
@@ -274,7 +276,46 @@ class Storage:
|
||||
return k, v
|
||||
|
||||
|
||||
class KvcacheView:
|
||||
class CacheView(ABC):
|
||||
"""Abstract view passed to attention layers for KV-cache I/O."""
|
||||
|
||||
@abstractmethod
|
||||
def write(self, layer_id: int, k: Tensor, v: Tensor): ...
|
||||
|
||||
@abstractmethod
|
||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]: ...
|
||||
|
||||
|
||||
class KVCache(ABC):
|
||||
"""Abstract KV-cache facade for scheduler/executor."""
|
||||
|
||||
@abstractmethod
|
||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool: ...
|
||||
|
||||
@abstractmethod
|
||||
def task_free(self, task_id: str): ...
|
||||
|
||||
@abstractmethod
|
||||
def task_extend(self, task_id: str, pos: int) -> bool: ...
|
||||
|
||||
@abstractmethod
|
||||
def bind_tasks(
|
||||
self,
|
||||
task_ids: List[str],
|
||||
total_len: int,
|
||||
device: torch.device,
|
||||
write_positions: Optional[Tensor] = None,
|
||||
) -> CacheView: ...
|
||||
|
||||
def task_cached(self, task_id: str) -> int:
|
||||
return 0
|
||||
|
||||
def task_record_hashes(
|
||||
self, task_id: str, prompt_ids: List[int], start_logical_page: int = 0
|
||||
): ...
|
||||
|
||||
|
||||
class PageCacheView(CacheView):
|
||||
"""Bundles Storage + page_table + total_len for attention layers."""
|
||||
|
||||
def __init__(self, storage: Storage, page_table: Tensor, total_len: int = 0):
|
||||
@@ -290,8 +331,8 @@ class KvcacheView:
|
||||
return self._storage.gather(layer_id, self._page_table, self._total_len)
|
||||
|
||||
|
||||
class KVCache:
|
||||
"""Facade: page management + KV-cache I/O for continuous batching."""
|
||||
class PageCache(KVCache):
|
||||
"""Paged KV-cache with prefix sharing."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -361,8 +402,132 @@ class KVCache:
|
||||
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_tasks(
|
||||
self,
|
||||
task_ids: List[str],
|
||||
total_len: int,
|
||||
device: torch.device,
|
||||
write_positions: Optional[Tensor] = None,
|
||||
) -> PageCacheView:
|
||||
page_table = self._table.table_tensor(task_ids, device)
|
||||
return PageCacheView(self._storage, page_table, total_len)
|
||||
|
||||
def bind(self, page_table: Tensor, total_len: int = 0) -> KvcacheView:
|
||||
return KvcacheView(self._storage, page_table, total_len)
|
||||
|
||||
class ContiguousCacheView(CacheView):
|
||||
"""Contiguous KV-cache view for attention layers."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
cache: "ContiguousCache",
|
||||
batch_indices: Tensor,
|
||||
total_len: int = 0,
|
||||
write_positions: Optional[Tensor] = None,
|
||||
):
|
||||
self._cache = cache
|
||||
self._batch_indices = batch_indices
|
||||
self._total_len = total_len
|
||||
self._write_positions = write_positions
|
||||
|
||||
def write(self, layer_id: int, k: Tensor, v: Tensor):
|
||||
seq_len = k.size(1)
|
||||
indices = self._batch_indices
|
||||
if self._write_positions is not None and seq_len == 1:
|
||||
pos = self._write_positions
|
||||
self._cache.k[layer_id, indices, pos] = k.squeeze(1)
|
||||
self._cache.v[layer_id, indices, pos] = v.squeeze(1)
|
||||
for s, p in zip(indices.tolist(), pos.tolist()):
|
||||
cur = self._cache._slot_len.get(s, 0)
|
||||
if p + 1 > cur:
|
||||
self._cache._slot_len[s] = p + 1
|
||||
else:
|
||||
start_pos = self._total_len - seq_len
|
||||
self._cache.k[layer_id, indices, start_pos : start_pos + seq_len] = k
|
||||
self._cache.v[layer_id, indices, start_pos : start_pos + seq_len] = v
|
||||
new_len = start_pos + seq_len
|
||||
for s in indices.tolist():
|
||||
cur = self._cache._slot_len.get(s, 0)
|
||||
if new_len > cur:
|
||||
self._cache._slot_len[s] = new_len
|
||||
|
||||
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
|
||||
max_len = max(
|
||||
self._cache._slot_len.get(int(s), 0) for s in self._batch_indices.tolist()
|
||||
)
|
||||
indices = self._batch_indices
|
||||
k = self._cache.k[layer_id, indices, :max_len]
|
||||
v = self._cache.v[layer_id, indices, :max_len]
|
||||
return k, v
|
||||
|
||||
|
||||
class ContiguousCache(KVCache):
|
||||
"""Contiguous per-slot KV cache (default implementation)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
n_layers: int,
|
||||
max_batch_size: int,
|
||||
max_seq_len: int,
|
||||
n_kv_heads: int,
|
||||
head_dim: int,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
self.max_seq_len = max_seq_len
|
||||
self.k = torch.zeros(
|
||||
n_layers,
|
||||
max_batch_size,
|
||||
max_seq_len,
|
||||
n_kv_heads,
|
||||
head_dim,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
self.v = torch.zeros(
|
||||
n_layers,
|
||||
max_batch_size,
|
||||
max_seq_len,
|
||||
n_kv_heads,
|
||||
head_dim,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
self._slot_len: Dict[int, int] = {}
|
||||
self._task_slot: Dict[str, int] = {}
|
||||
self._free_slots = list(range(max_batch_size))
|
||||
self._device = device
|
||||
|
||||
def task_alloc(self, task_id: str, prompt_ids: List[int]) -> bool:
|
||||
if not self._free_slots:
|
||||
return False
|
||||
slot = self._free_slots.pop(0)
|
||||
self._task_slot[task_id] = slot
|
||||
self._slot_len[slot] = 0
|
||||
return True
|
||||
|
||||
def task_free(self, task_id: str):
|
||||
slot = self._task_slot.pop(task_id, None)
|
||||
if slot is not None:
|
||||
self._slot_len.pop(slot, None)
|
||||
self._free_slots.append(slot)
|
||||
|
||||
def task_extend(self, task_id: str, pos: int) -> bool:
|
||||
return pos < self.max_seq_len
|
||||
|
||||
def task_cached(self, task_id: str) -> int:
|
||||
slot = self._task_slot.get(task_id)
|
||||
if slot is None:
|
||||
return 0
|
||||
return self._slot_len.get(slot, 0)
|
||||
|
||||
def bind_tasks(
|
||||
self,
|
||||
task_ids: List[str],
|
||||
total_len: int,
|
||||
device: torch.device,
|
||||
write_positions: Optional[Tensor] = None,
|
||||
) -> ContiguousCacheView:
|
||||
slots = [self._task_slot[tid] for tid in task_ids]
|
||||
batch_indices = torch.tensor(slots, dtype=torch.long, device=device)
|
||||
return ContiguousCacheView(
|
||||
self, batch_indices, total_len, write_positions=write_positions
|
||||
)
|
||||
|
||||
@@ -19,13 +19,13 @@ class Executor:
|
||||
self,
|
||||
model: AutoModel,
|
||||
tokenizer: AutoTokenizer,
|
||||
page_cache: KVCache,
|
||||
kv_cache: KVCache,
|
||||
device: Optional[str] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
self.model = model
|
||||
self.tokenizer = tokenizer
|
||||
self.page_cache = page_cache
|
||||
self.kv_cache = kv_cache
|
||||
self.device = device or next(model.parameters()).device
|
||||
self.dtype = dtype or next(model.parameters()).dtype
|
||||
|
||||
@@ -43,7 +43,6 @@ class Executor:
|
||||
)
|
||||
|
||||
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(
|
||||
@@ -53,7 +52,7 @@ class Executor:
|
||||
)
|
||||
.unsqueeze(0)
|
||||
.expand(batch_sz, -1),
|
||||
paged_cache=self.page_cache.bind(page_tables, total_len=prompt_len),
|
||||
paged_cache=self.kv_cache.bind_tasks(task_ids, prompt_len, self.device),
|
||||
)
|
||||
|
||||
def execute_decode(self, tasks: List[Task]) -> List[int]:
|
||||
@@ -72,16 +71,47 @@ class Executor:
|
||||
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)
|
||||
freq_penalties = torch.tensor(
|
||||
[t.frequency_penalty for t in tasks], device=self.device
|
||||
)
|
||||
|
||||
history_lists = []
|
||||
mask_lists = []
|
||||
for t in tasks:
|
||||
window = t.rep_window
|
||||
prompt_part = t.prompt_ids[-window:]
|
||||
ids = prompt_part + t.output_ids
|
||||
history_lists.append(ids)
|
||||
mask_lists.append([True] * len(ids))
|
||||
|
||||
max_len = max(len(h) for h in history_lists)
|
||||
padded_ids = torch.zeros(
|
||||
len(tasks), max_len, dtype=torch.long, device=self.device
|
||||
)
|
||||
padded_mask = torch.zeros(
|
||||
len(tasks), max_len, dtype=torch.bool, device=self.device
|
||||
)
|
||||
for i, (h, m) in enumerate(zip(history_lists, mask_lists)):
|
||||
padded_ids[i, : len(h)] = torch.tensor(
|
||||
h, dtype=torch.long, device=self.device
|
||||
)
|
||||
padded_mask[i, : len(m)] = torch.tensor(
|
||||
m, dtype=torch.bool, 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),
|
||||
paged_cache=self.kv_cache.bind_tasks(
|
||||
task_ids,
|
||||
total_len,
|
||||
self.device,
|
||||
write_positions=position_ids,
|
||||
),
|
||||
position_ids=position_ids.unsqueeze(1),
|
||||
)
|
||||
logits = outputs["logits"][:, -1, :]
|
||||
@@ -91,4 +121,7 @@ class Executor:
|
||||
temperature=temperatures,
|
||||
top_k=top_ks,
|
||||
top_p=top_ps,
|
||||
frequency_penalty=freq_penalties,
|
||||
input_ids=padded_ids,
|
||||
input_mask=padded_mask,
|
||||
).tolist()
|
||||
|
||||
@@ -4,7 +4,7 @@ from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.inference.core.cache import KVCache
|
||||
from astrai.inference.core.cache import ContiguousCache, KVCache
|
||||
from astrai.inference.core.executor import Executor
|
||||
from astrai.inference.core.task import STOP, Task, TaskManager, TaskStatus
|
||||
from astrai.model.automodel import AutoModel
|
||||
@@ -14,7 +14,7 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class InferenceScheduler:
|
||||
"""Four-phase continuous batching loop: cleanup -> refill -> prefill -> decode."""
|
||||
"""Continuous batching loop: cleanup -> refill -> prefill -> decode (all groups)."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -23,9 +23,9 @@ class InferenceScheduler:
|
||||
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,
|
||||
cache: Optional[KVCache] = None,
|
||||
):
|
||||
config = model.config
|
||||
|
||||
@@ -41,19 +41,20 @@ class InferenceScheduler:
|
||||
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
|
||||
head_dim = config.dim // config.n_heads
|
||||
|
||||
self._page_cache = KVCache(
|
||||
config.n_layers,
|
||||
n_pages,
|
||||
page_size,
|
||||
config.n_kv_heads,
|
||||
config.dim // config.n_heads,
|
||||
self.device,
|
||||
self.dtype,
|
||||
)
|
||||
if cache is not None:
|
||||
self._cache = cache
|
||||
else:
|
||||
self._cache = ContiguousCache(
|
||||
config.n_layers,
|
||||
max_batch_size,
|
||||
self.max_seq_len,
|
||||
config.n_kv_heads,
|
||||
head_dim,
|
||||
self.device,
|
||||
self.dtype,
|
||||
)
|
||||
|
||||
self._task_mgr = TaskManager(
|
||||
tokenizer=tokenizer,
|
||||
@@ -65,30 +66,32 @@ class InferenceScheduler:
|
||||
self._executor = Executor(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
page_cache=self._page_cache,
|
||||
kv_cache=self._cache,
|
||||
device=self.device,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
|
||||
self._running = False
|
||||
self._stop_event = threading.Event()
|
||||
self._loop_thread: Optional[threading.Thread] = None
|
||||
|
||||
def add_task(self, prompt: str, **kwargs) -> str:
|
||||
return self._task_mgr.add_task(prompt, **kwargs)
|
||||
|
||||
def remove_task(self, task_id: str):
|
||||
for task in self._task_mgr.remove_task(task_id):
|
||||
self._page_cache.task_free(task.task_id)
|
||||
self._cache.task_free(task.task_id)
|
||||
|
||||
def get_stats(self) -> Dict[str, Any]:
|
||||
return self._task_mgr.get_stats()
|
||||
|
||||
def _run_generation_loop(self):
|
||||
stop_ids = self._task_mgr.tokenizer.stop_ids
|
||||
cache = self._cache
|
||||
try:
|
||||
while self._running:
|
||||
while not self._stop_event.is_set():
|
||||
finished = self._task_mgr.remove_finished_tasks(stop_ids)
|
||||
for task in finished:
|
||||
self._page_cache.task_free(task.task_id)
|
||||
cache.task_free(task.task_id)
|
||||
|
||||
active = self._task_mgr.get_active_tasks()
|
||||
available = self._task_mgr.max_batch_size - len(active)
|
||||
@@ -96,7 +99,7 @@ class InferenceScheduler:
|
||||
candidates = self._task_mgr.pull_candidates(available)
|
||||
failed = []
|
||||
for task in candidates:
|
||||
if self._page_cache.task_alloc(task.task_id, task.prompt_ids):
|
||||
if cache.task_alloc(task.task_id, task.prompt_ids):
|
||||
self._task_mgr.activate(task)
|
||||
else:
|
||||
failed.append(task)
|
||||
@@ -111,7 +114,7 @@ class InferenceScheduler:
|
||||
t
|
||||
for t in self._task_mgr.get_active_tasks()
|
||||
if t.output_tokens == 0
|
||||
and self._page_cache.task_cached(t.task_id) < len(t.prompt_ids)
|
||||
and cache.task_cached(t.task_id) < len(t.prompt_ids)
|
||||
]
|
||||
if to_prefill:
|
||||
for t in to_prefill:
|
||||
@@ -121,85 +124,76 @@ class InferenceScheduler:
|
||||
for t in to_prefill:
|
||||
key = (
|
||||
len(t.prompt_ids),
|
||||
self._page_cache.task_cached(t.task_id),
|
||||
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
|
||||
start_logical_page = start_pos // getattr(
|
||||
cache, "page_size", 64
|
||||
)
|
||||
for t in group:
|
||||
self._page_cache.task_record_hashes(
|
||||
t.task_id,
|
||||
t.prompt_ids,
|
||||
start_logical_page=start_logical_page,
|
||||
cache.task_record_hashes(
|
||||
t.task_id, t.prompt_ids, 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)
|
||||
decode_tasks = self._task_mgr.get_active_tasks()
|
||||
|
||||
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 sorted(decode_tasks, key=lambda t: t.task_id):
|
||||
if cache.task_extend(t.task_id, t.next_pos):
|
||||
valid.append(t)
|
||||
else:
|
||||
t.status = TaskStatus.ABORTED
|
||||
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||
|
||||
valid: List[Task] = []
|
||||
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)
|
||||
|
||||
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
|
||||
new_text = t.decode_new_token(self._task_mgr.tokenizer)
|
||||
if new_text:
|
||||
self._task_mgr.invoke_callback(t.task_id, new_text)
|
||||
|
||||
for t, ntok in zip(valid, next_tokens):
|
||||
t.output_ids.append(ntok)
|
||||
t.output_tokens += 1
|
||||
pos = t.input_tokens + t.output_tokens
|
||||
extend_ok = self._page_cache.task_extend(t.task_id, pos)
|
||||
if t.stream_callback:
|
||||
t.stream_callback(
|
||||
self._task_mgr.tokenizer.decode([ntok])
|
||||
)
|
||||
if not extend_ok:
|
||||
t.status = TaskStatus.ABORTED
|
||||
if t.stream_callback:
|
||||
t.stream_callback(STOP)
|
||||
|
||||
for t in valid:
|
||||
if t.is_finished(stop_ids):
|
||||
if t.stream_callback:
|
||||
t.stream_callback(STOP)
|
||||
for t in valid:
|
||||
if t.is_finished(stop_ids):
|
||||
remaining = t.flush_remaining(self._task_mgr.tokenizer)
|
||||
if remaining:
|
||||
self._task_mgr.invoke_callback(t.task_id, remaining)
|
||||
self._task_mgr.invoke_callback(t.task_id, STOP)
|
||||
|
||||
except Exception as e:
|
||||
self._stop_event.set()
|
||||
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
|
||||
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.invoke_callback(task.task_id, STOP)
|
||||
cache.task_free(task.task_id)
|
||||
for task in self._task_mgr.get_waiting_tasks():
|
||||
if task.stream_callback:
|
||||
task.stream_callback(STOP)
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._task_mgr.clear_queues()
|
||||
raise
|
||||
|
||||
def start(self):
|
||||
if not self._running:
|
||||
self._running = True
|
||||
t = threading.Thread(target=self._run_generation_loop, daemon=True)
|
||||
t.start()
|
||||
self._loop_thread = t
|
||||
if self._loop_thread is not None and self._loop_thread.is_alive():
|
||||
return
|
||||
self._stop_event.clear()
|
||||
t = threading.Thread(target=self._run_generation_loop, daemon=True)
|
||||
t.start()
|
||||
self._loop_thread = t
|
||||
|
||||
def stop(self):
|
||||
self._running = False
|
||||
self._stop_event.set()
|
||||
self._task_mgr.wake()
|
||||
if hasattr(self, "_loop_thread"):
|
||||
if self._loop_thread is not None:
|
||||
self._loop_thread.join(timeout=2.0)
|
||||
self._loop_thread = None
|
||||
for task in self._task_mgr.get_active_tasks():
|
||||
self._page_cache.task_free(task.task_id)
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._cache.task_free(task.task_id)
|
||||
for task in self._task_mgr.get_waiting_tasks():
|
||||
self._task_mgr.invoke_callback(task.task_id, STOP)
|
||||
self._task_mgr.clear_queues()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
@@ -13,6 +13,40 @@ logger = logging.getLogger(__name__)
|
||||
STOP = object()
|
||||
|
||||
|
||||
class StreamDecoder:
|
||||
"""Incremental decoder for byte-level BPE streaming.
|
||||
|
||||
Byte-level BPE may split a single Unicode character (e.g. em-dash,
|
||||
smart quotes) across multiple tokens. Decoding such a token in
|
||||
isolation produces U+FFFD (replacement char). This decoder
|
||||
accumulates token IDs and only emits text once the trailing
|
||||
characters are complete, buffering incomplete multi-byte sequences
|
||||
until the next token arrives.
|
||||
"""
|
||||
|
||||
__slots__ = ("_tokenizer", "_ids", "_emitted")
|
||||
|
||||
def __init__(self, tokenizer: AutoTokenizer):
|
||||
self._tokenizer = tokenizer
|
||||
self._ids: List[int] = []
|
||||
self._emitted: str = ""
|
||||
|
||||
def push(self, token_id: int) -> str:
|
||||
"""Append a token ID and return newly completed text.
|
||||
|
||||
Returns "" while a multi-byte character is still incomplete.
|
||||
"""
|
||||
self._ids.append(token_id)
|
||||
full = self._tokenizer.decode(self._ids, skip_special_tokens=True)
|
||||
if full.endswith("\ufffd"):
|
||||
return ""
|
||||
if len(full) > len(self._emitted):
|
||||
diff = full[len(self._emitted) :]
|
||||
self._emitted = full
|
||||
return diff
|
||||
return ""
|
||||
|
||||
|
||||
class TaskStatus(Enum):
|
||||
"""Task lifecycle states."""
|
||||
|
||||
@@ -33,7 +67,8 @@ class Task:
|
||||
temperature: float = 1.0,
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
stream_callback: Optional[Callable[[str], None]] = None,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
):
|
||||
self.task_id = task_id
|
||||
self.prompt_ids = prompt_ids
|
||||
@@ -41,6 +76,8 @@ class Task:
|
||||
self.temperature = temperature
|
||||
self.top_p = top_p
|
||||
self.top_k = top_k
|
||||
self.frequency_penalty = frequency_penalty
|
||||
self.rep_window = rep_window
|
||||
|
||||
self.status = TaskStatus.PENDING
|
||||
self.output_ids: List[int] = []
|
||||
@@ -48,7 +85,34 @@ class Task:
|
||||
self.output_tokens: int = 0
|
||||
self.arrival_time = time.time()
|
||||
self.finish_time: Optional[float] = None
|
||||
self.stream_callback = stream_callback
|
||||
self._decoder: Optional[StreamDecoder] = None
|
||||
|
||||
def decode_new_token(self, tokenizer: AutoTokenizer) -> str:
|
||||
"""Decode the last appended output token, buffering incomplete
|
||||
multi-byte sequences across calls.
|
||||
|
||||
Lazily creates a :class:`StreamDecoder` on first use.
|
||||
"""
|
||||
if self._decoder is None:
|
||||
self._decoder = StreamDecoder(tokenizer)
|
||||
return self._decoder.push(self.output_ids[-1])
|
||||
|
||||
def flush_remaining(self, tokenizer: AutoTokenizer) -> str:
|
||||
"""Emit any text still buffered in the decoder.
|
||||
|
||||
Called when generation terminates (max_tokens reached, stop
|
||||
sequence, or external removal) to avoid dropping a final
|
||||
incomplete-looking fragment that is actually complete when
|
||||
adjacent to the stop token.
|
||||
"""
|
||||
if self._decoder is None or not self.output_ids:
|
||||
return ""
|
||||
full = tokenizer.decode(self.output_ids, skip_special_tokens=True)
|
||||
if len(full) > len(self._decoder._emitted):
|
||||
diff = full[len(self._decoder._emitted) :]
|
||||
self._decoder._emitted = full
|
||||
return diff
|
||||
return ""
|
||||
|
||||
@property
|
||||
def next_pos(self) -> int:
|
||||
@@ -79,6 +143,7 @@ class TaskManager:
|
||||
|
||||
self.waiting_queue: Deque[Task] = deque()
|
||||
self.active_tasks: List[Task] = []
|
||||
self._callbacks: Dict[str, Callable[[str], None]] = {}
|
||||
|
||||
self._task_event = threading.Event()
|
||||
self._lock = threading.Lock()
|
||||
@@ -93,6 +158,8 @@ class TaskManager:
|
||||
temperature: float = 1.0,
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
stream_callback: Optional[Callable[[str], None]] = None,
|
||||
) -> str:
|
||||
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
|
||||
@@ -117,12 +184,15 @@ class TaskManager:
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
stream_callback=stream_callback,
|
||||
frequency_penalty=frequency_penalty,
|
||||
rep_window=rep_window,
|
||||
)
|
||||
|
||||
with self._lock:
|
||||
self.waiting_queue.append(task)
|
||||
self._total_tasks += 1
|
||||
if stream_callback:
|
||||
self._callbacks[task_id] = stream_callback
|
||||
|
||||
self._task_event.set()
|
||||
return task_id
|
||||
@@ -134,8 +204,14 @@ class TaskManager:
|
||||
t for t in self.waiting_queue if t.task_id != task_id
|
||||
)
|
||||
self.active_tasks = [t for t in self.active_tasks if t.task_id != task_id]
|
||||
self._callbacks.pop(task_id, None)
|
||||
return removed_active
|
||||
|
||||
def invoke_callback(self, task_id: str, token: str):
|
||||
cb = self._callbacks.get(task_id)
|
||||
if cb:
|
||||
cb(token)
|
||||
|
||||
def get_stats(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"total_tasks": self._total_tasks,
|
||||
@@ -186,7 +262,10 @@ class TaskManager:
|
||||
return bool(self.active_tasks or self.waiting_queue)
|
||||
|
||||
def wait_for_tasks(self, timeout: float = 1.0):
|
||||
self._task_event.clear()
|
||||
with self._lock:
|
||||
if self.waiting_queue or self.active_tasks:
|
||||
return
|
||||
self._task_event.clear()
|
||||
self._task_event.wait(timeout=timeout)
|
||||
|
||||
def get_active_tasks(self) -> List[Task]:
|
||||
@@ -201,6 +280,7 @@ class TaskManager:
|
||||
with self._lock:
|
||||
self.waiting_queue.clear()
|
||||
self.active_tasks.clear()
|
||||
self._callbacks.clear()
|
||||
|
||||
def wake(self):
|
||||
self._task_event.set()
|
||||
|
||||
@@ -8,6 +8,7 @@ from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Tuple,
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from astrai.inference.core.cache import KVCache
|
||||
from astrai.inference.core.scheduler import InferenceScheduler
|
||||
from astrai.inference.core.task import STOP
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
@@ -73,20 +74,31 @@ class GenerationRequest:
|
||||
top_p: float = 1.0,
|
||||
temperature: float = 1.0,
|
||||
max_tokens: Optional[int] = None,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
stream: bool = False,
|
||||
):
|
||||
if not (isinstance(top_k, int) and top_k >= 0):
|
||||
raise ValueError("top_k must be a non-negative integer")
|
||||
if not (0.0 <= top_p <= 1.0):
|
||||
raise ValueError("top_p must be a float between 0.0 and 1.0")
|
||||
if not (isinstance(temperature, (int, float)) and temperature >= 0):
|
||||
raise ValueError("temperature must be a non-negative number")
|
||||
if not (isinstance(temperature, (int, float)) and temperature > 0):
|
||||
raise ValueError("temperature must be a positive number")
|
||||
if not (
|
||||
isinstance(frequency_penalty, (int, float))
|
||||
and -2.0 <= frequency_penalty <= 2.0
|
||||
):
|
||||
raise ValueError("frequency_penalty must be between -2.0 and 2.0")
|
||||
if not (isinstance(rep_window, int) and rep_window > 0):
|
||||
raise ValueError("rep_window must be a positive integer")
|
||||
|
||||
self.messages = messages
|
||||
self.top_k = top_k
|
||||
self.top_p = top_p
|
||||
self.temperature = temperature
|
||||
self.max_tokens = max_tokens
|
||||
self.frequency_penalty = frequency_penalty
|
||||
self.rep_window = rep_window
|
||||
self.stream = stream
|
||||
|
||||
|
||||
@@ -101,6 +113,7 @@ class InferenceEngine:
|
||||
max_seq_len: Optional[int] = None,
|
||||
max_prompt_len: int = 2048,
|
||||
page_size: int = 128,
|
||||
cache: Optional[KVCache] = None,
|
||||
):
|
||||
self.model = model
|
||||
self.tokenizer = tokenizer
|
||||
@@ -110,7 +123,7 @@ class InferenceEngine:
|
||||
max_batch_size=max_batch_size,
|
||||
max_seq_len=max_seq_len,
|
||||
max_prompt_len=max_prompt_len,
|
||||
page_size=page_size,
|
||||
cache=cache,
|
||||
)
|
||||
|
||||
self.scheduler.start()
|
||||
@@ -130,17 +143,33 @@ class InferenceEngine:
|
||||
temperature: float = 1.0,
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
) -> Union[Generator, str, List[str]]:
|
||||
is_batch = isinstance(prompt, list)
|
||||
prompts = prompt if is_batch else [prompt]
|
||||
|
||||
if stream:
|
||||
return self._generate_streaming(
|
||||
prompts, is_batch, max_tokens, temperature, top_p, top_k
|
||||
prompts,
|
||||
is_batch,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
else:
|
||||
return self._generate_non_streaming(
|
||||
prompts, is_batch, max_tokens, temperature, top_p, top_k
|
||||
prompts,
|
||||
is_batch,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
|
||||
def generate_async(
|
||||
@@ -150,9 +179,18 @@ class InferenceEngine:
|
||||
temperature: float = 1.0,
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
frequency_penalty: float = 0.0,
|
||||
rep_window: int = 64,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
sync_gen = self._generate_streaming(
|
||||
[prompt], False, max_tokens, temperature, top_p, top_k
|
||||
[prompt],
|
||||
False,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
|
||||
async def _agen():
|
||||
@@ -183,6 +221,8 @@ class InferenceEngine:
|
||||
temperature=request.temperature,
|
||||
top_p=request.top_p,
|
||||
top_k=request.top_k,
|
||||
frequency_penalty=request.frequency_penalty,
|
||||
rep_window=request.rep_window,
|
||||
)
|
||||
|
||||
def _submit_tasks(
|
||||
@@ -192,6 +232,8 @@ class InferenceEngine:
|
||||
temperature: float,
|
||||
top_p: float,
|
||||
top_k: int,
|
||||
frequency_penalty: float,
|
||||
rep_window: int,
|
||||
) -> Tuple[GenerateResult, List[str]]:
|
||||
n = len(prompts)
|
||||
result = GenerateResult(count=n)
|
||||
@@ -204,6 +246,8 @@ class InferenceEngine:
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
frequency_penalty=frequency_penalty,
|
||||
rep_window=rep_window,
|
||||
stream_callback=cb,
|
||||
)
|
||||
task_ids.append(task_id)
|
||||
@@ -224,9 +268,17 @@ class InferenceEngine:
|
||||
temperature: float,
|
||||
top_p: float,
|
||||
top_k: int,
|
||||
frequency_penalty: float,
|
||||
rep_window: int,
|
||||
) -> Generator:
|
||||
result, task_ids = self._submit_tasks(
|
||||
prompts, max_tokens, temperature, top_p, top_k
|
||||
prompts,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
n = len(prompts)
|
||||
remaining = n
|
||||
@@ -260,9 +312,17 @@ class InferenceEngine:
|
||||
temperature: float,
|
||||
top_p: float,
|
||||
top_k: int,
|
||||
frequency_penalty: float,
|
||||
rep_window: int,
|
||||
) -> Union[str, List[str]]:
|
||||
result, task_ids = self._submit_tasks(
|
||||
prompts, max_tokens, temperature, top_p, top_k
|
||||
prompts,
|
||||
max_tokens,
|
||||
temperature,
|
||||
top_p,
|
||||
top_k,
|
||||
frequency_penalty,
|
||||
rep_window,
|
||||
)
|
||||
|
||||
try:
|
||||
|
||||
+152
-16
@@ -1,15 +1,15 @@
|
||||
"""Composable sampling strategies for logit transformation.
|
||||
|
||||
Implements the Strategy pattern: each sampling technique
|
||||
(temperature, top-k, top-p) is a pluggable strategy that
|
||||
can be composed into a pipeline.
|
||||
(temperature, top-k, top-p, frequency penalty) is a pluggable
|
||||
strategy that can be composed into a pipeline.
|
||||
|
||||
All strategies accept both scalar and per-sample tensor
|
||||
parameters, so a single pipeline works for any batch size.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List, Union
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
@@ -19,16 +19,28 @@ class BaseSamplingStrategy(ABC):
|
||||
"""Abstract base for a logit transformation strategy."""
|
||||
|
||||
@abstractmethod
|
||||
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
"""Applies the strategy to logits.
|
||||
|
||||
Args:
|
||||
logits: Raw logits tensor (batch, vocab_size).
|
||||
filter_value: Value assigned to filtered-out positions.
|
||||
input_ids: Previously generated token IDs ``[batch, seq_len]``,
|
||||
padded with 0. Used by frequency penalty.
|
||||
input_mask: Boolean mask ``[batch, seq_len]``, True for real
|
||||
tokens, False for padding. Used to exclude padding from
|
||||
penalty computation.
|
||||
|
||||
Returns:
|
||||
Transformed logits tensor.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class TemperatureStrategy(BaseSamplingStrategy):
|
||||
@@ -41,13 +53,21 @@ class TemperatureStrategy(BaseSamplingStrategy):
|
||||
def __init__(self, temperature: Union[float, Tensor] = 1.0):
|
||||
self.temperature = temperature
|
||||
|
||||
def apply(self, logits, filter_value=-float("inf")):
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
t = self.temperature
|
||||
if isinstance(t, Tensor):
|
||||
t = t.to(logits.device, non_blocking=True).view(-1, 1)
|
||||
t = torch.clamp(t, min=1e-8)
|
||||
if (t != 1.0).any():
|
||||
logits = logits / t.to(logits.device, non_blocking=True).view(-1, 1)
|
||||
logits = logits / t
|
||||
elif t != 1.0:
|
||||
logits = logits / t
|
||||
logits = logits / max(t, 1e-8)
|
||||
return logits
|
||||
|
||||
|
||||
@@ -61,7 +81,13 @@ class TopKStrategy(BaseSamplingStrategy):
|
||||
def __init__(self, top_k: Union[int, Tensor] = 0):
|
||||
self.top_k = top_k
|
||||
|
||||
def apply(self, logits, filter_value=-float("inf")):
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
tk = self.top_k
|
||||
if isinstance(tk, Tensor):
|
||||
tk = tk.to(logits.device, non_blocking=True).long().clamp(min=0)
|
||||
@@ -98,7 +124,9 @@ class TopPStrategy(BaseSamplingStrategy):
|
||||
def __init__(self, top_p: Union[float, Tensor] = 1.0):
|
||||
self.top_p = top_p
|
||||
|
||||
def _apply(self, logits, top_p, filter_value):
|
||||
def _apply(
|
||||
self, logits: Tensor, top_p: Union[float, Tensor], filter_value: float
|
||||
) -> Tensor:
|
||||
sorted_logits, sorted_indices = torch.sort(logits, descending=True, dim=-1)
|
||||
cum_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
|
||||
remove = cum_probs > top_p
|
||||
@@ -109,7 +137,13 @@ class TopPStrategy(BaseSamplingStrategy):
|
||||
logits[mask] = filter_value
|
||||
return logits
|
||||
|
||||
def apply(self, logits, filter_value=-float("inf")):
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
tp = self.top_p
|
||||
if isinstance(tp, Tensor):
|
||||
tp = tp.to(logits.device, non_blocking=True)
|
||||
@@ -120,6 +154,84 @@ class TopPStrategy(BaseSamplingStrategy):
|
||||
return logits
|
||||
|
||||
|
||||
class FrequencyPenaltyStrategy(BaseSamplingStrategy):
|
||||
"""Penalizes tokens based on how many times they appeared in history.
|
||||
|
||||
Subtracts ``penalty * count(token)`` from each token's logit, where
|
||||
``count(token)`` is the number of occurrences in the generation history
|
||||
(prompt + output). A penalty of ``0.0`` disables the strategy.
|
||||
|
||||
Unlike repetition penalty (which only checks *presence*), frequency
|
||||
penalty scales linearly with occurrence count: the first use is
|
||||
penalized once, the third use three times. This allows natural
|
||||
repetition of common words while suppressing degenerate loops.
|
||||
|
||||
Reference: OpenAI API ``frequency_penalty`` parameter.
|
||||
|
||||
Args:
|
||||
penalty: Scalar or ``[batch]`` tensor (0.0 disables, range -2.0~2.0).
|
||||
"""
|
||||
|
||||
def __init__(self, penalty: Union[float, Tensor] = 0.0):
|
||||
self.penalty = penalty
|
||||
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
if input_ids is None:
|
||||
return logits
|
||||
|
||||
p = self.penalty
|
||||
if isinstance(p, Tensor):
|
||||
p = p.to(logits.device, non_blocking=True).view(-1, 1)
|
||||
if (p == 0.0).all():
|
||||
return logits
|
||||
elif p == 0.0:
|
||||
return logits
|
||||
|
||||
input_ids = input_ids.to(logits.device, non_blocking=True)
|
||||
|
||||
if input_mask is not None:
|
||||
input_mask = input_mask.to(logits.device, non_blocking=True)
|
||||
masked_ids = input_ids.clone()
|
||||
masked_ids[~input_mask] = -1
|
||||
else:
|
||||
masked_ids = input_ids
|
||||
|
||||
batch_sz, seq_len = masked_ids.shape
|
||||
vocab_size = logits.size(-1)
|
||||
|
||||
if isinstance(p, Tensor):
|
||||
penalty_per_row = p.expand(batch_sz, 1)
|
||||
else:
|
||||
penalty_per_row = torch.full(
|
||||
(batch_sz, 1), float(p), device=logits.device, dtype=logits.dtype
|
||||
)
|
||||
|
||||
counts = torch.zeros(
|
||||
batch_sz, vocab_size, device=logits.device, dtype=logits.dtype
|
||||
)
|
||||
valid_mask = masked_ids >= 0
|
||||
if valid_mask.any():
|
||||
valid_ids = masked_ids[valid_mask]
|
||||
row_indices = (
|
||||
torch.arange(batch_sz, device=logits.device)
|
||||
.unsqueeze(1)
|
||||
.expand_as(masked_ids)[valid_mask]
|
||||
)
|
||||
counts.index_put_(
|
||||
(row_indices, valid_ids),
|
||||
torch.ones_like(valid_ids, dtype=logits.dtype),
|
||||
accumulate=True,
|
||||
)
|
||||
|
||||
return logits - penalty_per_row * counts
|
||||
|
||||
|
||||
class SamplingPipeline(BaseSamplingStrategy):
|
||||
"""Composes multiple sampling strategies into a single transformation.
|
||||
|
||||
@@ -140,23 +252,39 @@ class SamplingPipeline(BaseSamplingStrategy):
|
||||
def __init__(self, strategies: List[BaseSamplingStrategy]):
|
||||
self.strategies = strategies
|
||||
|
||||
def apply(self, logits, filter_value=-float("inf")):
|
||||
def apply(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
for strategy in self.strategies:
|
||||
logits = strategy.apply(logits, filter_value)
|
||||
logits = strategy.apply(logits, filter_value, input_ids, input_mask)
|
||||
return logits
|
||||
|
||||
@torch.no_grad()
|
||||
def sample(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
|
||||
@torch.inference_mode()
|
||||
def sample(
|
||||
self,
|
||||
logits: Tensor,
|
||||
filter_value: float = -float("inf"),
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
"""Apply strategies then sample (softmax + multinomial).
|
||||
|
||||
Args:
|
||||
logits: Raw logits ``[batch, vocab_size]``.
|
||||
input_ids: Previously generated token IDs ``[batch, seq_len]``.
|
||||
input_mask: Boolean mask for ``input_ids`` padding.
|
||||
|
||||
Returns:
|
||||
Sampled token IDs ``[batch]``.
|
||||
"""
|
||||
return torch.multinomial(
|
||||
torch.softmax(self.apply(logits, filter_value), dim=-1),
|
||||
torch.softmax(
|
||||
self.apply(logits, filter_value, input_ids, input_mask), dim=-1
|
||||
),
|
||||
num_samples=1,
|
||||
).squeeze(-1)
|
||||
|
||||
@@ -167,6 +295,9 @@ def sample(
|
||||
temperature: Union[float, Tensor] = 1.0,
|
||||
top_k: Union[int, Tensor] = 0,
|
||||
top_p: Union[float, Tensor] = 1.0,
|
||||
frequency_penalty: Union[float, Tensor] = 0.0,
|
||||
input_ids: Optional[Tensor] = None,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
filter_value: float = -float("inf"),
|
||||
) -> Tensor:
|
||||
"""Apply sampling strategies then sample (softmax + multinomial).
|
||||
@@ -175,6 +306,10 @@ def sample(
|
||||
|
||||
Args:
|
||||
logits: Raw logits ``[batch, vocab_size]``.
|
||||
frequency_penalty: Penalty per occurrence for repeated tokens
|
||||
(0.0 disables, range -2.0~2.0).
|
||||
input_ids: Previously generated token IDs ``[batch, seq_len]``.
|
||||
input_mask: Boolean mask for ``input_ids`` padding.
|
||||
|
||||
Returns:
|
||||
Sampled token IDs ``[batch]``.
|
||||
@@ -184,5 +319,6 @@ def sample(
|
||||
TemperatureStrategy(temperature),
|
||||
TopKStrategy(top_k),
|
||||
TopPStrategy(top_p),
|
||||
FrequencyPenaltyStrategy(frequency_penalty),
|
||||
]
|
||||
).sample(logits, filter_value)
|
||||
).sample(logits, filter_value, input_ids, input_mask)
|
||||
|
||||
@@ -6,7 +6,7 @@ import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.inference.core.cache import KvcacheView
|
||||
from astrai.inference.core.cache import CacheView
|
||||
from astrai.model.components.linear import Linear
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
from astrai.model.components.rope import apply_rotary_emb
|
||||
@@ -24,9 +24,7 @@ def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
||||
|
||||
|
||||
class AttnFactory(BaseFactory[nn.Module]):
|
||||
@classmethod
|
||||
def create(cls, attn_type: str, **kwargs) -> nn.Module:
|
||||
return super().create(attn_type, **kwargs)
|
||||
pass
|
||||
|
||||
|
||||
@AttnFactory.register("gqa")
|
||||
@@ -40,6 +38,7 @@ class GQA(nn.Module):
|
||||
norm_eps: float,
|
||||
use_gated_attention: bool,
|
||||
layer_id: int,
|
||||
n_layers: int = 1,
|
||||
):
|
||||
super().__init__()
|
||||
assert dim % n_heads == 0
|
||||
@@ -57,7 +56,7 @@ class GQA(nn.Module):
|
||||
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)
|
||||
self.o_proj = Linear(dim, dim, init_std=0.02 / (2 * n_layers) ** 0.5)
|
||||
|
||||
if self.use_qk_norm:
|
||||
self.q_norm = RMSNorm(self.head_dim, norm_eps)
|
||||
@@ -76,7 +75,7 @@ class GQA(nn.Module):
|
||||
x: Tensor,
|
||||
rotary_emb: Tensor,
|
||||
attn_mask: Tensor = None,
|
||||
paged_cache: Optional[KvcacheView] = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
) -> Tensor:
|
||||
is_causal = attn_mask is None
|
||||
|
||||
@@ -123,6 +122,7 @@ class MLA(nn.Module):
|
||||
use_qk_norm: bool,
|
||||
use_gated_attention: bool,
|
||||
layer_id: int,
|
||||
n_layers: int = 1,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
@@ -150,7 +150,9 @@ class MLA(nn.Module):
|
||||
n_kv_heads * (2 * self.head_dim),
|
||||
)
|
||||
|
||||
self.o_proj = Linear(dim, dim, bias=False)
|
||||
self.o_proj = Linear(
|
||||
dim, dim, bias=False, init_std=0.02 / (2 * n_layers) ** 0.5
|
||||
)
|
||||
|
||||
if use_gated_attention:
|
||||
self.gate = Linear(dim, dim, bias=False)
|
||||
@@ -160,7 +162,7 @@ class MLA(nn.Module):
|
||||
x: Tensor,
|
||||
rotary_emb: Tensor,
|
||||
attn_mask: Tensor = None,
|
||||
paged_cache: Optional[KvcacheView] = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
) -> Tensor:
|
||||
bsz, seq_len, _ = x.size()
|
||||
is_causal = attn_mask is None
|
||||
|
||||
@@ -1,51 +1,31 @@
|
||||
from dataclasses import asdict
|
||||
from typing import Optional
|
||||
|
||||
import torch.nn as nn
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.inference.core.cache import KvcacheView
|
||||
from astrai.inference.core.cache import CacheView
|
||||
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,
|
||||
):
|
||||
def __init__(self, config, layer_id: int):
|
||||
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)
|
||||
cfg = asdict(config)
|
||||
cfg["down_init_std"] = 0.02 / (2 * config.n_layers) ** 0.5
|
||||
self.attention = AttnFactory.create(config.attn_type, **cfg, layer_id=layer_id)
|
||||
self.input_norm = RMSNorm(config.dim, config.norm_eps)
|
||||
self.post_attention_norm = RMSNorm(config.dim, config.norm_eps)
|
||||
self.mlp = FFNFactory.create(config.ffn_type, **cfg)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
rotary_emb: Tensor,
|
||||
attention_mask: Optional[Tensor] = None,
|
||||
paged_cache: Optional[KvcacheView] = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
) -> Tensor:
|
||||
attn_output = self.attention(
|
||||
self.input_norm(x),
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
@@ -5,12 +7,20 @@ from torch import Tensor
|
||||
|
||||
|
||||
class Embedding(nn.Module):
|
||||
def __init__(self, vocab_size: int, embedding_dim: int):
|
||||
def __init__(self, vocab_size: int, embedding_dim: int, neftune_alpha: float = 0.0):
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.empty((vocab_size, embedding_dim)))
|
||||
self.neftune_noise_alpha = neftune_alpha
|
||||
|
||||
def set_neftune_alpha(self, alpha: float):
|
||||
self.neftune_noise_alpha = alpha
|
||||
|
||||
def reset_parameters(self):
|
||||
nn.init.normal_(self.weight, mean=0.0, std=0.02)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return F.embedding(x, self.weight)
|
||||
out = F.embedding(x, self.weight)
|
||||
if self.training and self.neftune_noise_alpha > 0.0:
|
||||
eps = self.neftune_noise_alpha / math.sqrt(out.size(1))
|
||||
out = out + eps * torch.randn_like(out)
|
||||
return out
|
||||
|
||||
@@ -5,13 +5,16 @@ from torch import Tensor
|
||||
|
||||
|
||||
class Linear(nn.Module):
|
||||
def __init__(self, in_dim: int, out_dim: int, bias: bool = False):
|
||||
def __init__(
|
||||
self, in_dim: int, out_dim: int, bias: bool = False, init_std: float = 0.02
|
||||
):
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.empty((out_dim, in_dim)))
|
||||
self.bias = nn.Parameter(torch.zeros(out_dim)) if bias else None
|
||||
self.init_std = init_std
|
||||
|
||||
def reset_parameters(self):
|
||||
nn.init.kaiming_uniform_(self.weight, a=5**0.5)
|
||||
nn.init.normal_(self.weight, mean=0.0, std=self.init_std)
|
||||
if self.bias is not None:
|
||||
fan_in, _ = nn.init._calculate_fan_in_and_fan_out(self.weight)
|
||||
bound = 1 / (fan_in**0.5)
|
||||
|
||||
@@ -8,18 +8,16 @@ 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)
|
||||
pass
|
||||
|
||||
|
||||
@FFNFactory.register("mlp")
|
||||
class MLP(nn.Module):
|
||||
def __init__(self, dim: int, dim_ffn: int):
|
||||
def __init__(self, dim: int, dim_ffn: int, down_init_std: float = 0.02):
|
||||
super().__init__()
|
||||
self.up = Linear(dim, dim_ffn)
|
||||
self.gate = Linear(dim, dim_ffn)
|
||||
self.down = Linear(dim_ffn, dim)
|
||||
self.down = Linear(dim_ffn, dim, init_std=down_init_std)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
gated = self.up(x) * F.silu(self.gate(x))
|
||||
@@ -37,6 +35,7 @@ class DeepSeekMoE(nn.Module):
|
||||
n_shared_experts: int = 1,
|
||||
n_activated_experts: int = 2,
|
||||
topk_method: str = "greedy",
|
||||
n_layers: int = 1,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
@@ -46,12 +45,20 @@ class DeepSeekMoE(nn.Module):
|
||||
self.topk_method = topk_method
|
||||
|
||||
self.router = Linear(dim, n_routed_experts, bias=False)
|
||||
moe_scale = 1 / max(n_shared_experts, 1) + 1 / n_activated_experts
|
||||
down_init_std = 0.02 / (2 * n_layers * moe_scale) ** 0.5
|
||||
|
||||
self.shared_experts = nn.ModuleList(
|
||||
[MLP(dim, dim_ffn) for _ in range(n_shared_experts)]
|
||||
[
|
||||
MLP(dim, dim_ffn, down_init_std=down_init_std)
|
||||
for _ in range(n_shared_experts)
|
||||
]
|
||||
)
|
||||
self.routed_experts = nn.ModuleList(
|
||||
[MLP(dim, dim_ffn) for _ in range(n_routed_experts)]
|
||||
[
|
||||
MLP(dim, dim_ffn, down_init_std=down_init_std)
|
||||
for _ in range(n_routed_experts)
|
||||
]
|
||||
)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
|
||||
+4
-14
@@ -23,22 +23,12 @@ class EmbeddingEncoder(AutoModel):
|
||||
self.rotary_embedding = RotaryEmbedding(
|
||||
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling
|
||||
)
|
||||
self.embed_tokens = Embedding(config.vocab_size, config.dim)
|
||||
self.embed_tokens = Embedding(
|
||||
config.vocab_size, config.dim, neftune_alpha=config.neftune_alpha
|
||||
)
|
||||
|
||||
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)
|
||||
]
|
||||
[DecoderBlock(config, layer_id) for layer_id in range(config.n_layers)]
|
||||
)
|
||||
|
||||
self.norm = RMSNorm(config.dim, config.norm_eps)
|
||||
|
||||
+12
-34
@@ -5,7 +5,7 @@ import torch.nn as nn
|
||||
from torch import Tensor
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.inference.core.cache import KvcacheView
|
||||
from astrai.inference.core.cache import CacheView
|
||||
from astrai.model.automodel import AutoModel
|
||||
from astrai.model.components.decoder_block import DecoderBlock
|
||||
from astrai.model.components.embedding import Embedding
|
||||
@@ -26,24 +26,21 @@ def process_attention_mask(
|
||||
return input_mask
|
||||
|
||||
device = input_tensor.device
|
||||
dtype = input_tensor.dtype
|
||||
B, S = input_tensor.size()[:2]
|
||||
B = input_tensor.size(0)
|
||||
T = position_ids.max().item() + 1
|
||||
|
||||
if input_mask is None:
|
||||
if position_ids.min().item() == 0 and is_causal:
|
||||
return None
|
||||
pad = torch.ones(B, T, dtype=torch.bool, device=device)
|
||||
attend = torch.ones(B, 1, T, dtype=torch.bool, device=device)
|
||||
else:
|
||||
pad = input_mask[:, :T].to(device=device, dtype=torch.bool)
|
||||
attend = input_mask[:, :T].to(device=device, dtype=torch.bool).unsqueeze(1)
|
||||
|
||||
attend = pad.view(B, 1, T).expand(B, S, T).clone()
|
||||
if is_causal:
|
||||
attend &= position_ids.unsqueeze(-1) >= torch.arange(T, device=device)
|
||||
causal = position_ids.unsqueeze(-1) >= torch.arange(T, device=device)
|
||||
attend = attend & causal
|
||||
|
||||
return torch.full(
|
||||
(B, 1, S, T), -torch.finfo(dtype).max / 2, dtype=dtype, device=device
|
||||
).masked_fill_(attend.unsqueeze(1), 0.0)
|
||||
return attend.unsqueeze(1)
|
||||
|
||||
|
||||
@AutoModel.register("autoregressive_lm")
|
||||
@@ -62,31 +59,12 @@ class AutoRegressiveLM(AutoModel):
|
||||
self.rotary_embedding = RotaryEmbedding(
|
||||
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling
|
||||
)
|
||||
self.embed_tokens = Embedding(config.vocab_size, config.dim)
|
||||
self.embed_tokens = Embedding(
|
||||
config.vocab_size, config.dim, neftune_alpha=config.neftune_alpha
|
||||
)
|
||||
|
||||
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,
|
||||
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)
|
||||
]
|
||||
[DecoderBlock(config, layer_id) for layer_id in range(config.n_layers)]
|
||||
)
|
||||
|
||||
self.norm = RMSNorm(config.dim, config.norm_eps)
|
||||
@@ -134,7 +112,7 @@ class AutoRegressiveLM(AutoModel):
|
||||
self,
|
||||
input_ids: Tensor,
|
||||
input_mask: Optional[Tensor] = None,
|
||||
paged_cache: Optional[KvcacheView] = None,
|
||||
paged_cache: Optional[CacheView] = None,
|
||||
position_ids: Optional[Tensor] = None,
|
||||
) -> Dict[str, Tensor]:
|
||||
assert input_ids.ndim == 2
|
||||
|
||||
+58
-14
@@ -2,11 +2,14 @@
|
||||
|
||||
import contextlib
|
||||
import logging
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
from torch.distributed.fsdp import FullStateDictConfig, StateDictType
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||
from torch.optim import Optimizer
|
||||
@@ -115,8 +118,23 @@ class BaseExecutor:
|
||||
def backward(self, loss: torch.Tensor):
|
||||
loss.backward()
|
||||
|
||||
def unwrap_model(self, model: nn.Module) -> nn.Module:
|
||||
return model
|
||||
def unwrap_model(self, model: nn.Module):
|
||||
return model.state_dict()
|
||||
|
||||
@contextmanager
|
||||
def checkpoint_context(self, model: nn.Module):
|
||||
if self.use_distributed:
|
||||
dist.barrier()
|
||||
state_dict = self._gather_state_dict(model)
|
||||
yield state_dict
|
||||
if self.use_distributed:
|
||||
dist.barrier()
|
||||
|
||||
def _gather_state_dict(self, model: nn.Module):
|
||||
state_dict = self.unwrap_model(model)
|
||||
if self.use_distributed and get_rank() != 0:
|
||||
return None
|
||||
return state_dict
|
||||
|
||||
@property
|
||||
def use_distributed(self) -> bool:
|
||||
@@ -130,6 +148,19 @@ class BaseExecutor:
|
||||
def grad_accum_steps(self) -> int:
|
||||
return self.gradient_state.num_steps
|
||||
|
||||
def clip_grad_norm(self, model: nn.Module, max_norm: Optional[float]) -> float:
|
||||
if max_norm is None:
|
||||
total_norm = torch.norm(
|
||||
torch.stack(
|
||||
[p.grad.norm(2) for p in model.parameters() if p.grad is not None]
|
||||
)
|
||||
)
|
||||
return total_norm.item()
|
||||
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)
|
||||
if isinstance(total_norm, torch.Tensor):
|
||||
return total_norm.item()
|
||||
return total_norm
|
||||
|
||||
|
||||
class ExecutorFactory(BaseFactory[BaseExecutor]):
|
||||
pass
|
||||
@@ -180,7 +211,7 @@ class DDPExecutor(BaseExecutor):
|
||||
if not self.use_distributed:
|
||||
logger.warning("DDP backend selected but world_size=1, model not wrapped")
|
||||
return model
|
||||
local_rank = get_rank()
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", get_rank()))
|
||||
model = DDP(
|
||||
model,
|
||||
device_ids=[local_rank],
|
||||
@@ -195,10 +226,10 @@ class DDPExecutor(BaseExecutor):
|
||||
return model.no_sync()
|
||||
return contextlib.nullcontext()
|
||||
|
||||
def unwrap_model(self, model: nn.Module) -> nn.Module:
|
||||
def unwrap_model(self, model: nn.Module):
|
||||
if isinstance(model, DDP):
|
||||
return model.module
|
||||
return model
|
||||
return model.module.state_dict()
|
||||
return model.state_dict()
|
||||
|
||||
|
||||
@ExecutorFactory.register("fsdp")
|
||||
@@ -217,7 +248,6 @@ class FSDPExecutor(BaseExecutor):
|
||||
sync_module_states: bool = False,
|
||||
forward_prefetch: bool = False,
|
||||
limit_all_gathers: bool = True,
|
||||
use_orig_params: bool = False,
|
||||
ignored_states=None,
|
||||
device_mesh=None,
|
||||
):
|
||||
@@ -236,7 +266,7 @@ class FSDPExecutor(BaseExecutor):
|
||||
sync_module_states=sync_module_states,
|
||||
forward_prefetch=forward_prefetch,
|
||||
limit_all_gathers=limit_all_gathers,
|
||||
use_orig_params=use_orig_params,
|
||||
use_orig_params=True,
|
||||
ignored_states=ignored_states,
|
||||
device_mesh=device_mesh,
|
||||
).items()
|
||||
@@ -259,9 +289,23 @@ class FSDPExecutor(BaseExecutor):
|
||||
return model.no_sync()
|
||||
return contextlib.nullcontext()
|
||||
|
||||
def unwrap_model(self, model: nn.Module) -> nn.Module:
|
||||
if self._original_model is not None:
|
||||
return self._original_model
|
||||
if isinstance(model, FSDP):
|
||||
return model._fsdp_wrapped_module
|
||||
return model
|
||||
def clip_grad_norm(self, model: nn.Module, max_norm: Optional[float]) -> float:
|
||||
if max_norm is None:
|
||||
return super().clip_grad_norm(model, max_norm)
|
||||
if isinstance(model, FSDP) and self.use_distributed:
|
||||
total_norm = model.clip_grad_norm_(max_norm)
|
||||
if isinstance(total_norm, torch.Tensor):
|
||||
return total_norm.item()
|
||||
return total_norm
|
||||
return super().clip_grad_norm(model, max_norm)
|
||||
|
||||
def unwrap_model(self, model: nn.Module):
|
||||
if isinstance(model, FSDP) and self.use_distributed:
|
||||
with FSDP.state_dict_type(
|
||||
model,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
):
|
||||
return model.state_dict()
|
||||
|
||||
return model.state_dict()
|
||||
|
||||
+131
-54
@@ -1,13 +1,21 @@
|
||||
import os
|
||||
import socket
|
||||
from abc import ABC, abstractmethod
|
||||
from contextlib import contextmanager
|
||||
from functools import wraps
|
||||
from typing import Callable
|
||||
from typing import Callable, Optional
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.multiprocessing as mp
|
||||
|
||||
|
||||
def find_free_port() -> str:
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
s.bind(("", 0))
|
||||
return str(s.getsockname()[1])
|
||||
|
||||
|
||||
def get_current_device():
|
||||
return os.environ["LOCAL_DEVICE"]
|
||||
|
||||
@@ -30,6 +38,7 @@ def get_rank() -> int:
|
||||
def setup_parallel(
|
||||
rank: int,
|
||||
world_size: int,
|
||||
local_rank: int,
|
||||
backend: str = "nccl",
|
||||
master_addr: str = "localhost",
|
||||
master_port: str = "29500",
|
||||
@@ -41,20 +50,26 @@ def setup_parallel(
|
||||
return
|
||||
|
||||
if world_size <= 1:
|
||||
device_id = torch.device(device_type, local_rank)
|
||||
os.environ["LOCAL_RANK"] = str(local_rank)
|
||||
os.environ["WORLD_SIZE"] = "1"
|
||||
os.environ["LOCAL_DEVICE"] = str(device_id)
|
||||
yield None
|
||||
return
|
||||
|
||||
device_id = torch.device(device_type, rank)
|
||||
device_id = torch.device(device_type, local_rank)
|
||||
|
||||
os.environ["MASTER_ADDR"] = master_addr
|
||||
os.environ["MASTER_PORT"] = master_port
|
||||
os.environ["LOCAL_RANK"] = str(rank)
|
||||
os.environ["LOCAL_RANK"] = str(local_rank)
|
||||
os.environ["WORLD_SIZE"] = str(world_size)
|
||||
os.environ["LOCAL_DEVICE"] = str(device_id)
|
||||
|
||||
dist.init_process_group(
|
||||
rank=rank, world_size=world_size, backend=backend, device_id=device_id
|
||||
)
|
||||
pg_kwargs = dict(rank=rank, world_size=world_size, backend=backend)
|
||||
if backend in ("nccl", "ccl"):
|
||||
pg_kwargs["device_id"] = device_id
|
||||
|
||||
dist.init_process_group(**pg_kwargs)
|
||||
|
||||
try:
|
||||
if backend == "nccl" and torch.cuda.is_available():
|
||||
@@ -90,7 +105,7 @@ def only_on_rank(rank, sync=False):
|
||||
return decorator
|
||||
|
||||
|
||||
def wrapper_spawn_func(
|
||||
def _run_single_rank(
|
||||
rank: int,
|
||||
world_size: int,
|
||||
backend: str,
|
||||
@@ -100,20 +115,108 @@ def wrapper_spawn_func(
|
||||
func: Callable,
|
||||
kwargs: dict,
|
||||
):
|
||||
try:
|
||||
with setup_parallel(
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
local_rank=rank,
|
||||
backend=backend,
|
||||
master_addr=master_addr,
|
||||
master_port=master_port,
|
||||
device_type=device_type,
|
||||
):
|
||||
func(**kwargs)
|
||||
|
||||
|
||||
class LaunchStrategy(ABC):
|
||||
"""Strategy for launching a function in a distributed context."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
world_size: int,
|
||||
backend: str,
|
||||
master_addr: str,
|
||||
master_port: str,
|
||||
device_type: str,
|
||||
start_method: str,
|
||||
):
|
||||
self.world_size = world_size
|
||||
self.backend = backend
|
||||
self.master_addr = master_addr
|
||||
self.master_port = master_port
|
||||
self.device_type = device_type
|
||||
self.start_method = start_method
|
||||
|
||||
@abstractmethod
|
||||
def launch(self, func: Callable, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class TorchrunStrategy(LaunchStrategy):
|
||||
"""External orchestrator (torchrun, SLURM, K8s) — env vars pre-set."""
|
||||
|
||||
def launch(self, func: Callable, **kwargs):
|
||||
rank = int(os.environ["RANK"])
|
||||
world_size = int(os.environ["WORLD_SIZE"])
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", rank))
|
||||
with setup_parallel(
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
backend=backend,
|
||||
master_addr=master_addr,
|
||||
master_port=master_port,
|
||||
device_type=device_type,
|
||||
local_rank=local_rank,
|
||||
backend=self.backend,
|
||||
master_addr=os.environ.get("MASTER_ADDR", self.master_addr),
|
||||
master_port=os.environ.get("MASTER_PORT", self.master_port),
|
||||
device_type=self.device_type,
|
||||
):
|
||||
func(**kwargs)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error in rank {rank}: {e}")
|
||||
raise
|
||||
|
||||
class LocalStrategy(LaunchStrategy):
|
||||
"""Local launcher — single-process or mp.start_processes."""
|
||||
|
||||
def launch(self, func: Callable, **kwargs):
|
||||
args = (
|
||||
self.world_size,
|
||||
self.backend,
|
||||
self.master_addr,
|
||||
self.master_port,
|
||||
self.device_type,
|
||||
func,
|
||||
kwargs,
|
||||
)
|
||||
|
||||
if self.world_size == 1:
|
||||
_run_single_rank(0, *args)
|
||||
return
|
||||
|
||||
ctx = mp.start_processes(
|
||||
_run_single_rank,
|
||||
args=args,
|
||||
nprocs=self.world_size,
|
||||
start_method=self.start_method,
|
||||
join=False,
|
||||
)
|
||||
try:
|
||||
while not ctx.join():
|
||||
pass
|
||||
except BaseException:
|
||||
for p in ctx.processes:
|
||||
p.terminate()
|
||||
ctx.join()
|
||||
raise
|
||||
|
||||
|
||||
def _detect_launcher() -> str:
|
||||
"""Detect the distributed launcher from environment.
|
||||
|
||||
Returns one of: "torchelastic", "torchrun", "external", "local".
|
||||
"""
|
||||
if dist.is_torchelastic_launched():
|
||||
return "torchelastic"
|
||||
if "LOCAL_WORLD_SIZE" in os.environ:
|
||||
return "torchrun"
|
||||
if "RANK" in os.environ and "WORLD_SIZE" in os.environ:
|
||||
return "external"
|
||||
return "local"
|
||||
|
||||
|
||||
def spawn_parallel_fn(
|
||||
@@ -121,46 +224,20 @@ def spawn_parallel_fn(
|
||||
world_size: int,
|
||||
backend: str = "nccl",
|
||||
master_addr: str = "localhost",
|
||||
master_port: str = "29500",
|
||||
master_port: Optional[str] = None,
|
||||
device_type: str = "cuda",
|
||||
start_method: str = "spawn",
|
||||
**kwargs,
|
||||
):
|
||||
# clear environment variables
|
||||
for key in [
|
||||
"MASTER_ADDR",
|
||||
"MASTER_PORT",
|
||||
"RANK",
|
||||
"WORLD_SIZE",
|
||||
"LOCAL_RANK",
|
||||
"LOCAL_DEVICE",
|
||||
]:
|
||||
if key in os.environ:
|
||||
del os.environ[key]
|
||||
|
||||
if world_size == 1:
|
||||
device_id = torch.device(device_type, 0)
|
||||
os.environ["LOCAL_RANK"] = "0"
|
||||
os.environ["WORLD_SIZE"] = "1"
|
||||
os.environ["LOCAL_DEVICE"] = str(device_id)
|
||||
|
||||
func(**kwargs)
|
||||
return
|
||||
|
||||
wrapper_spawn_func_args = (
|
||||
world_size,
|
||||
backend,
|
||||
master_addr,
|
||||
master_port,
|
||||
device_type,
|
||||
func,
|
||||
kwargs,
|
||||
)
|
||||
|
||||
mp.start_processes(
|
||||
wrapper_spawn_func,
|
||||
args=wrapper_spawn_func_args,
|
||||
nprocs=world_size,
|
||||
start_method=start_method,
|
||||
join=True,
|
||||
)
|
||||
if master_port is None:
|
||||
master_port = find_free_port()
|
||||
launcher = _detect_launcher()
|
||||
if launcher in ("torchelastic", "torchrun", "external"):
|
||||
strategy = TorchrunStrategy(
|
||||
world_size, backend, master_addr, master_port, device_type, start_method
|
||||
)
|
||||
else:
|
||||
strategy = LocalStrategy(
|
||||
world_size, backend, master_addr, master_port, device_type, start_method
|
||||
)
|
||||
strategy.launch(func, **kwargs)
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
from astrai.preprocessing.builder import (
|
||||
BaseMaskBuilder,
|
||||
MaskBuilderFactory,
|
||||
MultiOutputMaskBuilder,
|
||||
SectionedMaskBuilder,
|
||||
SingleOutputMaskBuilder,
|
||||
)
|
||||
from astrai.preprocessing.packing import (
|
||||
PackingStrategy,
|
||||
PackingStrategyFactory,
|
||||
plan_bfd,
|
||||
)
|
||||
from astrai.preprocessing.pipeline import Pipeline, filter_by_length
|
||||
from astrai.preprocessing.position_id import (
|
||||
PositionIdStrategy,
|
||||
PositionIdStrategyFactory,
|
||||
)
|
||||
from astrai.preprocessing.transform import TokenizeTransform
|
||||
from astrai.preprocessing.writer import (
|
||||
StoreWriter,
|
||||
StoreWriterFactory,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"BaseMaskBuilder",
|
||||
"MaskBuilderFactory",
|
||||
"MultiOutputMaskBuilder",
|
||||
"PackingStrategy",
|
||||
"PackingStrategyFactory",
|
||||
"Pipeline",
|
||||
"PositionIdStrategy",
|
||||
"PositionIdStrategyFactory",
|
||||
"SectionedMaskBuilder",
|
||||
"SingleOutputMaskBuilder",
|
||||
"StoreWriter",
|
||||
"StoreWriterFactory",
|
||||
"TokenizeTransform",
|
||||
"filter_by_length",
|
||||
"plan_bfd",
|
||||
]
|
||||
@@ -0,0 +1,337 @@
|
||||
"""Mask building for preprocessing pipeline.
|
||||
|
||||
:class:`SectionRenderer` converts section specs into token ids and loss
|
||||
masks (template / text / value extraction). :class:`SingleOutputMaskBuilder`
|
||||
handles single-output (SFT / pretrain), :class:`MultiOutputMaskBuilder`
|
||||
handles multi-output (DPO / GRPO), and :class:`SectionedMaskBuilder`
|
||||
orchestrates both modes as a façade.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
|
||||
def _extract_domain(item: dict, domain_key: Optional[str]) -> str:
|
||||
if not domain_key:
|
||||
return "__default__"
|
||||
val = item.get(domain_key, "__default__")
|
||||
return val if isinstance(val, str) else "__default__"
|
||||
|
||||
|
||||
def _resolve_action(action: str, role: str, config) -> str:
|
||||
if action == "$role":
|
||||
return config.mask.get(role, config.mask_default)
|
||||
return action
|
||||
|
||||
|
||||
class SectionRenderer:
|
||||
"""Render section specs into ``(ids, loss_mask)`` tuples."""
|
||||
|
||||
def process_sections(
|
||||
self,
|
||||
item: dict,
|
||||
sections: list,
|
||||
config,
|
||||
tokenizer,
|
||||
*,
|
||||
is_top_level: bool = False,
|
||||
):
|
||||
all_ids: list[int] = []
|
||||
loss_mask: list[int] = []
|
||||
|
||||
has_template = any(s.get("template") for s in sections)
|
||||
is_text_config = not has_template and all(
|
||||
s["action"] == "train" for s in sections
|
||||
)
|
||||
|
||||
if is_top_level and has_template and tokenizer.bos_token_id is not None:
|
||||
all_ids.append(tokenizer.bos_token_id)
|
||||
loss_mask.append(0)
|
||||
|
||||
first_section = True
|
||||
for sec in sections:
|
||||
field = sec["field"]
|
||||
action = sec["action"]
|
||||
use_template = sec.get("template", False)
|
||||
add_special = sec.get(
|
||||
"add_special_tokens", not use_template and first_section
|
||||
)
|
||||
|
||||
if use_template:
|
||||
success = self._append_template(
|
||||
item, field, action, tokenizer, config, all_ids, loss_mask
|
||||
)
|
||||
if not success:
|
||||
continue
|
||||
else:
|
||||
success = self._append_text(
|
||||
item,
|
||||
field,
|
||||
action,
|
||||
tokenizer,
|
||||
add_special,
|
||||
is_text_config,
|
||||
config,
|
||||
all_ids,
|
||||
loss_mask,
|
||||
)
|
||||
if not success:
|
||||
continue
|
||||
|
||||
first_section = False
|
||||
|
||||
max_len = config.preprocessing.max_seq_len
|
||||
all_ids = all_ids[:max_len]
|
||||
loss_mask = loss_mask[: len(all_ids)]
|
||||
|
||||
if not all_ids:
|
||||
return None, None
|
||||
|
||||
if is_top_level and has_template and len(all_ids) <= 1:
|
||||
return None, None
|
||||
|
||||
return all_ids, loss_mask
|
||||
|
||||
def process_list_field(self, item: dict, sections: list, config, tokenizer):
|
||||
"""Tokenize a list-valued field, preserving per-element boundaries.
|
||||
|
||||
Returns ``(list_of_id_lists, list_of_mask_lists)`` where each
|
||||
inner list corresponds to one element of the source list. This
|
||||
is critical for GRPO where each response must stay a separate
|
||||
sequence so the strategy can form a ``[G, R]`` tensor.
|
||||
"""
|
||||
per_item_ids: list[list[int]] = []
|
||||
per_item_masks: list[list[int]] = []
|
||||
|
||||
for sec in sections:
|
||||
field = sec["field"]
|
||||
action = sec["action"]
|
||||
use_template = sec.get("template", False)
|
||||
|
||||
values = item.get(field)
|
||||
if not isinstance(values, list):
|
||||
continue
|
||||
|
||||
for val in values:
|
||||
ids: list[int] = []
|
||||
mask: list[int] = []
|
||||
if use_template:
|
||||
if isinstance(val, list):
|
||||
wrapper = {field: val}
|
||||
self._append_template(
|
||||
wrapper, field, action, tokenizer, config, ids, mask
|
||||
)
|
||||
else:
|
||||
wrapper = {field: str(val)}
|
||||
self._append_text(
|
||||
wrapper,
|
||||
field,
|
||||
action,
|
||||
tokenizer,
|
||||
False,
|
||||
False,
|
||||
config,
|
||||
ids,
|
||||
mask,
|
||||
)
|
||||
if ids:
|
||||
max_len = config.preprocessing.max_seq_len
|
||||
ids = ids[:max_len]
|
||||
mask = mask[: len(ids)]
|
||||
per_item_ids.append(ids)
|
||||
per_item_masks.append(mask)
|
||||
|
||||
if not per_item_ids:
|
||||
return None, None
|
||||
return per_item_ids, per_item_masks
|
||||
|
||||
@staticmethod
|
||||
def is_value_section(sections: list) -> bool:
|
||||
return len(sections) == 1 and sections[0].get("action") == "value"
|
||||
|
||||
@staticmethod
|
||||
def extract_raw_value(item: dict, sections: list):
|
||||
sec = sections[0]
|
||||
field = sec["field"]
|
||||
raw = item.get(field)
|
||||
if raw is None:
|
||||
return None
|
||||
if isinstance(raw, list):
|
||||
return [float(v) for v in raw]
|
||||
return [float(raw)]
|
||||
|
||||
def _append_template(
|
||||
self, item, field, action, tokenizer, config, all_ids, loss_mask
|
||||
):
|
||||
messages = item.get(field)
|
||||
if not isinstance(messages, list) or not messages:
|
||||
return False
|
||||
for msg in messages:
|
||||
role = msg.get("role", "")
|
||||
act = _resolve_action(action, role, config)
|
||||
rendered = tokenizer.apply_chat_template(
|
||||
[msg], tokenize=False, add_generation_prompt=False
|
||||
)
|
||||
ids = tokenizer.encode(rendered, add_special_tokens=False)
|
||||
all_ids.extend(ids)
|
||||
val = 1 if act == "train" else 0
|
||||
loss_mask.extend([val] * len(ids))
|
||||
return True
|
||||
|
||||
def _append_text(
|
||||
self,
|
||||
item,
|
||||
field,
|
||||
action,
|
||||
tokenizer,
|
||||
add_special,
|
||||
is_text_config,
|
||||
config,
|
||||
all_ids,
|
||||
loss_mask,
|
||||
):
|
||||
text = str(item.get(field, ""))
|
||||
if not text.strip():
|
||||
return False
|
||||
if is_text_config:
|
||||
pp = config.preprocessing
|
||||
if pp.min_chars > 0 and len(text) < pp.min_chars:
|
||||
return False
|
||||
if len(text) > pp.max_chars:
|
||||
return False
|
||||
ids = tokenizer.encode(text, add_special_tokens=add_special)
|
||||
all_ids.extend(ids)
|
||||
val = 1 if action == "train" else 0
|
||||
loss_mask.extend([val] * len(ids))
|
||||
return True
|
||||
|
||||
|
||||
class BaseMaskBuilder(ABC):
|
||||
"""Convert a JSONL item into token ids and optional loss_mask."""
|
||||
|
||||
@abstractmethod
|
||||
def build(self, item: dict, config, tokenizer) -> Optional[dict]: ...
|
||||
|
||||
|
||||
class MaskBuilderFactory(BaseFactory["BaseMaskBuilder"]):
|
||||
pass
|
||||
|
||||
|
||||
@MaskBuilderFactory.register("single")
|
||||
class SingleOutputMaskBuilder(BaseMaskBuilder):
|
||||
"""Build a single output sequence with optional loss mask.
|
||||
|
||||
Expects ``config.input.sections`` (list of section specs).
|
||||
"""
|
||||
|
||||
def __init__(self, renderer: Optional[SectionRenderer] = None):
|
||||
self.renderer = renderer or SectionRenderer()
|
||||
|
||||
def build(self, item: dict, config, tokenizer) -> Optional[dict]:
|
||||
sections = config.input.sections
|
||||
if not sections:
|
||||
return None
|
||||
|
||||
ids, mask = self.renderer.process_sections(
|
||||
item, sections, config, tokenizer, is_top_level=True
|
||||
)
|
||||
if ids is None:
|
||||
return None
|
||||
|
||||
result: dict = {
|
||||
"sequence": ids,
|
||||
"domain": _extract_domain(item, config.output.domain_key),
|
||||
}
|
||||
if not all(m == 1 for m in mask):
|
||||
result["loss_mask"] = mask
|
||||
return result
|
||||
|
||||
|
||||
@MaskBuilderFactory.register("multi")
|
||||
class MultiOutputMaskBuilder(BaseMaskBuilder):
|
||||
"""Build multiple output sequences (DPO / GRPO).
|
||||
|
||||
Expects ``config.input.sources`` (dict of output_key → spec).
|
||||
"""
|
||||
|
||||
def __init__(self, renderer: Optional[SectionRenderer] = None):
|
||||
self.renderer = renderer or SectionRenderer()
|
||||
|
||||
def build(self, item: dict, config, tokenizer) -> Optional[dict]:
|
||||
sources_spec = getattr(config.input, "sources", None)
|
||||
if not sources_spec:
|
||||
return None
|
||||
|
||||
result: dict = {}
|
||||
any_output = False
|
||||
|
||||
for output_key, spec in sources_spec.items():
|
||||
sections = spec.get("sections", [])
|
||||
if not sections:
|
||||
continue
|
||||
|
||||
if self.renderer.is_value_section(sections):
|
||||
ids = self.renderer.extract_raw_value(item, sections)
|
||||
if ids is None:
|
||||
continue
|
||||
result[output_key] = ids
|
||||
any_output = True
|
||||
continue
|
||||
|
||||
list_field = spec.get("list_field", False)
|
||||
mask_key = spec.get("mask_key", f"{output_key}_mask")
|
||||
|
||||
if list_field:
|
||||
ids, mask = self.renderer.process_list_field(
|
||||
item, sections, config, tokenizer
|
||||
)
|
||||
if ids is None:
|
||||
continue
|
||||
# ids is List[List[int]] — preserve per-response structure
|
||||
result[output_key] = ids
|
||||
if mask is not None:
|
||||
result[mask_key] = mask
|
||||
any_output = True
|
||||
continue
|
||||
|
||||
ids, mask = self.renderer.process_sections(
|
||||
item, sections, config, tokenizer, is_top_level=True
|
||||
)
|
||||
|
||||
if ids is None:
|
||||
continue
|
||||
|
||||
result[output_key] = ids
|
||||
if not all(m == 1 for m in mask):
|
||||
result[mask_key] = mask
|
||||
elif "mask_key" in spec:
|
||||
result[mask_key] = mask
|
||||
|
||||
any_output = True
|
||||
|
||||
if not any_output:
|
||||
return None
|
||||
|
||||
result["domain"] = _extract_domain(item, config.output.domain_key)
|
||||
return result
|
||||
|
||||
|
||||
@MaskBuilderFactory.register("sectioned")
|
||||
class SectionedMaskBuilder(BaseMaskBuilder):
|
||||
"""Façade that dispatches to SingleOutputMaskBuilder or MultiOutputMaskBuilder.
|
||||
|
||||
Preserves backward compatibility for existing configs and code that rely
|
||||
on the ``"sectioned"`` factory name.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._single = SingleOutputMaskBuilder()
|
||||
self._multi = MultiOutputMaskBuilder()
|
||||
|
||||
def build(self, item: dict, config, tokenizer) -> Optional[dict]:
|
||||
sources_spec = getattr(config.input, "sources", None)
|
||||
if sources_spec:
|
||||
return self._multi.build(item, config, tokenizer)
|
||||
return self._single.build(item, config, tokenizer)
|
||||
@@ -0,0 +1,124 @@
|
||||
"""Shared preprocessing kernel used by both :class:`Pipeline` and
|
||||
:class:`TokenizeTransform`.
|
||||
|
||||
The two entry points previously duplicated ~60 % of their logic:
|
||||
record iteration, mask-builder invocation, primary-id extraction,
|
||||
per-key accumulation, dtype inference and position-id generation.
|
||||
This module factors out the common core as pure functions so that
|
||||
the online (``TokenizeTransform``) and offline (``Pipeline``) paths
|
||||
stay in lockstep.
|
||||
"""
|
||||
|
||||
from itertools import chain
|
||||
from typing import Dict, Iterator, List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
from astrai.preprocessing.builder import MaskBuilderFactory
|
||||
from astrai.preprocessing.position_id import PositionIdStrategyFactory
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
|
||||
def build_preprocessing_components(config: PipelineConfig, tokenizer_path: str):
|
||||
"""Load tokenizer, mask builder and position-id strategy together.
|
||||
|
||||
Both ``Pipeline`` and ``TokenizeTransform`` need the same triple;
|
||||
centralising the construction avoids drift (e.g. one path forgetting
|
||||
to create the position-id strategy).
|
||||
"""
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
||||
mask_builder = MaskBuilderFactory.create("sectioned")
|
||||
position_strategy = PositionIdStrategyFactory.create(
|
||||
config.output.position_ids_mode
|
||||
)
|
||||
return tokenizer, mask_builder, position_strategy
|
||||
|
||||
|
||||
def primary_ids(result: dict) -> List[int]:
|
||||
"""Return the first flat int-list value in *result*.
|
||||
|
||||
Used for token counting and position-id generation when the
|
||||
primary key name is not known (DPO uses ``chosen``, GRPO uses
|
||||
``prompts``, SFT uses ``sequence``).
|
||||
"""
|
||||
for val in result.values():
|
||||
if isinstance(val, list) and val and isinstance(val[0], int):
|
||||
return val
|
||||
return []
|
||||
|
||||
|
||||
def infer_dtype(ids: List) -> torch.dtype:
|
||||
"""Float values become float32, everything else int32."""
|
||||
if ids and isinstance(ids[0], float):
|
||||
return torch.float32
|
||||
return torch.int32
|
||||
|
||||
|
||||
def iter_raw_records(
|
||||
records: List[dict],
|
||||
mask_builder,
|
||||
config: PipelineConfig,
|
||||
tokenizer,
|
||||
) -> Iterator[dict]:
|
||||
"""Yield mask-builder output dicts for each record, skipping failures.
|
||||
|
||||
Drops ``domain`` from the result (callers that need it should read
|
||||
it before calling this). Each yielded dict maps a key
|
||||
(``sequence``, ``chosen``, ``responses``…) to either a flat
|
||||
``List[int]`` or a nested ``List[List[int]]`` (GRPO responses/masks).
|
||||
"""
|
||||
for item in records:
|
||||
result = mask_builder.build(item, config, tokenizer)
|
||||
if result is None:
|
||||
continue
|
||||
result.pop("domain", None)
|
||||
if not primary_ids(result):
|
||||
continue
|
||||
yield result
|
||||
|
||||
|
||||
def to_per_record_tensors(
|
||||
raw: Dict[str, list],
|
||||
) -> Dict[str, List[torch.Tensor]]:
|
||||
"""Convert an accumulated ``{key: [per-record ids]}`` dict to tensors.
|
||||
|
||||
Handles three shapes transparently:
|
||||
|
||||
- ``List[int]`` per record (``sequence``, ``chosen``…) → one tensor per record.
|
||||
- ``List[List[int]]`` per record (GRPO ``responses``/``masks``) → one
|
||||
``List[Tensor]`` per record (nested), preserving the per-response
|
||||
boundary so downstream code can index responses individually.
|
||||
- ``List[int]`` for the whole shard (pre-packed keys) → single tensor.
|
||||
|
||||
The detection mirrors the previous inline logic in
|
||||
``Pipeline._flush`` and ``TokenizeTransform.apply``.
|
||||
"""
|
||||
tensors: Dict[str, List[torch.Tensor]] = {}
|
||||
for key, ids_list in raw.items():
|
||||
if ids_list and isinstance(ids_list[0], list):
|
||||
tensors[key] = [
|
||||
[torch.tensor(sub, dtype=infer_dtype(sub)) for sub in ids]
|
||||
if ids and isinstance(ids[0], list)
|
||||
else torch.tensor(ids, dtype=infer_dtype(ids))
|
||||
for ids in ids_list
|
||||
]
|
||||
else:
|
||||
tensors[key] = [
|
||||
torch.tensor(list(chain.from_iterable(ids_list)), dtype=torch.int32)
|
||||
]
|
||||
return tensors
|
||||
|
||||
|
||||
def build_position_ids(
|
||||
sequences: List[List[int]],
|
||||
strategy,
|
||||
) -> Optional[List[int]]:
|
||||
"""Generate position ids for *sequences* using *strategy*.
|
||||
|
||||
Returns ``None`` when the strategy produces no ids (e.g. ``none``
|
||||
mode), so callers can skip attaching the key instead of storing
|
||||
an empty list.
|
||||
"""
|
||||
pos_ids = strategy.generate(sequences)
|
||||
return pos_ids or None
|
||||
@@ -0,0 +1,176 @@
|
||||
"""Sequence packing strategies for shard-level reordering and truncation.
|
||||
|
||||
Each strategy receives the accumulated ``{key: [list of token lists]}``
|
||||
dict for a shard and returns a reordered / truncated version. The
|
||||
pipeline later flattens the result into contiguous tensors.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Dict, List
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
|
||||
def _truncate(seq: List[int], max_len: int, mode: str) -> List[int]:
|
||||
if len(seq) <= max_len:
|
||||
return seq
|
||||
if mode == "keep_end":
|
||||
return seq[-max_len:]
|
||||
return seq[:max_len]
|
||||
|
||||
|
||||
def plan_bfd(
|
||||
sequences: List[List[int]], max_packed_len: int, truncation_mode: str = "keep_start"
|
||||
) -> List[List[int]]:
|
||||
"""Best-Fit Decreasing bin packing of *sequences* into bins.
|
||||
|
||||
Returns a list of bins, each bin a list of original indices into
|
||||
*sequences*. Bin capacities are respected on the *truncated*
|
||||
length of each sequence (so a sequence longer than
|
||||
*max_packed_len* counts at *max_packed_len*).
|
||||
|
||||
Pure index-based so callers can apply the same plan to any
|
||||
aligned key (``loss_mask``, ``position_ids``…).
|
||||
"""
|
||||
n = len(sequences)
|
||||
order = sorted(range(n), key=lambda i: len(sequences[i]), reverse=True)
|
||||
bins: List[List[int]] = []
|
||||
bin_lengths: List[int] = []
|
||||
|
||||
for orig_idx in order:
|
||||
seq_len = len(_truncate(sequences[orig_idx], max_packed_len, truncation_mode))
|
||||
best_bin = None
|
||||
best_remain = max_packed_len + 1
|
||||
for i, bl in enumerate(bin_lengths):
|
||||
remain = max_packed_len - bl
|
||||
if seq_len <= remain < best_remain:
|
||||
best_remain = remain
|
||||
best_bin = i
|
||||
if best_bin is not None:
|
||||
bins[best_bin].append(orig_idx)
|
||||
bin_lengths[best_bin] += seq_len
|
||||
else:
|
||||
bins.append([orig_idx])
|
||||
bin_lengths.append(seq_len)
|
||||
|
||||
return bins
|
||||
|
||||
|
||||
class PackingStrategy(ABC):
|
||||
"""Reorder and truncate sequences within a shard."""
|
||||
|
||||
@abstractmethod
|
||||
def apply(
|
||||
self,
|
||||
keys: Dict[str, List[List[int]]],
|
||||
max_packed_len: int,
|
||||
truncation_mode: str,
|
||||
) -> Dict[str, List[List[int]]]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class PackingStrategyFactory(BaseFactory["PackingStrategy"]):
|
||||
pass
|
||||
|
||||
|
||||
@PackingStrategyFactory.register("simple")
|
||||
class SimplePacking(PackingStrategy):
|
||||
def apply(
|
||||
self,
|
||||
keys: Dict[str, List[List[int]]],
|
||||
max_packed_len: int,
|
||||
truncation_mode: str,
|
||||
) -> Dict[str, List[List[int]]]:
|
||||
return {
|
||||
k: [_truncate(v, max_packed_len, truncation_mode) for v in vals]
|
||||
for k, vals in keys.items()
|
||||
}
|
||||
|
||||
|
||||
@PackingStrategyFactory.register("bfd")
|
||||
class BFDPacking(PackingStrategy):
|
||||
"""Best-Fit Decreasing bin packing.
|
||||
|
||||
Assigns sequences to bins using a best-fit heuristic (sorted by
|
||||
decreasing length) and concatenates sequences within each bin into
|
||||
a single packed sequence. Packed sequences are truncated to
|
||||
*max_packed_len* so that each packed bin fits within one context
|
||||
window during training.
|
||||
"""
|
||||
|
||||
def apply(
|
||||
self,
|
||||
keys: Dict[str, List[List[int]]],
|
||||
max_packed_len: int,
|
||||
truncation_mode: str,
|
||||
) -> Dict[str, List[List[int]]]:
|
||||
sequences = keys.get("sequence", [])
|
||||
if not sequences:
|
||||
return keys
|
||||
bins = plan_bfd(sequences, max_packed_len, truncation_mode)
|
||||
|
||||
packed: Dict[str, List[List[int]]] = {}
|
||||
for k, vals in keys.items():
|
||||
packed[k] = [
|
||||
_truncate(
|
||||
self._concat_bin(vals, bin_indices),
|
||||
max_packed_len,
|
||||
truncation_mode,
|
||||
)
|
||||
for bin_indices in bins
|
||||
]
|
||||
return packed
|
||||
|
||||
@staticmethod
|
||||
def _concat_bin(vals: List[List[int]], indices: List[int]) -> List[int]:
|
||||
result: List[int] = []
|
||||
for i in indices:
|
||||
result.extend(vals[i])
|
||||
return result
|
||||
|
||||
|
||||
@PackingStrategyFactory.register("bfd_split")
|
||||
class BFDSplitPacking(BFDPacking):
|
||||
"""BFD packing with over-length sequences split into chunks.
|
||||
|
||||
Sequences longer than *max_packed_len* are split into consecutive
|
||||
chunks of at most *max_packed_len* tokens instead of being
|
||||
truncated. Each chunk becomes an independent sequence that enters
|
||||
BFD planning. All keys (``loss_mask``, ``position_ids``, …) are
|
||||
split in lockstep so per-token alignment is preserved.
|
||||
|
||||
Note: because each chunk is treated as a separate document, the
|
||||
second chunk of a split sequence loses the preceding context.
|
||||
"""
|
||||
|
||||
def apply(
|
||||
self,
|
||||
keys: Dict[str, List[List[int]]],
|
||||
max_packed_len: int,
|
||||
truncation_mode: str,
|
||||
) -> Dict[str, List[List[int]]]:
|
||||
sequences = keys.get("sequence", [])
|
||||
if not sequences:
|
||||
return keys
|
||||
if max_packed_len <= 0:
|
||||
return super().apply(keys, max_packed_len, truncation_mode)
|
||||
|
||||
split_keys = self._split_all(keys, max_packed_len)
|
||||
return super().apply(split_keys, max_packed_len, truncation_mode)
|
||||
|
||||
@staticmethod
|
||||
def _split_all(
|
||||
keys: Dict[str, List[List[int]]], max_packed_len: int
|
||||
) -> Dict[str, List[List[int]]]:
|
||||
"""Split every sequence exceeding *max_packed_len* into chunks,
|
||||
applying the same chunk boundaries to all keys."""
|
||||
sequences = keys["sequence"]
|
||||
chunk_bounds = [list(range(0, len(s), max_packed_len)) for s in sequences]
|
||||
result: Dict[str, List[List[int]]] = {}
|
||||
for key, vals in keys.items():
|
||||
split_vals: List[List[int]] = []
|
||||
for val, starts in zip(vals, chunk_bounds):
|
||||
for start in starts:
|
||||
split_vals.append(val[start : start + max_packed_len])
|
||||
result[key] = split_vals
|
||||
return result
|
||||
@@ -0,0 +1,248 @@
|
||||
"""Config-driven JSONL preprocessing pipeline.
|
||||
|
||||
Composes a :class:`BaseMaskBuilder` (selected by ``input.type``) with
|
||||
sharding and flush to ``.h5`` / ``.bin`` storage. Packing, position-id
|
||||
generation and storage writing are each delegated to pluggable strategies,
|
||||
dispatched by configuration keys.
|
||||
|
||||
Record iteration, mask building, primary-id extraction and per-key
|
||||
accumulation are shared with :class:`TokenizeTransform` via the
|
||||
:mod:`astrai.preprocessing.core` helpers.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from collections import defaultdict
|
||||
from itertools import chain
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
import torch
|
||||
import tqdm
|
||||
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
from astrai.preprocessing.core import (
|
||||
build_preprocessing_components,
|
||||
iter_raw_records,
|
||||
primary_ids,
|
||||
)
|
||||
from astrai.preprocessing.packing import PackingStrategyFactory
|
||||
from astrai.preprocessing.writer import StoreWriterFactory
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_STR_TO_DTYPE: dict[str, torch.dtype] = {
|
||||
"bool": torch.bool,
|
||||
"uint8": torch.uint8,
|
||||
"int8": torch.int8,
|
||||
"int16": torch.int16,
|
||||
"int32": torch.int32,
|
||||
"int64": torch.int64,
|
||||
"float16": torch.float16,
|
||||
"float32": torch.float32,
|
||||
"float64": torch.float64,
|
||||
}
|
||||
|
||||
|
||||
def filter_by_length(text: str, min_len: int = 50, max_len: int = 2_000_000) -> bool:
|
||||
return min_len <= len(text) <= max_len
|
||||
|
||||
|
||||
class Pipeline:
|
||||
"""Tokenization pipeline driven by a declarative :class:`PipelineConfig`.
|
||||
|
||||
Usage::
|
||||
|
||||
config = PipelineConfig.from_file("sft_pipeline.json")
|
||||
Pipeline(config, ["data.jsonl"], output_dir="out", tokenizer_path="params").run()
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: PipelineConfig,
|
||||
input_paths: list[str],
|
||||
output_dir: str,
|
||||
tokenizer_path: str,
|
||||
):
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
self.config = config
|
||||
self.paths = input_paths
|
||||
self.output_dir = output_dir
|
||||
self.tokenizer_path = tokenizer_path
|
||||
|
||||
self.tokenizer, self.mask_builder, self._position_id = (
|
||||
build_preprocessing_components(config, tokenizer_path)
|
||||
)
|
||||
self._packer = PackingStrategyFactory.create(
|
||||
config.preprocessing.packing_strategy
|
||||
)
|
||||
self._writer = StoreWriterFactory.create(config.output.storage_format)
|
||||
|
||||
def transform(self, item: dict) -> Optional[dict]:
|
||||
return self.mask_builder.build(item, self.config, self.tokenizer)
|
||||
|
||||
def run(self):
|
||||
domains: dict = defaultdict(lambda: defaultdict(list))
|
||||
total_tokens = 0
|
||||
shard_idx: dict[str, int] = defaultdict(int)
|
||||
count = 0
|
||||
|
||||
pp = self.config.preprocessing
|
||||
|
||||
for item in tqdm.tqdm(
|
||||
self._iter_items(), desc="Tokenizing", unit="docs", mininterval=0.5
|
||||
):
|
||||
if pp.max_items and count >= pp.max_items:
|
||||
break
|
||||
|
||||
try:
|
||||
result = self.transform(item)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to process item #%d, skipping", count + 1, exc_info=True
|
||||
)
|
||||
continue
|
||||
if result is None:
|
||||
continue
|
||||
|
||||
domain = result.pop("domain", "__default__")
|
||||
ids = primary_ids(result)
|
||||
if not ids:
|
||||
continue
|
||||
|
||||
bucket = domains[domain]
|
||||
self._align_bucket(bucket, result, ids)
|
||||
for key, val in result.items():
|
||||
bucket[key].append(val)
|
||||
|
||||
count += 1
|
||||
total_tokens += len(ids)
|
||||
|
||||
if total_tokens >= self.config.output.max_tokens_per_shard:
|
||||
self._flush(domains, shard_idx)
|
||||
domains.clear()
|
||||
total_tokens = 0
|
||||
|
||||
if total_tokens > 0:
|
||||
self._flush(domains, shard_idx)
|
||||
|
||||
@staticmethod
|
||||
def _align_bucket(bucket: dict, result: dict, ids: list):
|
||||
"""Pad previously-accumulated keys that are missing from *result*."""
|
||||
for key in list(bucket.keys()):
|
||||
if key in result:
|
||||
continue
|
||||
bucket[key].append([0] * len(ids))
|
||||
|
||||
def _iter_items(self):
|
||||
for path in self.paths:
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
if path.endswith(".json"):
|
||||
data = json.load(f)
|
||||
if isinstance(data, dict):
|
||||
yield data
|
||||
elif isinstance(data, list):
|
||||
yield from data
|
||||
else:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
yield json.loads(line)
|
||||
|
||||
def _flush(self, domains, shard_idx):
|
||||
for domain, keys in domains.items():
|
||||
idx = shard_idx[domain]
|
||||
|
||||
pp = self.config.preprocessing
|
||||
original_sequences = keys.get("sequence", [])
|
||||
mode = self.config.output.position_ids_mode
|
||||
|
||||
keys = self._inject_doc_reset_position_ids(keys, mode, original_sequences)
|
||||
keys = self._packer.apply(dict(keys), pp.max_packed_len, pp.truncation_mode)
|
||||
tensors = self._to_tensors(keys)
|
||||
tensors = self._inject_continuous_position_ids(
|
||||
tensors, mode, keys.get("sequence", [])
|
||||
)
|
||||
|
||||
self._writer.save(self.output_dir, domain, idx, tensors)
|
||||
shard_idx[domain] = idx + 1
|
||||
|
||||
first_key = "sequence" if "sequence" in tensors else next(iter(tensors))
|
||||
tqdm.tqdm.write(
|
||||
f" saved {domain}/shard_{idx:04d} "
|
||||
f"({tensors[first_key][0].numel():,} tokens)"
|
||||
)
|
||||
|
||||
def _inject_doc_reset_position_ids(
|
||||
self,
|
||||
keys: Dict[str, list],
|
||||
mode: str,
|
||||
original_sequences: List[List[int]],
|
||||
) -> Dict[str, list]:
|
||||
"""Attach per-document position_ids before packing (``doc_reset``).
|
||||
|
||||
``doc_reset`` position ids must enter the packer so that each
|
||||
packed bin concatenates the per-doc ranges in bin order. The
|
||||
per-record structure ``[range(len(s)) for s in seqs]`` is required
|
||||
by the packer (it concatenates per-record lists per bin); the
|
||||
``PositionIdStrategy.generate`` flattens, so it cannot be used
|
||||
directly here — it is only consulted for the ``continuous``
|
||||
post-packing path.
|
||||
"""
|
||||
if mode != "doc_reset" or not original_sequences:
|
||||
return keys
|
||||
keys["position_ids"] = [list(range(len(s))) for s in original_sequences]
|
||||
return keys
|
||||
|
||||
def _inject_continuous_position_ids(
|
||||
self,
|
||||
tensors: Dict[str, List[torch.Tensor]],
|
||||
mode: str,
|
||||
packed_sequences: List[List[int]],
|
||||
) -> Dict[str, List[torch.Tensor]]:
|
||||
"""Attach a single continuous position_ids tensor after packing.
|
||||
|
||||
``continuous`` mode spans the whole shard (post-packing), so it
|
||||
cannot participate in bin packing — it is computed from the
|
||||
packed sequences and appended directly to the tensor dict.
|
||||
"""
|
||||
if mode != "continuous" or not packed_sequences:
|
||||
return tensors
|
||||
pos_ids = self._position_id.generate(packed_sequences)
|
||||
if pos_ids:
|
||||
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
||||
return tensors
|
||||
|
||||
def _to_tensors(self, keys: Dict[str, list]) -> Dict[str, List[torch.Tensor]]:
|
||||
"""Convert packed per-key id lists to tensors.
|
||||
|
||||
Honours ``config.output.dtype`` overrides per key; falls back to
|
||||
``int32``. Handles three shapes (see
|
||||
:func:`astrai.preprocessing.core.to_per_record_tensors` for the
|
||||
equivalent online-path helper):
|
||||
- ``List[int]`` per record → one tensor per record.
|
||||
- ``List[List[int]]`` per record (GRPO responses/masks) → one tensor
|
||||
per record, inner lists flattened.
|
||||
- ``List[int]`` for the whole shard (pre-packed keys) → single tensor.
|
||||
"""
|
||||
tensors: Dict[str, List[torch.Tensor]] = {}
|
||||
for key, ids_list in keys.items():
|
||||
dt = _STR_TO_DTYPE.get(
|
||||
self.config.output.dtype.get(key, "int32"), torch.int32
|
||||
)
|
||||
if ids_list and isinstance(ids_list[0], list):
|
||||
tensors[key] = [
|
||||
torch.tensor(
|
||||
list(chain.from_iterable(ids))
|
||||
if ids and isinstance(ids[0], list)
|
||||
else ids,
|
||||
dtype=dt,
|
||||
)
|
||||
for ids in ids_list
|
||||
]
|
||||
else:
|
||||
tensors[key] = [
|
||||
torch.tensor(list(chain.from_iterable(ids_list)), dtype=dt)
|
||||
]
|
||||
return tensors
|
||||
@@ -0,0 +1,46 @@
|
||||
"""Position-id generation strategies for packed sequences.
|
||||
|
||||
Each strategy takes the list of per-document token sequences after packing
|
||||
and returns a flat list of position ids (same total length as all
|
||||
sequences combined). The pipeline wraps the result into a tensor and
|
||||
attaches it as ``position_ids``.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
|
||||
class PositionIdStrategy(ABC):
|
||||
"""Generate ``position_ids`` for packed sequences."""
|
||||
|
||||
@abstractmethod
|
||||
def generate(self, sequences: List[List[int]]) -> List[int]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class PositionIdStrategyFactory(BaseFactory["PositionIdStrategy"]):
|
||||
pass
|
||||
|
||||
|
||||
@PositionIdStrategyFactory.register("none")
|
||||
class NoPositionId(PositionIdStrategy):
|
||||
def generate(self, sequences: List[List[int]]) -> List[int]:
|
||||
return []
|
||||
|
||||
|
||||
@PositionIdStrategyFactory.register("doc_reset")
|
||||
class DocResetPositionId(PositionIdStrategy):
|
||||
def generate(self, sequences: List[List[int]]) -> List[int]:
|
||||
pos_ids = []
|
||||
for seq in sequences:
|
||||
pos_ids.extend(range(len(seq)))
|
||||
return pos_ids
|
||||
|
||||
|
||||
@PositionIdStrategyFactory.register("continuous")
|
||||
class ContinuousPositionId(PositionIdStrategy):
|
||||
def generate(self, sequences: List[List[int]]) -> List[int]:
|
||||
total = sum(len(seq) for seq in sequences)
|
||||
return list(range(total))
|
||||
@@ -0,0 +1,92 @@
|
||||
"""Tokenization transform for JSONL record streams.
|
||||
|
||||
Bridges the Reader layer (``JsonlStore`` reads raw JSON records) and the
|
||||
Dataset layer (expects per-record tensors). Holds the tokenizer,
|
||||
mask-builder and position-id strategy together so that I/O code stays
|
||||
free of model dependencies.
|
||||
|
||||
The record-processing core (mask building, primary-id extraction,
|
||||
per-key tensorisation, position-id generation) is shared with
|
||||
:class:`astrai.preprocessing.pipeline.Pipeline` via the
|
||||
:mod:`astrai.preprocessing.core` helpers.
|
||||
"""
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Dict, List
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.config.preprocess_config import PipelineConfig
|
||||
from astrai.preprocessing.core import (
|
||||
build_position_ids,
|
||||
build_preprocessing_components,
|
||||
iter_raw_records,
|
||||
to_per_record_tensors,
|
||||
)
|
||||
|
||||
|
||||
class TokenizeTransform:
|
||||
"""Tokenize raw JSONL record dicts into per-key tensor lists.
|
||||
|
||||
Owns the three preprocessing concerns that were previously inlined in
|
||||
``JsonlStore``: tokenization, loss-mask construction and position-id
|
||||
generation. Constructing it loads the tokenizer, so it is intentionally
|
||||
cheap to pass around once built.
|
||||
|
||||
Args:
|
||||
config: Pipeline config describing sections / masks / position mode.
|
||||
tokenizer_path: Path passed to ``AutoTokenizer.from_pretrained``.
|
||||
"""
|
||||
|
||||
def __init__(self, config: PipelineConfig, tokenizer_path: str):
|
||||
self.config = config
|
||||
self.tokenizer, self.mask_builder, self.position_strategy = (
|
||||
build_preprocessing_components(config, tokenizer_path)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_config_file(cls, config_path: str) -> "TokenizeTransform":
|
||||
"""Build from a ``dataset_config.json`` file path.
|
||||
|
||||
The config file follows :class:`PipelineConfig` schema with an
|
||||
extra ``tokenizer_path`` field. When omitted, the config's
|
||||
parent directory is used as the tokenizer path.
|
||||
"""
|
||||
root = Path(config_path).parent
|
||||
with open(config_path, "r", encoding="utf-8") as f:
|
||||
raw_config = json.load(f)
|
||||
tokenizer_path = raw_config.pop("tokenizer_path", None) or str(root)
|
||||
config = PipelineConfig.from_dict(raw_config)
|
||||
return cls(config, tokenizer_path)
|
||||
|
||||
def apply(self, records: List[dict]) -> Dict[str, list]:
|
||||
"""Tokenize a list of raw record dicts.
|
||||
|
||||
Returns a dict mapping key (``sequence``, ``chosen``, ``responses``,
|
||||
…) to a list of per-record tensors (or nested tensor lists for
|
||||
multi-response keys such as GRPO ``responses``).
|
||||
"""
|
||||
raw: Dict[str, list] = {}
|
||||
doc_sequences: List[List[int]] = []
|
||||
|
||||
for result in iter_raw_records(
|
||||
records, self.mask_builder, self.config, self.tokenizer
|
||||
):
|
||||
primary = None
|
||||
for val in result.values():
|
||||
if isinstance(val, list) and val and isinstance(val[0], int):
|
||||
primary = val
|
||||
break
|
||||
if primary is not None:
|
||||
doc_sequences.append(primary)
|
||||
for key, ids in result.items():
|
||||
raw.setdefault(key, []).append(ids)
|
||||
|
||||
tensors = to_per_record_tensors(raw)
|
||||
|
||||
pos_ids = build_position_ids(doc_sequences, self.position_strategy)
|
||||
if pos_ids is not None:
|
||||
tensors["position_ids"] = [torch.tensor(pos_ids, dtype=torch.int32)]
|
||||
|
||||
return tensors
|
||||
@@ -0,0 +1,75 @@
|
||||
"""Storage writer strategies for pipeline output.
|
||||
|
||||
The :class:`StoreWriter` abstraction decouples the pipeline from the
|
||||
concrete storage format (bin / h5). The pipeline builds a ``{key:
|
||||
List[Tensor]}`` dict and delegates the write to the writer selected
|
||||
by ``output.storage_format``.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Dict, List
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.serialization import save_bin, save_h5
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class StoreWriter(ABC):
|
||||
"""Write pre-tokenized tensors to disk in a format-specific way."""
|
||||
|
||||
@abstractmethod
|
||||
def save(
|
||||
self,
|
||||
output_dir: str,
|
||||
domain: str,
|
||||
shard_idx: int,
|
||||
tensors: Dict[str, List[torch.Tensor]],
|
||||
) -> None: ...
|
||||
|
||||
|
||||
class StoreWriterFactory(BaseFactory["StoreWriter"]):
|
||||
pass
|
||||
|
||||
|
||||
@StoreWriterFactory.register("bin")
|
||||
class BinWriter(StoreWriter):
|
||||
def save(self, output_dir, domain, shard_idx, tensors):
|
||||
shard_path = os.path.join(output_dir, domain, f"shard_{shard_idx:04d}")
|
||||
try:
|
||||
save_bin(shard_path, tensors)
|
||||
except Exception:
|
||||
if os.path.exists(shard_path):
|
||||
shutil.rmtree(shard_path, ignore_errors=True)
|
||||
logger.error(
|
||||
"Failed to write shard %s/%s_%04d, cleaned up partial output",
|
||||
domain,
|
||||
"shard",
|
||||
shard_idx,
|
||||
exc_info=True,
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
@StoreWriterFactory.register("h5")
|
||||
class H5Writer(StoreWriter):
|
||||
def save(self, output_dir, domain, shard_idx, tensors):
|
||||
chunk_dir = os.path.join(output_dir, domain)
|
||||
file_path = os.path.join(chunk_dir, f"data_{shard_idx:04d}.h5")
|
||||
try:
|
||||
save_h5(chunk_dir, f"data_{shard_idx:04d}", tensors)
|
||||
except Exception:
|
||||
if os.path.exists(file_path):
|
||||
os.remove(file_path)
|
||||
logger.error(
|
||||
"Failed to write shard %s/data_%04d.h5, cleaned up partial output",
|
||||
domain,
|
||||
shard_idx,
|
||||
exc_info=True,
|
||||
)
|
||||
raise
|
||||
@@ -0,0 +1,45 @@
|
||||
"""Serialization utilities for models and datasets.
|
||||
|
||||
This package re-exports checkpoint helpers and dataset storage helpers so
|
||||
that existing imports from ``astrai.serialization`` continue to work.
|
||||
"""
|
||||
|
||||
from astrai.serialization.checkpoint import (
|
||||
Checkpoint,
|
||||
load_json,
|
||||
load_model_config,
|
||||
load_model_weights,
|
||||
load_safetensors,
|
||||
load_state_dict,
|
||||
load_torch,
|
||||
save_json,
|
||||
save_model,
|
||||
save_safetensors,
|
||||
save_torch,
|
||||
)
|
||||
from astrai.serialization.dataset import (
|
||||
load_bin,
|
||||
load_bin_offsets,
|
||||
load_h5,
|
||||
save_bin,
|
||||
save_h5,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Checkpoint",
|
||||
"load_json",
|
||||
"load_model_config",
|
||||
"load_model_weights",
|
||||
"load_safetensors",
|
||||
"load_state_dict",
|
||||
"load_torch",
|
||||
"save_json",
|
||||
"save_model",
|
||||
"save_safetensors",
|
||||
"save_torch",
|
||||
"load_bin",
|
||||
"load_bin_offsets",
|
||||
"load_h5",
|
||||
"save_bin",
|
||||
"save_h5",
|
||||
]
|
||||
@@ -1,9 +1,11 @@
|
||||
"""Model checkpoint serialization helpers."""
|
||||
|
||||
import io
|
||||
import json
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Union
|
||||
from typing import Any, Dict, Optional, Union
|
||||
|
||||
import safetensors.torch as st
|
||||
import torch
|
||||
@@ -136,7 +138,7 @@ def load_state_dict(path: Union[str, Path], broadcast: bool = False) -> dict:
|
||||
class Checkpoint:
|
||||
state_dict: Dict[str, Any] = field(default_factory=dict)
|
||||
epoch: int = 0
|
||||
iteration: int = 0
|
||||
consumed_samples: int = 0
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
meta: Dict[str, Any] = field(default_factory=dict)
|
||||
config: Dict[str, Any] = field(default_factory=dict)
|
||||
@@ -145,12 +147,9 @@ class Checkpoint:
|
||||
save_path = Path(save_dir)
|
||||
save_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if get_rank() != 0:
|
||||
return
|
||||
|
||||
meta = {
|
||||
"epoch": self.epoch,
|
||||
"iteration": self.iteration,
|
||||
"consumed_samples": self.consumed_samples,
|
||||
"timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"),
|
||||
**self.meta,
|
||||
}
|
||||
@@ -176,7 +175,27 @@ class Checkpoint:
|
||||
return cls(
|
||||
state_dict=state_dict,
|
||||
epoch=meta.get("epoch", 0),
|
||||
iteration=meta.get("iteration", 0),
|
||||
consumed_samples=meta.get("consumed_samples", 0),
|
||||
extra=extra,
|
||||
meta=meta,
|
||||
config=config,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def load_any(cls, save_dir: str, broadcast: bool = False) -> Optional["Checkpoint"]:
|
||||
save_path = Path(save_dir)
|
||||
meta_path = save_path / _META_FILE
|
||||
weights_path = save_path / _WEIGHTS_FILE
|
||||
|
||||
if meta_path.exists():
|
||||
return cls.load(save_dir, broadcast=broadcast)
|
||||
|
||||
if weights_path.exists():
|
||||
state_dict = load_state_dict(weights_path, broadcast=broadcast)
|
||||
config = {}
|
||||
config_path = save_path / _CONFIG_FILE
|
||||
if config_path.exists():
|
||||
config = load_json(config_path, broadcast)
|
||||
return cls(state_dict=state_dict, config=config)
|
||||
|
||||
return None
|
||||
@@ -0,0 +1,123 @@
|
||||
"""Dataset storage serialization helpers (HDF5 / memory-mapped binary)."""
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import h5py
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
def save_h5(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]):
|
||||
os.makedirs(file_path, exist_ok=True)
|
||||
full_file_path = os.path.join(file_path, f"{file_name}.h5")
|
||||
with h5py.File(full_file_path, "w") as f:
|
||||
for key, tensors in tensor_group.items():
|
||||
grp = f.create_group(key)
|
||||
for idx, tensor in enumerate(tensors):
|
||||
arr = tensor.cpu().numpy()
|
||||
grp.create_dataset(f"data_{idx}", data=arr)
|
||||
|
||||
|
||||
def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
|
||||
tensor_group: Dict[str, List[Tensor]] = {}
|
||||
|
||||
root_path = Path(file_path)
|
||||
if root_path.is_file() and root_path.suffix in (".h5", ".hdf5"):
|
||||
h5_files = [root_path]
|
||||
else:
|
||||
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_bin(
|
||||
file_path: str,
|
||||
tensor_group: Dict[str, List[Tensor]],
|
||||
record_keys: Optional[List[str]] = None,
|
||||
):
|
||||
"""Save tensors as memory-mapped binary files.
|
||||
|
||||
When *record_keys* is provided, those keys are written with per-record
|
||||
cumulative offsets in ``meta.json`` so that ``MmapStore.fetch_record``
|
||||
can slice individual records from the concatenated binary without
|
||||
cross-record concatenation. Keys not in *record_keys* (e.g. SEQ
|
||||
``sequence``) are written as a single contiguous stream without
|
||||
offsets, preserving backward compatibility.
|
||||
|
||||
Nested keys (``List[List[Tensor]]`` such as GRPO ``responses``) are
|
||||
not supported in bin format — use H5 for those.
|
||||
"""
|
||||
os.makedirs(file_path, exist_ok=True)
|
||||
record_keys = set(record_keys or [])
|
||||
meta = {}
|
||||
for key, tensors in tensor_group.items():
|
||||
if tensors and isinstance(tensors[0], list):
|
||||
raise ValueError(
|
||||
f"Nested key '{key}' (List[List[Tensor]]) is not supported "
|
||||
f"in bin format. Use H5 or JSONL storage instead."
|
||||
)
|
||||
cat = torch.cat(tensors, dim=0)
|
||||
entry: Dict[str, Any] = {
|
||||
"shape": list(cat.shape),
|
||||
"dtype": str(cat.dtype).split(".")[-1],
|
||||
}
|
||||
if key in record_keys:
|
||||
offsets = [0]
|
||||
for t in tensors:
|
||||
offsets.append(offsets[-1] + t.shape[0])
|
||||
entry["offsets"] = offsets
|
||||
meta[key] = entry
|
||||
np.asarray(cat.cpu().numpy()).tofile(os.path.join(file_path, f"{key}.bin"))
|
||||
with open(os.path.join(file_path, "meta.json"), "w") as f:
|
||||
json.dump(meta, f)
|
||||
|
||||
|
||||
def load_bin(file_path: str) -> Dict[str, List[Tensor]]:
|
||||
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
||||
meta = json.load(f)
|
||||
segments: Dict[str, List[Tensor]] = {}
|
||||
for key, info in meta.items():
|
||||
arr = np.memmap(
|
||||
os.path.join(file_path, f"{key}.bin"),
|
||||
dtype=info["dtype"],
|
||||
mode="r+",
|
||||
shape=tuple(info["shape"]),
|
||||
)
|
||||
segments[key] = [torch.from_numpy(arr)]
|
||||
return segments
|
||||
|
||||
|
||||
def load_bin_offsets(file_path: str) -> Dict[str, List[int]]:
|
||||
"""Read per-record cumulative offsets from ``meta.json``.
|
||||
|
||||
Returns an empty dict when no key has offsets (legacy bin files),
|
||||
in which case record-mode access falls back to per-record segment
|
||||
indexing (H5/JSONL layout).
|
||||
"""
|
||||
with open(os.path.join(file_path, "meta.json"), "r") as f:
|
||||
meta = json.load(f)
|
||||
offsets: Dict[str, List[int]] = {}
|
||||
for key, info in meta.items():
|
||||
if "offsets" in info:
|
||||
offsets[key] = info["offsets"]
|
||||
return offsets
|
||||
@@ -1,13 +1,11 @@
|
||||
from dataclasses import dataclass
|
||||
from functools import cached_property
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from jinja2 import Template
|
||||
|
||||
# Message type for chat messages
|
||||
type MessageType = Dict[str, Any]
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChatTemplate:
|
||||
"""A chat template with Jinja2 rendering support.
|
||||
|
||||
@@ -15,23 +13,36 @@ class ChatTemplate:
|
||||
name: Unique identifier for the template.
|
||||
template_str: Jinja2 template string.
|
||||
description: Optional description.
|
||||
default_variables: Optional dictionary of default variable values
|
||||
that will be passed to the template if not overridden during rendering.
|
||||
default_variables: Optional dictionary of default variable values.
|
||||
special_tokens: Optional dictionary mapping token names to their string values.
|
||||
These tokens are automatically added to the template variables.
|
||||
"""
|
||||
|
||||
name: str
|
||||
template_str: str
|
||||
description: str = ""
|
||||
default_variables: Dict[str, Any] = None
|
||||
special_tokens: Dict[str, str] = None
|
||||
def __init__(
|
||||
self,
|
||||
name: str = "",
|
||||
template_str: str = "",
|
||||
description: str = "",
|
||||
default_variables: Optional[Dict[str, Any]] = None,
|
||||
special_tokens: Optional[Dict[str, str]] = None,
|
||||
):
|
||||
self.name = name
|
||||
self.template_str = template_str
|
||||
self.description = description
|
||||
self.default_variables = default_variables or {}
|
||||
self.special_tokens = special_tokens or {}
|
||||
|
||||
def __post_init__(self):
|
||||
if self.default_variables is None:
|
||||
self.default_variables = {}
|
||||
if self.special_tokens is None:
|
||||
self.special_tokens = {}
|
||||
@cached_property
|
||||
def _compiled(self) -> Template:
|
||||
"""Lazy-compiled Jinja2 template, cached on first access.
|
||||
|
||||
The compiled :class:`~jinja2.Template` holds a dynamically-generated
|
||||
``root`` render function whose ``__module__`` is ``None``; under
|
||||
``pickle`` it falls back to ``__main__`` and breaks ``spawn``-based
|
||||
multiprocessing. By deferring compilation to first access, the
|
||||
default pickle protocol serialises only ``template_str``; each
|
||||
worker rebuilds the cache on first render.
|
||||
"""
|
||||
return Template(self.template_str)
|
||||
|
||||
@classmethod
|
||||
def from_string(
|
||||
@@ -43,7 +54,7 @@ class ChatTemplate:
|
||||
) -> "ChatTemplate":
|
||||
"""Create a ChatTemplate instance directly from a template string."""
|
||||
return cls(
|
||||
name="", # empty name for ad‑hoc templates
|
||||
name="",
|
||||
template_str=template_str,
|
||||
description=description,
|
||||
default_variables=default_variables,
|
||||
@@ -73,5 +84,4 @@ class ChatTemplate:
|
||||
if system_prompt is not None:
|
||||
variables["system_prompt"] = system_prompt
|
||||
|
||||
jinja_template = Template(self.template_str)
|
||||
return jinja_template.render(**variables)
|
||||
return self._compiled.render(**variables)
|
||||
|
||||
@@ -164,7 +164,14 @@ class AutoTokenizer:
|
||||
- tokenizer.bos_token → returns string
|
||||
- tokenizer.bos_token_id → returns corresponding integer ID
|
||||
- tokenizer.stop_ids → returns list of corresponding integer IDs for all special tokens
|
||||
|
||||
Internal/private attrs are not intercepted: during unpickle
|
||||
``__dict__`` is empty, so probing ``self._special_token_map``
|
||||
would recurse infinitely.
|
||||
"""
|
||||
if key.startswith("_"):
|
||||
raise AttributeError(key)
|
||||
|
||||
# Handle stop_ids - return IDs for all special tokens
|
||||
if key == "stop_ids":
|
||||
stop_ids = []
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
from astrai.trainer.optim import Muon
|
||||
from astrai.trainer.schedule import BaseScheduler, SchedulerFactory
|
||||
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
||||
from astrai.trainer.train_callback import (
|
||||
@@ -10,8 +9,6 @@ from astrai.trainer.trainer import Trainer
|
||||
__all__ = [
|
||||
# Main trainer
|
||||
"Trainer",
|
||||
# Optimizer
|
||||
"Muon",
|
||||
# Strategy factory
|
||||
"StrategyFactory",
|
||||
"BaseStrategy",
|
||||
|
||||
@@ -1,42 +1,25 @@
|
||||
from typing import Any, Callable, Dict
|
||||
from typing import Dict
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
def _grad_stat(
|
||||
model: nn.Module, fn: Callable[[torch.Tensor], Any], default: Any
|
||||
) -> dict:
|
||||
results = {}
|
||||
for name, param in model.named_parameters():
|
||||
results[name] = default
|
||||
if param.grad is not None:
|
||||
results[name] = fn(param.grad.data)
|
||||
return results
|
||||
def grad_norm(model: nn.Module, per_param: bool = False) -> float | Dict[str, float]:
|
||||
grads = [p.grad.detach() for p in model.parameters() if p.grad is not None]
|
||||
if not grads:
|
||||
return 0.0
|
||||
|
||||
|
||||
def grad_norm(model: nn.Module, norm_type: int = 2) -> Dict[str, float]:
|
||||
return _grad_stat(model, lambda g: g.norm(norm_type).item(), 0.0)
|
||||
|
||||
|
||||
def grad_std(model: nn.Module) -> Dict[str, float]:
|
||||
return _grad_stat(model, lambda g: g.std().item(), 0.0)
|
||||
|
||||
|
||||
def grad_max(model: nn.Module) -> Dict[str, float]:
|
||||
return _grad_stat(model, lambda g: g.max().item(), -float("inf"))
|
||||
|
||||
|
||||
def grad_min(model: nn.Module) -> Dict[str, float]:
|
||||
return _grad_stat(model, lambda g: g.min().item(), float("inf"))
|
||||
|
||||
|
||||
def grad_mean(model: nn.Module) -> Dict[str, float]:
|
||||
return _grad_stat(model, lambda g: g.mean().item(), 0.0)
|
||||
|
||||
|
||||
def grad_nan_num(model: nn.Module) -> Dict[str, int]:
|
||||
return _grad_stat(model, lambda g: g.isnan().sum().item(), 0)
|
||||
total_sq = torch.stack([g.pow(2).sum() for g in grads]).sum()
|
||||
if per_param:
|
||||
norms = {}
|
||||
for name, param in model.named_parameters():
|
||||
if param.grad is not None:
|
||||
norms[name] = param.grad.norm(2).item()
|
||||
else:
|
||||
norms[name] = 0.0
|
||||
norms["total"] = total_sq.sqrt().item()
|
||||
return norms
|
||||
return total_sq.sqrt().item()
|
||||
|
||||
|
||||
def ctx_get_loss(ctx):
|
||||
@@ -52,24 +35,4 @@ def ctx_get_val_loss(ctx):
|
||||
|
||||
|
||||
def ctx_get_grad_norm(ctx):
|
||||
return grad_norm(ctx.model)
|
||||
|
||||
|
||||
def ctx_get_grad_std(ctx):
|
||||
return grad_std(ctx.model)
|
||||
|
||||
|
||||
def ctx_get_grad_max(ctx):
|
||||
return grad_max(ctx.model)
|
||||
|
||||
|
||||
def ctx_get_grad_min(ctx):
|
||||
return grad_min(ctx.model)
|
||||
|
||||
|
||||
def ctx_get_grad_mean(ctx):
|
||||
return grad_mean(ctx.model)
|
||||
|
||||
|
||||
def ctx_get_grad_nan_num(ctx):
|
||||
return grad_nan_num(ctx.model)
|
||||
return ctx.grad_norm
|
||||
|
||||
@@ -1,143 +0,0 @@
|
||||
import torch
|
||||
from torch.optim import Optimizer
|
||||
|
||||
|
||||
def _zeropower_via_newtonschulz(G: torch.Tensor, steps: int = 5):
|
||||
assert G.ndim == 2
|
||||
X = G
|
||||
scale = max(1, G.size(0) / G.size(1)) ** 0.5
|
||||
X = X / (X.norm() + 1e-7) * scale
|
||||
if steps == 0:
|
||||
return X
|
||||
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
|
||||
|
||||
|
||||
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:
|
||||
params_2d, params_1d = [], []
|
||||
grads_2d, grads_1d = [], []
|
||||
|
||||
for p in group["params"]:
|
||||
if p.grad is None:
|
||||
continue
|
||||
if p.grad.is_sparse:
|
||||
raise RuntimeError("Muon does not support sparse gradients")
|
||||
if p.ndim >= 2:
|
||||
params_2d.append(p)
|
||||
grads_2d.append(p.grad)
|
||||
else:
|
||||
params_1d.append(p)
|
||||
grads_1d.append(p.grad)
|
||||
|
||||
if params_2d:
|
||||
self._muon_update_foreach(params_2d, grads_2d, group)
|
||||
if params_1d:
|
||||
self._adamw_update_foreach(params_1d, grads_1d, group)
|
||||
|
||||
return loss
|
||||
|
||||
def _muon_update_foreach(self, params_2d, grads_2d, group):
|
||||
lr = group["lr"]
|
||||
momentum = group["momentum"]
|
||||
wd = group["weight_decay"]
|
||||
nesterov = group["nesterov"]
|
||||
ns_steps = group["ns_steps"]
|
||||
|
||||
if wd != 0:
|
||||
torch._foreach_mul_(params_2d, 1 - lr * wd)
|
||||
|
||||
if nesterov:
|
||||
grads_2d = torch._foreach_add(grads_2d, params_2d, alpha=wd)
|
||||
|
||||
bufs = []
|
||||
for p, grad in zip(params_2d, grads_2d):
|
||||
state = self.state[p]
|
||||
if "momentum_buffer" not in state:
|
||||
state["momentum_buffer"] = torch.zeros_like(grad)
|
||||
bufs.append(state["momentum_buffer"])
|
||||
|
||||
torch._foreach_lerp_(bufs, grads_2d, 1 - momentum)
|
||||
|
||||
for p, buf in zip(params_2d, bufs):
|
||||
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_foreach(self, params_1d, grads_1d, group):
|
||||
lr = group["adamw_lr"]
|
||||
betas = group["adamw_betas"]
|
||||
eps = group["adamw_eps"]
|
||||
wd = group["adamw_wd"]
|
||||
|
||||
steps: list[int] = []
|
||||
exp_avgs, exp_avg_sqs = [], []
|
||||
has_state = []
|
||||
for p in params_1d:
|
||||
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)
|
||||
has_state.append(False)
|
||||
else:
|
||||
has_state.append(True)
|
||||
state["step"] += 1
|
||||
steps.append(state["step"])
|
||||
exp_avgs.append(state["exp_avg"])
|
||||
exp_avg_sqs.append(state["exp_avg_sq"])
|
||||
|
||||
beta1, beta2 = betas
|
||||
|
||||
torch._foreach_lerp_(exp_avgs, grads_1d, 1 - beta1)
|
||||
grads_sq = torch._foreach_mul(grads_1d, grads_1d)
|
||||
torch._foreach_lerp_(exp_avg_sqs, grads_sq, 1 - beta2)
|
||||
|
||||
bias_correction1 = [1 - beta1**s for s in steps]
|
||||
bias_correction2 = [1 - beta2**s for s in steps]
|
||||
|
||||
if wd != 0:
|
||||
torch._foreach_mul_(params_1d, 1 - lr * wd)
|
||||
|
||||
exp_avg_corrected = torch._foreach_div(exp_avgs, bias_correction1)
|
||||
denom = torch._foreach_div(exp_avg_sqs, bias_correction2)
|
||||
denom = torch._foreach_sqrt(denom)
|
||||
torch._foreach_add_(denom, eps)
|
||||
torch._foreach_addcdiv_(params_1d, exp_avg_corrected, denom, value=-lr)
|
||||
+75
-34
@@ -2,7 +2,7 @@
|
||||
|
||||
import math
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, List, Type
|
||||
from typing import Any, Dict, List
|
||||
|
||||
from torch.optim.lr_scheduler import LRScheduler
|
||||
|
||||
@@ -31,7 +31,6 @@ class SchedulerFactory(BaseFactory["BaseScheduler"]):
|
||||
"""Factory class for creating learning rate schedulers.
|
||||
|
||||
Supports decorator-based registration for extensible scheduler types.
|
||||
Also supports creation from ScheduleConfig objects.
|
||||
|
||||
Example usage:
|
||||
@SchedulerFactory.register("custom")
|
||||
@@ -41,33 +40,6 @@ class SchedulerFactory(BaseFactory["BaseScheduler"]):
|
||||
scheduler = SchedulerFactory.create("custom", optimizer, **kwargs)
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def _validate_component(cls, scheduler_cls: Type[BaseScheduler]):
|
||||
"""Validate that the scheduler class inherits from BaseScheduler."""
|
||||
if not issubclass(scheduler_cls, BaseScheduler):
|
||||
raise TypeError(f"{scheduler_cls.__name__} must inherit from BaseScheduler")
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
cls, optimizer, schedule_type: str = "none", **kwargs
|
||||
) -> "BaseScheduler":
|
||||
"""Create a scheduler instance by type name.
|
||||
|
||||
Args:
|
||||
optimizer: PyTorch optimizer
|
||||
schedule_type: Type of scheduler ("cosine", "sgdr")
|
||||
**kwargs: Arguments passed to the scheduler constructor
|
||||
|
||||
Returns:
|
||||
Scheduler instance
|
||||
"""
|
||||
return super().create(schedule_type, optimizer, **kwargs)
|
||||
|
||||
@classmethod
|
||||
def available_types(cls) -> list:
|
||||
"""Return list of registered scheduler type names."""
|
||||
return cls.list_registered()
|
||||
|
||||
|
||||
# ----------- Scheduler implementations -----------
|
||||
|
||||
@@ -81,7 +53,7 @@ class CosineScheduler(BaseScheduler):
|
||||
optimizer,
|
||||
warmup_steps: int,
|
||||
lr_decay_steps: int,
|
||||
min_rate: float = 0.05,
|
||||
min_rate: float = 0.01,
|
||||
last_epoch: int = -1,
|
||||
):
|
||||
self.warmup_steps = warmup_steps
|
||||
@@ -93,11 +65,15 @@ class CosineScheduler(BaseScheduler):
|
||||
def get_lr(self) -> List[float]:
|
||||
# warmup
|
||||
if self.last_epoch < self.warmup_steps:
|
||||
warmup_factor = max(self.min_rate, self.last_epoch / self.warmup_steps)
|
||||
warmup_factor = max(
|
||||
self.min_rate, self.last_epoch / max(self.warmup_steps, 1)
|
||||
)
|
||||
return [base_lr * warmup_factor for base_lr in self.base_lrs]
|
||||
|
||||
# cosine decay
|
||||
decay_progress = (self.last_epoch - self.warmup_steps) / self.lr_decay_steps
|
||||
decay_progress = (self.last_epoch - self.warmup_steps) / max(
|
||||
self.lr_decay_steps, 1
|
||||
)
|
||||
decay_progress = min(decay_progress, 1.0)
|
||||
cosine_decay = 0.5 * (1.0 + math.cos(math.pi * decay_progress))
|
||||
decay_factor = max(self.min_rate, cosine_decay)
|
||||
@@ -132,7 +108,7 @@ class SGDRScheduler(BaseScheduler):
|
||||
optimizer,
|
||||
warmup_steps: int,
|
||||
cycle_length: int,
|
||||
min_rate: float = 0.05,
|
||||
min_rate: float = 0.01,
|
||||
t_mult: int = 2,
|
||||
last_epoch: int = -1,
|
||||
):
|
||||
@@ -146,7 +122,9 @@ class SGDRScheduler(BaseScheduler):
|
||||
def get_lr(self):
|
||||
# warmup
|
||||
if self.last_epoch < self.warmup_steps:
|
||||
warmup_factor = max(self.min_rate, self.last_epoch / self.warmup_steps)
|
||||
warmup_factor = max(
|
||||
self.min_rate, self.last_epoch / max(self.warmup_steps, 1)
|
||||
)
|
||||
return [base_lr * warmup_factor for base_lr in self.base_lrs]
|
||||
|
||||
# SGDR
|
||||
@@ -192,3 +170,66 @@ class SGDRScheduler(BaseScheduler):
|
||||
self.min_rate = state_dict.pop("min_rate")
|
||||
self.t_mult = state_dict.pop("t_mult")
|
||||
super().load_state_dict(state_dict)
|
||||
|
||||
|
||||
@SchedulerFactory.register("wsd")
|
||||
class WSDScheduler(BaseScheduler):
|
||||
"""WSD (Warmup-Stable-Decay) scheduler with sqrt cooldown.
|
||||
|
||||
warmup_steps: linear warmup from min_rate to 1.0
|
||||
stable_steps: constant at base_lr
|
||||
decay_steps: sqrt decay from base_lr to min_rate
|
||||
min_rate: minimum lr as fraction of base_lr (default 0.0)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
optimizer,
|
||||
warmup_steps: int,
|
||||
stable_steps: int,
|
||||
decay_steps: int,
|
||||
min_rate: float = 0.01,
|
||||
last_epoch: int = -1,
|
||||
):
|
||||
self.warmup_steps = warmup_steps
|
||||
self.stable_steps = stable_steps
|
||||
self.decay_steps = decay_steps
|
||||
self.min_rate = min_rate
|
||||
self.total_steps = warmup_steps + stable_steps + decay_steps
|
||||
super().__init__(optimizer, last_epoch)
|
||||
|
||||
def get_lr(self) -> List[float]:
|
||||
if self.last_epoch < self.warmup_steps:
|
||||
factor = max(self.min_rate, self.last_epoch / max(self.warmup_steps, 1))
|
||||
return [base_lr * factor for base_lr in self.base_lrs]
|
||||
|
||||
offset = self.last_epoch - self.warmup_steps
|
||||
|
||||
if offset < self.stable_steps:
|
||||
return list(self.base_lrs)
|
||||
|
||||
decay_ratio = (offset - self.stable_steps) / max(self.decay_steps, 1)
|
||||
decay_ratio = min(decay_ratio, 1.0)
|
||||
factor = (1.0 - self.min_rate) * (1.0 - decay_ratio) ** 2 + self.min_rate
|
||||
return [base_lr * factor for base_lr in self.base_lrs]
|
||||
|
||||
def state_dict(self):
|
||||
state = super().state_dict()
|
||||
state.update(
|
||||
{
|
||||
"warmup_steps": self.warmup_steps,
|
||||
"stable_steps": self.stable_steps,
|
||||
"decay_steps": self.decay_steps,
|
||||
"min_rate": self.min_rate,
|
||||
"total_steps": self.total_steps,
|
||||
}
|
||||
)
|
||||
return state
|
||||
|
||||
def load_state_dict(self, state_dict):
|
||||
self.warmup_steps = state_dict.pop("warmup_steps")
|
||||
self.stable_steps = state_dict.pop("stable_steps")
|
||||
self.decay_steps = state_dict.pop("decay_steps")
|
||||
self.min_rate = state_dict.pop("min_rate")
|
||||
self.total_steps = state_dict.pop("total_steps")
|
||||
super().load_state_dict(state_dict)
|
||||
|
||||
+117
-88
@@ -1,41 +1,28 @@
|
||||
"""Training strategy implementations with factory pattern."""
|
||||
|
||||
import copy
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Callable, Dict, Union
|
||||
from typing import Callable, Dict, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
|
||||
def unwrap_model(model: nn.Module) -> nn.Module:
|
||||
if isinstance(model, DDP):
|
||||
return model.module
|
||||
if isinstance(model, FSDP):
|
||||
return model._fsdp_wrapped_module
|
||||
return model
|
||||
|
||||
|
||||
def create_ref_model(model: nn.Module) -> nn.Module:
|
||||
"""Create a reference model for DPO/GRPO training.
|
||||
|
||||
Handles DDP-wrapped models safely by unwrapping first,
|
||||
then creating a deep copy with frozen gradients.
|
||||
"""
|
||||
original_model = unwrap_model(model)
|
||||
ref_model = copy.deepcopy(original_model)
|
||||
def create_ref_model(
|
||||
model_fn: Callable[[], nn.Module], state_dict: Dict[str, Tensor]
|
||||
) -> nn.Module:
|
||||
"""Create a frozen reference model from model_fn + full state dict."""
|
||||
ref_model = model_fn()
|
||||
ref_model.load_state_dict(state_dict)
|
||||
ref_model.requires_grad_(False)
|
||||
ref_model.eval()
|
||||
return ref_model
|
||||
|
||||
|
||||
def move_to_device(batch: Dict[str, Tensor], device: str) -> Any:
|
||||
def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
|
||||
"""Move batch tensors to specified device with non-blocking transfer."""
|
||||
return {key: value.to(device, non_blocking=True) for key, value in batch.items()}
|
||||
|
||||
@@ -45,7 +32,7 @@ def get_logprobs(
|
||||
input_ids: Tensor,
|
||||
mask: Tensor,
|
||||
reduction: str,
|
||||
):
|
||||
) -> Tensor:
|
||||
"""Compute token-wise log probabilities from model outputs.
|
||||
|
||||
Args:
|
||||
@@ -83,14 +70,34 @@ def get_logprobs(
|
||||
return token_logprobs * shifted_mask
|
||||
|
||||
|
||||
def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
|
||||
S = position_ids.size(1)
|
||||
device = position_ids.device
|
||||
boundaries = position_ids[:, 1:] <= position_ids[:, :-1]
|
||||
doc_ids = torch.cat(
|
||||
[
|
||||
torch.zeros(position_ids.size(0), 1, dtype=torch.long, device=device),
|
||||
boundaries.long().cumsum(dim=1),
|
||||
],
|
||||
dim=1,
|
||||
)
|
||||
same_doc = doc_ids.unsqueeze(-1) == doc_ids.unsqueeze(-2)
|
||||
causal = torch.tril(torch.ones(S, S, dtype=torch.bool, device=device))
|
||||
return (same_doc & causal).unsqueeze(1)
|
||||
|
||||
|
||||
class BaseStrategy(ABC):
|
||||
"""Abstract base class for training strategies."""
|
||||
|
||||
def __init__(
|
||||
self, model: Union[Callable[..., Dict[str, Tensor]]], device: str, **kwargs
|
||||
self,
|
||||
model: Union[nn.Module, Callable[..., Dict[str, Tensor]]],
|
||||
device: str,
|
||||
**kwargs,
|
||||
):
|
||||
self.model = model
|
||||
self.device = device
|
||||
self.executor = kwargs.pop("executor", None)
|
||||
self.extra_kwargs = kwargs
|
||||
|
||||
@abstractmethod
|
||||
@@ -124,32 +131,6 @@ class StrategyFactory(BaseFactory["BaseStrategy"]):
|
||||
strategy = StrategyFactory.create("custom", model, device)
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def _validate_component(cls, strategy_cls: type):
|
||||
"""Validate that the strategy class inherits from BaseStrategy."""
|
||||
if not issubclass(strategy_cls, BaseStrategy):
|
||||
raise TypeError(f"{strategy_cls.__name__} must inherit from BaseStrategy")
|
||||
|
||||
@classmethod
|
||||
def create(cls, train_type: str, model, device: str, **kwargs) -> "BaseStrategy":
|
||||
"""Create a strategy instance based on training type.
|
||||
|
||||
Args:
|
||||
train_type: Type of training ("seq", "sft", "dpo", "grpo")
|
||||
model: Model instance for the strategy
|
||||
device: Device to run the strategy on
|
||||
**kwargs: Additional arguments passed to strategy constructor
|
||||
|
||||
Returns:
|
||||
Strategy instance
|
||||
"""
|
||||
return super().create(train_type, model, device, **kwargs)
|
||||
|
||||
@classmethod
|
||||
def available_strategies(cls) -> list:
|
||||
"""Return list of registered strategy names."""
|
||||
return cls.list_registered()
|
||||
|
||||
|
||||
# ============== Strategy Classes ==============
|
||||
# All strategies are registered at class definition time using the decorator
|
||||
@@ -162,7 +143,13 @@ class SEQStrategy(BaseStrategy):
|
||||
Computes cross-entropy loss for next token prediction.
|
||||
"""
|
||||
|
||||
def __init__(self, model, device, label_smoothing: float = 0.0, **kwargs):
|
||||
def __init__(
|
||||
self,
|
||||
model: Union[nn.Module, Callable[..., Dict[str, Tensor]]],
|
||||
device: str,
|
||||
label_smoothing: float = 0.0,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(model, device, **kwargs)
|
||||
self.label_smoothing = label_smoothing
|
||||
|
||||
@@ -187,21 +174,31 @@ class SFTStrategy(BaseStrategy):
|
||||
Applies cross-entropy loss only to tokens where loss_mask is True.
|
||||
"""
|
||||
|
||||
def __init__(self, model, device, label_smoothing: float = 0.0, **kwargs):
|
||||
def __init__(
|
||||
self,
|
||||
model: Union[nn.Module, Callable[..., Dict[str, Tensor]]],
|
||||
device: str,
|
||||
label_smoothing: float = 0.0,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(model, device, **kwargs)
|
||||
self.label_smoothing = label_smoothing
|
||||
|
||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||
batch = move_to_device(batch, self.device)
|
||||
input_ids, target_ids, loss_mask = (
|
||||
input_ids, target_ids, position_ids, loss_mask = (
|
||||
batch["input_ids"],
|
||||
batch["target_ids"],
|
||||
batch["position_ids"],
|
||||
batch["loss_mask"],
|
||||
)
|
||||
|
||||
ignore_index = -100
|
||||
logits = self.model(input_ids=input_ids)["logits"]
|
||||
target_ids = target_ids.masked_fill(loss_mask == 0, ignore_index)
|
||||
input_mask = make_doc_boundary_mask(position_ids)
|
||||
target_ids = target_ids.masked_fill(~loss_mask, ignore_index)
|
||||
logits = self.model(
|
||||
input_ids=input_ids, position_ids=position_ids, input_mask=input_mask
|
||||
)["logits"]
|
||||
|
||||
loss = F.cross_entropy(
|
||||
input=logits.flatten(0, 1).float(),
|
||||
@@ -225,12 +222,13 @@ class DPOStrategy(BaseStrategy):
|
||||
self,
|
||||
model: nn.Module,
|
||||
device: str,
|
||||
ref_model: nn.Module,
|
||||
beta: float = 0.1,
|
||||
reduction: str = "mean",
|
||||
reduction: str = "sum",
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(model, device, **kwargs)
|
||||
self.ref_model = create_ref_model(model)
|
||||
self.ref_model = ref_model
|
||||
self.beta = beta
|
||||
self.reduction = reduction
|
||||
|
||||
@@ -267,41 +265,45 @@ class DPOStrategy(BaseStrategy):
|
||||
class GRPOStrategy(BaseStrategy):
|
||||
"""Group Relative Policy Optimization strategy.
|
||||
|
||||
On-policy GRPO following DeepSeek-R1: the policy model is updated while
|
||||
a frozen ref_model stores the old-policy log-probs. ratio = exp(logπ_θ - logπ_ref),
|
||||
clipped PPO objective. Call ``sync_ref_model()`` after each data-generation round.
|
||||
Implements GRPO following DeepSeek-R1 with token-level PPO clipping.
|
||||
Advantages are group-normalized from scalar per-response rewards and
|
||||
broadcast across all response tokens. The loss is computed **only on
|
||||
response tokens** — prompt tokens are masked out.
|
||||
|
||||
Three model roles are distinguished:
|
||||
|
||||
* **Policy** ``self.model`` — the model being trained.
|
||||
* **Old policy** ``self.old_model`` — the behaviour policy that generated
|
||||
the responses. Used for the importance sampling ratio
|
||||
``ρ = π_θ / π_old``. Synced externally after each data-generation round.
|
||||
* **Reference model** ``self.ref_model`` — a frozen copy of the initial
|
||||
policy (typically the SFT checkpoint) used **only** for the KL
|
||||
regularisation term. It is never updated during training.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: nn.Module,
|
||||
device: str,
|
||||
old_model: nn.Module,
|
||||
ref_model: nn.Module,
|
||||
clip_eps: float = 0.2,
|
||||
kl_coef: float = 0.01,
|
||||
group_size: int = 4,
|
||||
reduction: str = "mean",
|
||||
sync_interval: int = 200,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(model, device, **kwargs)
|
||||
self.ref_model = create_ref_model(model)
|
||||
self.old_model = old_model
|
||||
self.ref_model = ref_model
|
||||
self.clip_eps = clip_eps
|
||||
self.kl_coef = kl_coef
|
||||
self.group_size = group_size
|
||||
self.reduction = reduction
|
||||
self.sync_interval = sync_interval
|
||||
self._step = 0
|
||||
|
||||
def sync_ref_model(self):
|
||||
"""Copy current model weights to ref model."""
|
||||
ref_state = self.model.state_dict()
|
||||
self.ref_model.load_state_dict(ref_state)
|
||||
def sync_old_model(self):
|
||||
"""Copy current policy weights to old model."""
|
||||
self.old_model.load_state_dict(self.executor.unwrap_model(self.model))
|
||||
|
||||
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
|
||||
self._step += 1
|
||||
if self._step % self.sync_interval == 0:
|
||||
self.sync_ref_model()
|
||||
|
||||
batch = move_to_device(batch, self.device)
|
||||
prompts = batch["prompts"]
|
||||
responses = batch["responses"]
|
||||
@@ -312,33 +314,60 @@ class GRPOStrategy(BaseStrategy):
|
||||
responses_flat = responses.view(-1, response_len)
|
||||
masks_flat = masks.view(-1, response_len)
|
||||
prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1)
|
||||
prompt_len = prompt_expanded.size(1)
|
||||
|
||||
full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1)
|
||||
full_masks = torch.cat([torch.ones_like(prompt_expanded), masks_flat], dim=-1)
|
||||
|
||||
log_probs_policy = get_logprobs(
|
||||
self.model, full_sequences, full_masks, self.reduction
|
||||
)
|
||||
log_probs_policy = log_probs_policy.view(batch_size, group_size)
|
||||
# Prompt tokens are masked out (0) so logprobs are computed only for
|
||||
# response tokens. get_logprobs shifts the mask by one position, so
|
||||
# the first response token's logprob (predicted from the last prompt
|
||||
# token) is correctly included.
|
||||
full_masks = torch.cat([torch.zeros_like(prompt_expanded), masks_flat], dim=-1)
|
||||
|
||||
# get_logprobs returns [B*G, S-1] (S = prompt_len + response_len).
|
||||
# Response token logprobs occupy the last ``response_len`` positions
|
||||
# (the first response token is predicted from the last prompt token).
|
||||
token_log_probs_policy = get_logprobs(
|
||||
self.model, full_sequences, full_masks, "none"
|
||||
)[:, prompt_len - 1 :]
|
||||
with torch.no_grad():
|
||||
log_probs_ref = get_logprobs(
|
||||
self.ref_model, full_sequences, full_masks, self.reduction
|
||||
)
|
||||
log_probs_ref = log_probs_ref.view(batch_size, group_size)
|
||||
token_log_probs_old = get_logprobs(
|
||||
self.old_model, full_sequences, full_masks, "none"
|
||||
)[:, prompt_len - 1 :]
|
||||
token_log_probs_ref = get_logprobs(
|
||||
self.ref_model, full_sequences, full_masks, "none"
|
||||
)[:, prompt_len - 1 :]
|
||||
|
||||
eps = torch.finfo(log_probs_policy.dtype).eps
|
||||
# Reshape to [B, G, response_len]
|
||||
token_log_probs_policy = token_log_probs_policy.view(batch_size, group_size, -1)
|
||||
token_log_probs_old = token_log_probs_old.view(batch_size, group_size, -1)
|
||||
token_log_probs_ref = token_log_probs_ref.view(batch_size, group_size, -1)
|
||||
token_masks = masks_flat.view(batch_size, group_size, -1).float()
|
||||
|
||||
# Group-normalized advantages from scalar per-response rewards.
|
||||
eps = 1e-8
|
||||
mean = rewards.mean(dim=-1, keepdim=True)
|
||||
std = rewards.std(dim=-1, keepdim=True)
|
||||
std = rewards.std(dim=-1, keepdim=True, unbiased=False)
|
||||
advantages = (rewards - mean) / (std + eps)
|
||||
# Broadcast scalar advantage to every response token: [B, G, 1]
|
||||
advantages = advantages.unsqueeze(-1)
|
||||
|
||||
ratio = torch.exp(log_probs_policy - log_probs_ref)
|
||||
# Token-level ratio (π_θ / π_old) and PPO clipping.
|
||||
log_ratio = token_log_probs_policy - token_log_probs_old
|
||||
ratio = torch.exp(log_ratio)
|
||||
|
||||
surr1 = ratio * advantages
|
||||
surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * advantages
|
||||
per_token_policy_loss = -torch.min(surr1, surr2)
|
||||
token_count = token_masks.sum().clamp(min=1.0)
|
||||
policy_loss = (per_token_policy_loss * token_masks).sum() / token_count
|
||||
|
||||
# KL penalty to frozen reference model with k1 estimator (non-negative):
|
||||
# k1 = π_ref / π_θ - log(π_ref / π_θ) - 1, where π_ref / π_θ = exp(log_ref - log_policy).
|
||||
log_ref_ratio = token_log_probs_ref - token_log_probs_policy
|
||||
r = torch.exp(log_ref_ratio)
|
||||
kl_per_token = r - torch.log(r + eps) - 1.0
|
||||
kl_penalty = self.kl_coef * (kl_per_token * token_masks).sum() / token_count
|
||||
|
||||
policy_loss = -torch.min(surr1, surr2).mean()
|
||||
kl_penalty = self.kl_coef * (log_probs_policy - log_probs_ref).square().mean()
|
||||
total_loss = policy_loss + kl_penalty
|
||||
|
||||
return total_loss
|
||||
|
||||
@@ -9,21 +9,15 @@ from typing import IO, Callable, List, Optional, Protocol, runtime_checkable
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
from torch.nn.utils import clip_grad_norm_
|
||||
from torch.utils.checkpoint import checkpoint as torch_checkpoint
|
||||
from tqdm import tqdm
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.parallel import only_on_rank
|
||||
from astrai.parallel.setup import get_current_device, get_rank
|
||||
from astrai.parallel.setup import get_current_device
|
||||
from astrai.serialization import Checkpoint
|
||||
from astrai.trainer.metric_util import (
|
||||
ctx_get_grad_max,
|
||||
ctx_get_grad_mean,
|
||||
ctx_get_grad_min,
|
||||
ctx_get_grad_nan_num,
|
||||
ctx_get_grad_norm,
|
||||
ctx_get_grad_std,
|
||||
ctx_get_loss,
|
||||
ctx_get_lr,
|
||||
ctx_get_val_loss,
|
||||
@@ -86,7 +80,9 @@ class GradientClippingCallback(TrainCallback):
|
||||
self.max_grad_norm = max_grad_norm
|
||||
|
||||
def on_optimizer_step(self, context: TrainContext):
|
||||
clip_grad_norm_(context.model.parameters(), self.max_grad_norm)
|
||||
context.grad_norm = context.executor.clip_grad_norm(
|
||||
context.model, self.max_grad_norm
|
||||
)
|
||||
|
||||
|
||||
@CallbackFactory.register("gradient_checkpointing")
|
||||
@@ -143,33 +139,38 @@ class CheckpointCallback(TrainCallback):
|
||||
self.interval = interval
|
||||
self.weight_only = weight_only
|
||||
self.save_extra_fn = save_extra_fn or CheckpointCallback.save_extra
|
||||
self.last_ckpt_iter = 0
|
||||
self.last_ckpt_step = None
|
||||
|
||||
def on_train_begin(self, context: TrainContext):
|
||||
self.last_ckpt_step = context.optimizer_step
|
||||
|
||||
def _save_checkpoint(self, context: TrainContext):
|
||||
unwrapped = context.executor.unwrap_model(context.model)
|
||||
state_dict = unwrapped.state_dict()
|
||||
self.last_ckpt_iter = context.iteration
|
||||
self.last_ckpt_step = context.optimizer_step
|
||||
|
||||
if get_rank() == 0:
|
||||
save_path = os.path.join(
|
||||
self.save_dir, f"epoch_{context.epoch}_iter_{context.iteration}"
|
||||
)
|
||||
extra = self.save_extra_fn(context)
|
||||
context.checkpoint = Checkpoint(
|
||||
state_dict=state_dict,
|
||||
epoch=context.epoch,
|
||||
iteration=context.iteration,
|
||||
extra=extra,
|
||||
config=context.model_config,
|
||||
)
|
||||
context.checkpoint.save(save_path)
|
||||
with context.executor.checkpoint_context(context.model) as state_dict:
|
||||
if state_dict is not None:
|
||||
save_path = os.path.join(
|
||||
self.save_dir,
|
||||
f"epoch_{context.epoch}_step_{context.optimizer_step}",
|
||||
)
|
||||
extra = self.save_extra_fn(context)
|
||||
meta = context.config.to_dict()
|
||||
context.checkpoint = Checkpoint(
|
||||
state_dict=state_dict,
|
||||
epoch=context.epoch,
|
||||
consumed_samples=context.consumed_samples,
|
||||
config=context.model_config,
|
||||
extra=extra,
|
||||
meta=meta,
|
||||
)
|
||||
context.checkpoint.save(save_path)
|
||||
|
||||
def on_batch_end(self, context: TrainContext):
|
||||
if context.iteration - self.last_ckpt_iter >= self.interval:
|
||||
if context.optimizer_step - self.last_ckpt_step >= self.interval:
|
||||
self._save_checkpoint(context)
|
||||
|
||||
def on_train_end(self, context: TrainContext):
|
||||
if context.iteration != self.last_ckpt_iter:
|
||||
if context.optimizer_step != self.last_ckpt_step:
|
||||
self._save_checkpoint(context)
|
||||
|
||||
def on_error(self, context: TrainContext):
|
||||
@@ -201,20 +202,24 @@ class ProgressBarCallback(TrainCallback):
|
||||
|
||||
@only_on_rank(0)
|
||||
def on_epoch_begin(self, context: TrainContext):
|
||||
total_steps = len(context.dataloader) // context.executor.grad_accum_steps
|
||||
self.progress_bar = tqdm(
|
||||
context.dataloader,
|
||||
total=total_steps,
|
||||
desc=f"Epoch {context.epoch + 1}/{self.num_epoch}",
|
||||
dynamic_ncols=True,
|
||||
file=self.file or sys.stdout,
|
||||
)
|
||||
|
||||
@only_on_rank(0)
|
||||
def on_batch_end(self, context: TrainContext):
|
||||
def on_optimizer_step(self, context: TrainContext):
|
||||
postfix = {
|
||||
"step": f"{context.optimizer_step:d}",
|
||||
"loss": f"{context.loss:.4f}",
|
||||
"lr": f"{context.optimizer.param_groups[-1]['lr']:.2e}",
|
||||
}
|
||||
if context.val_loss > 0:
|
||||
if context.grad_norm is not None:
|
||||
postfix["grad_norm"] = f"{context.grad_norm:.2f}"
|
||||
if context.val_loss is not None:
|
||||
postfix["val_loss"] = f"{context.val_loss:.4f}"
|
||||
self.progress_bar.set_postfix(postfix)
|
||||
self.progress_bar.update(1)
|
||||
@@ -226,19 +231,20 @@ class ProgressBarCallback(TrainCallback):
|
||||
self.progress_bar.close()
|
||||
|
||||
|
||||
@CallbackFactory.register("metric_logger")
|
||||
class MetricLoggerCallback(TrainCallback):
|
||||
@CallbackFactory.register("metric")
|
||||
class MetricCallback(TrainCallback):
|
||||
def __init__(
|
||||
self,
|
||||
log_dir: str,
|
||||
save_interval: int,
|
||||
log_interval: int = 10,
|
||||
metrics: List[str] = None,
|
||||
val_step: int = 0,
|
||||
):
|
||||
self.last_log_iter = 0
|
||||
self.last_log_flush_step = None
|
||||
self.save_interval = save_interval
|
||||
self.log_interval = log_interval
|
||||
self.metrics = metrics or ["loss", "lr"]
|
||||
self.val_step = val_step
|
||||
self._next_val_step = 0
|
||||
|
||||
self.log_dir = Path(log_dir) if log_dir else Path.cwd() / "logs"
|
||||
self.log_dir.mkdir(parents=True, exist_ok=True)
|
||||
@@ -250,53 +256,28 @@ class MetricLoggerCallback(TrainCallback):
|
||||
"lr": ctx_get_lr,
|
||||
"val_loss": ctx_get_val_loss,
|
||||
"grad_norm": ctx_get_grad_norm,
|
||||
"grad_std": ctx_get_grad_std,
|
||||
"grad_max": ctx_get_grad_max,
|
||||
"grad_min": ctx_get_grad_min,
|
||||
"grad_mean": ctx_get_grad_mean,
|
||||
"grad_nan_num": ctx_get_grad_nan_num,
|
||||
}
|
||||
|
||||
def _get_log_data(self, context: TrainContext):
|
||||
def _metrics(self, context: TrainContext, names):
|
||||
return {
|
||||
m: self._metric_funcs[m](context)
|
||||
for m in names
|
||||
if self._metric_funcs[m](context) is not None
|
||||
}
|
||||
|
||||
@only_on_rank(0)
|
||||
def _append(self, event_type: str, context: TrainContext, **extra):
|
||||
entry = {
|
||||
"type": event_type,
|
||||
"timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"),
|
||||
"epoch": context.epoch,
|
||||
"iter": context.iteration,
|
||||
**{m: self._metric_funcs[m](context) for m in self.metrics},
|
||||
"step": context.optimizer_step,
|
||||
"consumed_samples": context.consumed_samples,
|
||||
**extra,
|
||||
}
|
||||
self.log_cache.append(entry)
|
||||
|
||||
@only_on_rank(0)
|
||||
def _add_log(self, log_data):
|
||||
self.log_cache.append(log_data)
|
||||
|
||||
@only_on_rank(0)
|
||||
def _save_log(self, epoch, iter):
|
||||
log_file = self.log_dir / f"epoch_{epoch}_iter_{iter}_metric.jsonl"
|
||||
|
||||
with open(log_file, "w") as f:
|
||||
for log in self.log_cache:
|
||||
f.write(json.dumps(log) + "\n")
|
||||
|
||||
def on_batch_end(self, context):
|
||||
if context.iteration % self.log_interval == 0:
|
||||
log_data = self._get_log_data(context)
|
||||
self._add_log(log_data)
|
||||
|
||||
if context.iteration - self.last_log_iter >= self.save_interval:
|
||||
self._save_log(context.epoch, context.iteration)
|
||||
self.last_log_iter = context.iteration
|
||||
|
||||
def on_train_end(self, context):
|
||||
if context.iteration != self.last_log_iter:
|
||||
self._save_log(context.epoch, context.iteration)
|
||||
|
||||
def on_error(self, context):
|
||||
self._save_log(context.epoch, context.iteration)
|
||||
|
||||
|
||||
@CallbackFactory.register("validation")
|
||||
class ValidationCallback(TrainCallback):
|
||||
def _run_validation(self, context: TrainContext):
|
||||
def _run_validation(self, context: TrainContext) -> float:
|
||||
context.model.eval()
|
||||
|
||||
total_loss = 0.0
|
||||
@@ -308,27 +289,56 @@ class ValidationCallback(TrainCallback):
|
||||
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()
|
||||
stats = torch.tensor(
|
||||
[total_loss, float(num_batches)], device=get_current_device()
|
||||
)
|
||||
dist.all_reduce(stats, op=dist.ReduceOp.SUM)
|
||||
avg_loss = (stats[0] / stats[1]).item()
|
||||
else:
|
||||
avg_loss = total_loss / max(num_batches, 1)
|
||||
|
||||
context.val_loss = avg_loss
|
||||
context.model.train()
|
||||
return avg_loss
|
||||
|
||||
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_train_begin(self, context: TrainContext):
|
||||
self.last_log_flush_step = context.optimizer_step
|
||||
|
||||
def on_optimizer_step(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)
|
||||
@only_on_rank(0)
|
||||
def _flush(self, epoch, step):
|
||||
log_file = self.log_dir / f"epoch_{epoch}_step_{step}_metric.jsonl"
|
||||
log_file.parent.mkdir(parents=True, exist_ok=True)
|
||||
with open(log_file, "w") as f:
|
||||
for log in self.log_cache:
|
||||
f.write(json.dumps(log) + "\n")
|
||||
|
||||
def on_optimizer_step(self, context):
|
||||
if (
|
||||
context.val_dataloader is not None
|
||||
and self.val_step > 0
|
||||
and context.optimizer_step >= self._next_val_step
|
||||
):
|
||||
context.val_loss = self._run_validation(context)
|
||||
self._next_val_step = context.optimizer_step + self.val_step
|
||||
self._append("validation", context, val_loss=context.val_loss)
|
||||
|
||||
step_metrics = [m for m in self.metrics if m != "val_loss"]
|
||||
self._append("step", context, **self._metrics(context, step_metrics))
|
||||
|
||||
if context.optimizer_step - self.last_log_flush_step >= self.save_interval:
|
||||
self._flush(context.epoch, context.optimizer_step)
|
||||
self.last_log_flush_step = context.optimizer_step
|
||||
|
||||
def on_epoch_end(self, context):
|
||||
self._append("epoch", context)
|
||||
|
||||
def on_train_end(self, context):
|
||||
if (
|
||||
self.last_log_flush_step is None
|
||||
or context.optimizer_step != self.last_log_flush_step
|
||||
):
|
||||
self._flush(context.epoch, context.optimizer_step)
|
||||
self.last_log_flush_step = context.optimizer_step
|
||||
|
||||
def on_error(self, context):
|
||||
self._flush(context.epoch, context.optimizer_step)
|
||||
|
||||
@@ -1,18 +1,19 @@
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Optional, Self
|
||||
from typing import Any, Dict, Optional, Self
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.utils.data import DataLoader
|
||||
from torch.utils.data import DataLoader, random_split
|
||||
|
||||
from astrai.config.train_config import TrainConfig
|
||||
from astrai.dataset import ResumableDistributedSampler
|
||||
from astrai.dataset import RDSampler
|
||||
from astrai.model.components.lora import inject_lora
|
||||
from astrai.parallel.executor import BaseExecutor, ExecutorFactory
|
||||
from astrai.parallel.setup import get_current_device, get_rank, get_world_size
|
||||
from astrai.protocols import OptimizerProtocol, SchedulerProtocol
|
||||
from astrai.serialization import Checkpoint, load_json, load_model_weights
|
||||
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
||||
from astrai.serialization import Checkpoint, load_json
|
||||
from astrai.trainer.strategy import BaseStrategy, StrategyFactory, create_ref_model
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -28,14 +29,23 @@ class TrainContext:
|
||||
executor: BaseExecutor = field(default=None)
|
||||
|
||||
epoch: int = field(default=0)
|
||||
iteration: int = field(default=0)
|
||||
consumed_samples: int = field(default=0)
|
||||
loss: float = field(default=0.0)
|
||||
val_dataloader: DataLoader = field(default=None)
|
||||
val_loss: float = field(default=0.0)
|
||||
grad_norm: Optional[float] = field(default=None)
|
||||
val_dataloader: Optional[DataLoader] = field(default=None)
|
||||
val_loss: Optional[float] = field(default=None)
|
||||
|
||||
world_size: int = field(default=1)
|
||||
rank: int = field(default=0)
|
||||
kwargs: dict = field(default_factory=dict)
|
||||
kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
@property
|
||||
def optimizer_step(self) -> int:
|
||||
return self.consumed_samples // (
|
||||
self.config.batch_per_device
|
||||
* self.world_size
|
||||
* self.config.grad_accum_steps
|
||||
)
|
||||
|
||||
|
||||
class TrainContextBuilder:
|
||||
@@ -44,10 +54,12 @@ class TrainContextBuilder:
|
||||
config: TrainConfig,
|
||||
):
|
||||
self.config = config
|
||||
self._resume_dir: Optional[str] = None
|
||||
self._param_path: Optional[str] = None
|
||||
self._resume: bool = False
|
||||
|
||||
def with_resume_dir(self, resume_dir: Optional[str]) -> Self:
|
||||
self._resume_dir = resume_dir
|
||||
def with_param_path(self, param_path: Optional[str], resume: bool = False) -> Self:
|
||||
self._param_path = param_path
|
||||
self._resume = resume
|
||||
return self
|
||||
|
||||
def build(self) -> TrainContext:
|
||||
@@ -64,8 +76,8 @@ class TrainContextBuilder:
|
||||
model = model.to(device=device)
|
||||
|
||||
model_config = {}
|
||||
if self._resume_dir:
|
||||
config_path = Path(self._resume_dir) / "config.json"
|
||||
if self._param_path:
|
||||
config_path = Path(self._param_path) / "config.json"
|
||||
if config_path.exists():
|
||||
model_config = load_json(config_path)
|
||||
|
||||
@@ -81,21 +93,29 @@ class TrainContextBuilder:
|
||||
executor=executor,
|
||||
)
|
||||
|
||||
if self._resume_dir is not None:
|
||||
resume_path = Path(self._resume_dir)
|
||||
if (resume_path / "meta.json").exists():
|
||||
checkpoint = Checkpoint.load(self._resume_dir)
|
||||
state_dict = checkpoint.state_dict
|
||||
if self._param_path:
|
||||
checkpoint = Checkpoint.load_any(self._param_path)
|
||||
if checkpoint is not None:
|
||||
model.load_state_dict(checkpoint.state_dict, strict=False)
|
||||
if checkpoint.config:
|
||||
context.model_config = checkpoint.config
|
||||
else:
|
||||
checkpoint = None
|
||||
state_dict = load_model_weights(self._resume_dir)
|
||||
model.load_state_dict(state_dict, strict=False)
|
||||
if checkpoint is not None:
|
||||
context.epoch = cfg.start_epoch
|
||||
context.iteration = cfg.start_batch
|
||||
context.checkpoint = checkpoint
|
||||
|
||||
if self._resume:
|
||||
context.epoch = checkpoint.epoch or cfg.start_epoch
|
||||
if checkpoint.consumed_samples > 0:
|
||||
per_step = (
|
||||
cfg.batch_per_device
|
||||
* context.world_size
|
||||
* cfg.grad_accum_steps
|
||||
)
|
||||
context.consumed_samples = (
|
||||
checkpoint.consumed_samples // per_step
|
||||
) * per_step
|
||||
else:
|
||||
context.consumed_samples = (
|
||||
cfg.start_samples * context.world_size
|
||||
)
|
||||
context.checkpoint = checkpoint
|
||||
|
||||
if cfg.lora is not None:
|
||||
inject_lora(
|
||||
@@ -108,37 +128,51 @@ class TrainContextBuilder:
|
||||
context.optimizer = cfg.optimizer_fn(model)
|
||||
context.scheduler = cfg.scheduler_fn(context.optimizer)
|
||||
|
||||
sampler_offset = context.iteration * cfg.batch_per_device
|
||||
sampler = ResumableDistributedSampler(
|
||||
data_source=cfg.dataset,
|
||||
train_dataset = cfg.dataset
|
||||
val_dataset = cfg.val_dataset
|
||||
|
||||
if val_dataset is None and cfg.val_split is not None:
|
||||
n_total = len(cfg.dataset)
|
||||
n_val = max(1, int(n_total * cfg.val_split))
|
||||
n_train = n_total - n_val
|
||||
generator = torch.Generator().manual_seed(cfg.random_seed)
|
||||
train_dataset, val_dataset = random_split(
|
||||
cfg.dataset, [n_train, n_val], generator=generator
|
||||
)
|
||||
|
||||
sampler_offset = context.consumed_samples // context.world_size
|
||||
sampler = RDSampler(
|
||||
data_source=train_dataset,
|
||||
start_epoch=context.epoch,
|
||||
start_iter=sampler_offset,
|
||||
seed=cfg.random_seed,
|
||||
)
|
||||
context.dataloader = DataLoader(
|
||||
cfg.dataset,
|
||||
train_dataset,
|
||||
batch_size=cfg.batch_per_device,
|
||||
sampler=sampler,
|
||||
num_workers=cfg.num_workers,
|
||||
pin_memory=cfg.pin_memory,
|
||||
prefetch_factor=cfg.prefetch_factor,
|
||||
collate_fn=cfg.collate_fn,
|
||||
)
|
||||
|
||||
if cfg.val_dataset is not None:
|
||||
val_sampler = ResumableDistributedSampler(
|
||||
data_source=cfg.val_dataset,
|
||||
if val_dataset is not None:
|
||||
val_sampler = RDSampler(
|
||||
data_source=val_dataset,
|
||||
start_epoch=0,
|
||||
start_iter=0,
|
||||
seed=cfg.random_seed,
|
||||
shuffle=False,
|
||||
)
|
||||
context.val_dataloader = DataLoader(
|
||||
cfg.val_dataset,
|
||||
val_dataset,
|
||||
batch_size=cfg.batch_per_device,
|
||||
sampler=val_sampler,
|
||||
num_workers=cfg.num_workers,
|
||||
pin_memory=cfg.pin_memory,
|
||||
prefetch_factor=cfg.prefetch_factor,
|
||||
collate_fn=cfg.collate_fn,
|
||||
)
|
||||
|
||||
context.model, context.optimizer, context.dataloader, context.scheduler = (
|
||||
@@ -158,11 +192,26 @@ class TrainContextBuilder:
|
||||
if obj is not None:
|
||||
obj.load_state_dict(extra[name])
|
||||
|
||||
strategy_kwargs = dict(cfg.extra_kwargs)
|
||||
|
||||
if cfg.strategy in ("dpo", "grpo"):
|
||||
ref_model = create_ref_model(
|
||||
cfg.model_fn, executor.unwrap_model(context.model)
|
||||
).to(device=device)
|
||||
strategy_kwargs["ref_model"] = ref_model
|
||||
|
||||
if cfg.strategy == "grpo":
|
||||
old_model = create_ref_model(
|
||||
cfg.model_fn, executor.unwrap_model(context.model)
|
||||
).to(device=device)
|
||||
strategy_kwargs["old_model"] = old_model
|
||||
|
||||
context.strategy = StrategyFactory.create(
|
||||
cfg.strategy,
|
||||
model=context.model,
|
||||
train_type=cfg.strategy,
|
||||
device=device,
|
||||
**cfg.extra_kwargs,
|
||||
executor=executor,
|
||||
**strategy_kwargs,
|
||||
)
|
||||
|
||||
return context
|
||||
|
||||
+13
-10
@@ -35,15 +35,14 @@ class Trainer:
|
||||
cfg.ckpt_interval,
|
||||
),
|
||||
CallbackFactory.create(
|
||||
"metric_logger",
|
||||
"metric",
|
||||
log_dir=cfg.log_dir,
|
||||
save_interval=cfg.ckpt_interval,
|
||||
log_interval=cfg.log_interval,
|
||||
metrics=cfg.metrics,
|
||||
val_step=cfg.val_step,
|
||||
),
|
||||
CallbackFactory.create("progress_bar", cfg.n_epoch),
|
||||
CallbackFactory.create("gradient_clipping", cfg.max_grad_norm),
|
||||
CallbackFactory.create("validation"),
|
||||
]
|
||||
return callbacks
|
||||
|
||||
@@ -53,9 +52,11 @@ class Trainer:
|
||||
if method:
|
||||
method(context)
|
||||
|
||||
def _trainer_loop(self, resume_dir: Optional[str] = None):
|
||||
def _trainer_loop(self, param_path: Optional[str] = None, resume: bool = False):
|
||||
context = (
|
||||
TrainContextBuilder(self.train_config).with_resume_dir(resume_dir).build()
|
||||
TrainContextBuilder(self.train_config)
|
||||
.with_param_path(param_path, resume=resume)
|
||||
.build()
|
||||
)
|
||||
executor = context.executor
|
||||
self._call_callbacks("on_train_begin", context)
|
||||
@@ -68,14 +69,15 @@ class Trainer:
|
||||
self._call_callbacks("on_epoch_begin", context)
|
||||
|
||||
for batch in context.dataloader:
|
||||
self._call_callbacks("on_batch_begin", context)
|
||||
|
||||
with executor.accumulate(context.model):
|
||||
self._call_callbacks("on_batch_begin", context)
|
||||
loss = context.strategy(batch)
|
||||
context.loss = loss.item()
|
||||
stand_loss = loss / executor.grad_accum_steps
|
||||
executor.backward(stand_loss)
|
||||
context.iteration += 1
|
||||
context.consumed_samples += (
|
||||
context.config.batch_per_device * context.world_size
|
||||
)
|
||||
self._call_callbacks("on_batch_end", context)
|
||||
|
||||
if executor.sync_gradients:
|
||||
@@ -95,7 +97,7 @@ class Trainer:
|
||||
finally:
|
||||
self._call_callbacks("on_train_end", context)
|
||||
|
||||
def train(self, resume_dir: Optional[str] = None):
|
||||
def train(self, param_path: Optional[str] = None, resume: bool = False):
|
||||
cfg = self.train_config
|
||||
spawn_parallel_fn(
|
||||
self._trainer_loop,
|
||||
@@ -105,5 +107,6 @@ class Trainer:
|
||||
master_port=cfg.master_port,
|
||||
device_type=cfg.device_type,
|
||||
start_method=cfg.start_method,
|
||||
resume_dir=resume_dir,
|
||||
param_path=param_path,
|
||||
resume=resume,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
# Source directory for CUDA kernels — build-time only.
|
||||
# Compiled .so files live in astrAI/_ext/.
|
||||
@@ -0,0 +1,48 @@
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _arch_flags() -> list[str]:
|
||||
import torch
|
||||
|
||||
if torch.cuda.is_available():
|
||||
cap = torch.cuda.get_device_capability()
|
||||
else:
|
||||
cap = (8, 0)
|
||||
ver = f"{cap[0]}{cap[1]}"
|
||||
flags = [f"-gencode=arch=compute_{ver},code=sm_{ver}"]
|
||||
# tensor-core mma path (mma.sync.m16n8k16.bf16) requires sm_80+; decide the
|
||||
# kernel dispatch at build time via this define rather than at runtime.
|
||||
if cap[0] < 8:
|
||||
flags.append("-DASTRAI_NO_MMA")
|
||||
return flags
|
||||
|
||||
|
||||
_kernels_dir = Path("csrc/kernels")
|
||||
REGISTRY: dict[str, dict] = {}
|
||||
|
||||
CXX_FLAGS = ["-O3", "-funroll-loops"]
|
||||
NVCC_FLAGS = [
|
||||
"-O3",
|
||||
"--expt-relaxed-constexpr",
|
||||
"--use_fast_math",
|
||||
"--ptxas-options=-O3,-v",
|
||||
"--extra-device-vectorization",
|
||||
"--threads=8",
|
||||
]
|
||||
|
||||
|
||||
def register(name: str, sources: list[str] | None = None, **kwargs):
|
||||
if sources is None:
|
||||
sources = [str(_kernels_dir / f"{name}.cu")]
|
||||
REGISTRY[name] = {
|
||||
"sources": sources,
|
||||
"cxx_flags": [*CXX_FLAGS],
|
||||
"nvcc_flags": [*NVCC_FLAGS, *_arch_flags()],
|
||||
"extra_link_args": kwargs.pop("extra_link_args", []),
|
||||
**kwargs,
|
||||
}
|
||||
|
||||
|
||||
register("attn_decode")
|
||||
register("attn_prefill")
|
||||
register("attn_paged_decode")
|
||||
@@ -0,0 +1,68 @@
|
||||
#pragma once
|
||||
|
||||
|
||||
template<typename T, typename AT = float>
|
||||
struct AttentionParams {
|
||||
int batch;
|
||||
int q_head;
|
||||
int kv_head;
|
||||
int q_len;
|
||||
int kv_len;
|
||||
int head_dim;
|
||||
int use_mask;
|
||||
int causal_offset; // -1 = non-causal; >=0 = absolute position of first Q token
|
||||
int num_splits;
|
||||
float scale;
|
||||
|
||||
// Q strides (element offsets for each dim — layout-agnostic)
|
||||
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
|
||||
// KV strides (K and V share the same layout — only base pointers differ)
|
||||
int kv_stride_b, kv_stride_h, kv_stride_l, kv_stride_d;
|
||||
|
||||
// Mask: 2D [batch, kv_len] (mask_q_stride=0) or 3D [batch, q_len, kv_len]
|
||||
int mask_b_stride; // = kv_len (both 2D and 3D)
|
||||
int mask_q_stride; // 2D: 0 (all q rows share); 3D: kv_len
|
||||
|
||||
const T* __restrict__ q;
|
||||
const T* __restrict__ k;
|
||||
const T* __restrict__ v;
|
||||
const bool* __restrict__ mask;
|
||||
|
||||
T* __restrict__ o;
|
||||
AT* __restrict__ o_part;
|
||||
AT* __restrict__ ml_part;
|
||||
};
|
||||
|
||||
template<typename T, typename AT = float>
|
||||
struct PagedAttentionParams {
|
||||
int batch;
|
||||
int q_head;
|
||||
int kv_head;
|
||||
int q_len;
|
||||
int kv_len;
|
||||
int head_dim;
|
||||
int use_mask;
|
||||
int causal_offset;
|
||||
float scale;
|
||||
|
||||
int num_splits;
|
||||
int page_size;
|
||||
int max_pages;
|
||||
|
||||
// Q strides (layout-agnostic)
|
||||
int q_stride_b, q_stride_h, q_stride_l, q_stride_d;
|
||||
|
||||
// Mask strides (2D or 3D)
|
||||
int mask_b_stride;
|
||||
int mask_q_stride;
|
||||
|
||||
const T* __restrict__ q;
|
||||
const T* __restrict__ k_cache;
|
||||
const T* __restrict__ v_cache;
|
||||
const bool* __restrict__ mask;
|
||||
const int64_t* __restrict__ page_table;
|
||||
|
||||
T* __restrict__ o;
|
||||
AT* __restrict__ o_part;
|
||||
AT* __restrict__ ml_part;
|
||||
};
|
||||
@@ -0,0 +1,82 @@
|
||||
#include "attn_decode_split_kv.cuh"
|
||||
#include "attn_entry_utils.cuh"
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "attn_decode_split_kv_mma.cuh"
|
||||
#endif
|
||||
|
||||
// Scalar fallback: one warp per query head, split-KV across grid.z.
|
||||
static void launch_scalar_decode(AttentionParams<bf16>& p) {
|
||||
int group_size = p.q_head / p.kv_head;
|
||||
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||
alloc_split_partials(p);
|
||||
|
||||
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
attn_decode_split_kv_kernel<<<dim3(p.batch * p.kv_head, 1, p.num_splits), dim3(32, group_size), smem>>>(p);
|
||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
// MMA head-packing requires G <= 16 (BR=16 rows). sm_80+ tensor-core
|
||||
// + cp.async wins even at G=1 (decode is memory-bound, not compute-bound).
|
||||
// STAGES=2 (double-buffer) for D<=128 (smem 16 KB); STAGES=1 for D=256
|
||||
// (double-buffer would be 32 KB, near the 48 KB static cap — keep single
|
||||
// to preserve occupancy).
|
||||
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
|
||||
static void launch_mma_decode(AttentionParams<bf16>& p) {
|
||||
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||
alloc_split_partials(p);
|
||||
|
||||
attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES><<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
|
||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
#endif
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static void dispatch_decode(AttentionParams<bf16>& p) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
int G = p.q_head / p.kv_head;
|
||||
if (G >= 1 && G <= 16) {
|
||||
launch_mma_decode<HEAD_DIM, 32>(p);
|
||||
return;
|
||||
}
|
||||
#endif
|
||||
launch_scalar_decode(p);
|
||||
}
|
||||
|
||||
torch::Tensor attn_decode(
|
||||
torch::Tensor q,
|
||||
torch::Tensor k,
|
||||
torch::Tensor v,
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
int64_t layout
|
||||
) {
|
||||
AttentionParams<bf16> p;
|
||||
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
|
||||
TORCH_CHECK(p.q_len == 1, "Q seq_len must be 1");
|
||||
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
||||
|
||||
// O matches Q's original layout
|
||||
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
||||
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
||||
p.o = (bf16*)O_view.data_ptr();
|
||||
|
||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p);
|
||||
return O;
|
||||
}
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def("attn_decode", &attn_decode,
|
||||
py::arg("q"),
|
||||
py::arg("k"),
|
||||
py::arg("v"),
|
||||
py::arg("mask") = py::none(),
|
||||
py::arg("causal_offset") = -1,
|
||||
py::arg("scale") = 0.0,
|
||||
py::arg("layout") = 0,
|
||||
"GQA decode (tensor-core head-packing on sm_80+, scalar fallback)");
|
||||
}
|
||||
@@ -0,0 +1,132 @@
|
||||
#pragma once
|
||||
#include <cuda_bf16.h>
|
||||
#include <float.h>
|
||||
#include "attn_common.h"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
constexpr int DC_CHUNK = 64;
|
||||
|
||||
__device__ inline float warp_reduce_sum(float val) {
|
||||
for (int offset = 16; offset > 0; offset >>= 1)
|
||||
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
|
||||
return val;
|
||||
}
|
||||
|
||||
__global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
||||
int batch = blockIdx.x / p.kv_head;
|
||||
int kv_head = blockIdx.x % p.kv_head;
|
||||
int split = blockIdx.z;
|
||||
int group_size = blockDim.y;
|
||||
int q_head = kv_head * group_size + threadIdx.y;
|
||||
int lane = threadIdx.x;
|
||||
int hd_per_thread = p.head_dim / 32;
|
||||
|
||||
// Q: [batch, q_head, q_len=1, head_dim] — stride-based
|
||||
float q_reg[8];
|
||||
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
||||
+ lane * hd_per_thread * p.q_stride_d;
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
|
||||
|
||||
// KV: [batch, kv_head, kv_len, head_dim] — stride-based base
|
||||
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||
int mask_base = batch * p.mask_b_stride;
|
||||
|
||||
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
|
||||
|
||||
extern __shared__ __align__(16) bf16 k_smem[];
|
||||
|
||||
// Split-KV: each split processes a contiguous subset of chunks
|
||||
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||
int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits;
|
||||
int ch_begin = split * chunks_per_split;
|
||||
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
|
||||
|
||||
for (int ci = ch_begin; ci < ch_end; ci++) {
|
||||
int chunk_start = ci * DC_CHUNK;
|
||||
int this_chunk = min(DC_CHUNK, p.kv_len - chunk_start);
|
||||
|
||||
// Load K into shared memory (gather from strided global)
|
||||
int total = this_chunk * p.head_dim;
|
||||
for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y) {
|
||||
int s = i / p.head_dim;
|
||||
int d_dim = i % p.head_dim;
|
||||
int kv_idx = chunk_start + s;
|
||||
int g_off = kv_base + kv_idx * p.kv_stride_l + d_dim * p.kv_stride_d;
|
||||
k_smem[i] = p.k[g_off];
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
for (int s = 0; s < this_chunk; s++) {
|
||||
float partial = 0.0f;
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
||||
partial = warp_reduce_sum(partial) * p.scale;
|
||||
|
||||
int kv_idx = chunk_start + s;
|
||||
if (p.use_mask && p.mask && !p.mask[mask_base + kv_idx])
|
||||
partial = -FLT_MAX;
|
||||
if (p.causal_offset >= 0 && kv_idx > p.causal_offset)
|
||||
partial = -FLT_MAX;
|
||||
|
||||
float new_m = fmaxf(m, partial);
|
||||
float alpha = expf(m - new_m);
|
||||
float beta = expf(partial - new_m);
|
||||
d = d * alpha + beta;
|
||||
|
||||
// V: stride-based read
|
||||
int v_off = kv_base + kv_idx * p.kv_stride_l + lane * hd_per_thread * p.kv_stride_d;
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v[v_off + i * p.kv_stride_d]) * beta;
|
||||
m = new_m;
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
// ---- write UN-normalised partials for this split ----
|
||||
size_t bh = (size_t)batch * p.q_head + q_head;
|
||||
size_t slot = bh * p.num_splits + split;
|
||||
int d0 = lane * hd_per_thread;
|
||||
for (int i = 0; i < hd_per_thread; i++) {
|
||||
int dd = d0 + i;
|
||||
p.o_part[slot * p.head_dim + dd] = acc_reg[i];
|
||||
}
|
||||
if (lane == 0) {
|
||||
p.ml_part[slot * 2] = m;
|
||||
p.ml_part[slot * 2 + 1] = d;
|
||||
}
|
||||
}
|
||||
|
||||
// Reduce split-K partials into the final bf16 output. One block per (batch,
|
||||
// q_head); each thread folds across all splits with a single-pass
|
||||
// online-rescale reduction (expf + FMA counts halved vs 3-pass original).
|
||||
__global__ void attn_decode_combine_kernel(AttentionParams<bf16> p) {
|
||||
int bh = blockIdx.x;
|
||||
int d = threadIdx.x;
|
||||
if (d >= p.head_dim) return;
|
||||
|
||||
int batch = bh / p.q_head;
|
||||
int q_head = bh % p.q_head;
|
||||
|
||||
size_t split_base = (size_t)bh * p.num_splits;
|
||||
const float* mlp = p.ml_part + split_base * 2;
|
||||
const float* op = p.o_part + split_base * p.head_dim;
|
||||
|
||||
float m = -FLT_MAX, l = 0.0f, acc = 0.0f;
|
||||
for (int s = 0; s < p.num_splits; s++) {
|
||||
float mi = mlp[s * 2];
|
||||
if (mi <= -FLT_MAX) continue;
|
||||
float li = mlp[s * 2 + 1];
|
||||
float nm = fmaxf(m, mi);
|
||||
float corr = __expf(m - nm);
|
||||
float e = __expf(mi - nm);
|
||||
acc = acc * corr + op[s * p.head_dim + d] * e;
|
||||
l = l * corr + li * e;
|
||||
m = nm;
|
||||
}
|
||||
|
||||
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
||||
// Stride-based output write (q_len=1 for decode, so stride_l not needed)
|
||||
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d;
|
||||
p.o[o_off] = __float2bfloat16(acc * inv);
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
#pragma once
|
||||
#include <cfloat>
|
||||
#include <cuda_bf16.h>
|
||||
#include "attn_common.h"
|
||||
#include "attn_mma_utils.cuh"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
// Split-K (FlashDecoding) tensor-core decode via GQA head-packing.
|
||||
//
|
||||
// Decode has q_len == 1, so S = q @ K^T is a GEMV per head — no tensor-core
|
||||
// work on its own. But GQA gives us G = q_head / kv_head query heads that all
|
||||
// share one kv_head. We pack those G heads into the M=16 rows of
|
||||
// mma.sync.m16n8k16, turning G independent GEMVs into a single GEMM that
|
||||
// reuses each loaded K/V tile across all G heads (K/V load is the decode
|
||||
// bottleneck, so the reuse is the win, not the flops). The KV sequence is
|
||||
// partitioned across gridDim.z blocks so that a decode with only
|
||||
// batch*kv_head independent tasks can fill all SMs. Each (batch, kv_head,
|
||||
// split) block computes an UN-normalised partial (Oacc, m, l) over its KV
|
||||
// slice; the combine kernel below reduces across splits. Fixes the "grid too
|
||||
// small" bottleneck (0.04 waves/SM → many blocks) for long-context,
|
||||
// small-batch decode.
|
||||
|
||||
template <int HEAD_DIM, int BC, int STAGES = 2>
|
||||
__global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
||||
constexpr int KD = HEAD_DIM / 16;
|
||||
constexpr int NC8 = BC / 8;
|
||||
constexpr int KT2 = BC / 16;
|
||||
constexpr int DN8 = HEAD_DIM / 8;
|
||||
constexpr int LD = HEAD_DIM;
|
||||
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
|
||||
constexpr int VEC = 8;
|
||||
constexpr int TOTAL = BC * HEAD_DIM;
|
||||
|
||||
const int lane = threadIdx.x;
|
||||
const int gid = lane >> 2;
|
||||
const int tid4 = lane & 3;
|
||||
|
||||
const int kv_head = blockIdx.x;
|
||||
const int batch = blockIdx.y;
|
||||
const int split = blockIdx.z;
|
||||
const int G = p.q_head / p.kv_head;
|
||||
const int q_head0 = kv_head * G;
|
||||
|
||||
// Double-buffered shared memory for K/V (no sQ needed — Q goes direct
|
||||
// from global to registers).
|
||||
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
|
||||
__shared__ __align__(16) bf16 sV[STAGES * BC * LD];
|
||||
|
||||
// ---- Load Q directly from global into mma A-operand registers ----
|
||||
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
|
||||
const int qra = gid;
|
||||
const int qrb = gid + 8;
|
||||
const bool va = qra < G, vb = qrb < G;
|
||||
unsigned Qa[KD][4];
|
||||
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
|
||||
qra, qrb, va, vb, tid4, Qa);
|
||||
|
||||
float Oacc[DN8][4];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < DN8; j++)
|
||||
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||
|
||||
// KV: stride-based base — [batch, kv_head, kv_len, head_dim]
|
||||
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||
const int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
|
||||
const int ti_begin = split * tiles_per_split;
|
||||
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
|
||||
const int has_mask = p.use_mask && p.mask;
|
||||
|
||||
// ---- Load tile lambda: predicated cp.async, unified full/partial ----
|
||||
auto load_tile = [&](int ti, int buf) {
|
||||
int kv0 = ti * BC;
|
||||
bf16* dK = sK + buf * BC * LD;
|
||||
bf16* dV = sV + buf * BC * LD;
|
||||
#pragma unroll
|
||||
for (int i = lane * VEC; i < TOTAL; i += 32 * VEC) {
|
||||
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
||||
int kc = kv0 + r;
|
||||
bool valid = kc < p.kv_len;
|
||||
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
|
||||
// KV stride-based: contiguous within head_dim (stride_d == 1 typically)
|
||||
int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
|
||||
cp_async_16_pred(&dK[off], &p.k[g_off], valid);
|
||||
cp_async_16_pred(&dV[off], &p.v[g_off], valid);
|
||||
}
|
||||
cp_async_commit();
|
||||
};
|
||||
|
||||
// ---- Prologue: issue first tile load ----
|
||||
if (ti_begin < ti_end) {
|
||||
load_tile(ti_begin, 0);
|
||||
}
|
||||
|
||||
for (int ti = ti_begin; ti < ti_end; ti++) {
|
||||
constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0;
|
||||
int buf = (ti - ti_begin) & BUF_MASK;
|
||||
|
||||
// Wait for current tile, then issue next tile's prefetch (overlaps
|
||||
// with this tile's compute). Single syncwarp covers both hazards.
|
||||
// When STAGES==1, no prefetch — load happens at end of prior iter.
|
||||
cp_async_wait_group<0>();
|
||||
__syncwarp();
|
||||
if constexpr (STAGES > 1) {
|
||||
if (ti + 1 < ti_end)
|
||||
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
|
||||
}
|
||||
|
||||
const bf16* bK = sK + buf * BC * LD;
|
||||
const bf16* bV = sV + buf * BC * LD;
|
||||
int kv0 = ti * BC;
|
||||
|
||||
float Sacc[NC8][4];
|
||||
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
|
||||
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < NC8; n8++)
|
||||
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||
|
||||
// Decode: q_len=1, so qrow0=qrow1=0, mask_q_stride irrelevant
|
||||
int maxc = (p.causal_offset >= 0) ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
||||
mma_softmax_tile<NC8, DN8>(kv0, maxc, maxc,
|
||||
0, 0,
|
||||
p.mask_b_stride, 0,
|
||||
batch,
|
||||
p.mask, has_mask,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
|
||||
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
|
||||
__syncwarp();
|
||||
|
||||
if constexpr (STAGES == 1) {
|
||||
if (ti + 1 < ti_end)
|
||||
load_tile(ti + 1, 0);
|
||||
}
|
||||
}
|
||||
|
||||
// ---- write UN-normalised partials for this split ----
|
||||
auto split_slot = [&](int h) -> size_t {
|
||||
size_t bh = (size_t)batch * p.q_head + h;
|
||||
return bh * p.num_splits + split;
|
||||
};
|
||||
#pragma unroll
|
||||
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
||||
int d = dn8 * 8 + 2 * tid4;
|
||||
int r0 = gid, r1 = gid + 8;
|
||||
if (r0 < G) {
|
||||
int h = q_head0 + r0;
|
||||
float* op = p.o_part + split_slot(h) * HEAD_DIM;
|
||||
op[d] = Oacc[dn8][0];
|
||||
op[d + 1] = Oacc[dn8][1];
|
||||
}
|
||||
if (r1 < G) {
|
||||
int h = q_head0 + r1;
|
||||
float* op = p.o_part + split_slot(h) * HEAD_DIM;
|
||||
op[d] = Oacc[dn8][2];
|
||||
op[d + 1] = Oacc[dn8][3];
|
||||
}
|
||||
}
|
||||
if (tid4 == 0) {
|
||||
int r0 = gid, r1 = gid + 8;
|
||||
if (r0 < G) {
|
||||
int h = q_head0 + r0;
|
||||
float* mp = p.ml_part + split_slot(h) * 2;
|
||||
mp[0] = m0; mp[1] = l0;
|
||||
}
|
||||
if (r1 < G) {
|
||||
int h = q_head0 + r1;
|
||||
float* mp = p.ml_part + split_slot(h) * 2;
|
||||
mp[0] = m1; mp[1] = l1;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,177 @@
|
||||
#pragma once
|
||||
#include <torch/extension.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include "attn_common.h"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
inline int compute_num_splits(int base_blocks, int tiles_total) {
|
||||
int sm_count = 0;
|
||||
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
|
||||
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
|
||||
return std::max(1, std::min(n, std::min(tiles_total, 32)));
|
||||
}
|
||||
|
||||
// Dispatch head_dim: shared macro — avoids C++20 lambda template syntax.
|
||||
// Usage: DISPATCH_HEAD_DIM(hd, fn, arg)
|
||||
// Expands to: fn<32>(arg); fn<64>(arg); etc.
|
||||
#define DISPATCH_HEAD_DIM(hd, fn, arg) \
|
||||
switch (hd) { \
|
||||
case 32: fn<32>(arg); break; \
|
||||
case 64: fn<64>(arg); break; \
|
||||
case 128: fn<128>(arg); break; \
|
||||
case 256: fn<256>(arg); break; \
|
||||
default: \
|
||||
TORCH_CHECK(false, "unsupported head_dim ", hd, \
|
||||
" (supported: 32, 64, 128, 256)"); \
|
||||
}
|
||||
|
||||
template<typename P>
|
||||
inline void alloc_split_partials(P& p) {
|
||||
auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA);
|
||||
auto o_part = torch::empty({p.batch, p.q_head, p.num_splits, p.head_dim}, fopt);
|
||||
auto ml_part = torch::empty({p.batch, p.q_head, p.num_splits, 2}, fopt);
|
||||
p.o_part = (float*)o_part.data_ptr();
|
||||
p.ml_part = (float*)ml_part.data_ptr();
|
||||
}
|
||||
|
||||
// ---- Shared Q-dims + strides extraction ----
|
||||
template <typename P>
|
||||
inline void extract_q_dims_and_strides(torch::Tensor& q, int64_t layout, P& p) {
|
||||
if (layout == 1) q = q.transpose(1, 2);
|
||||
p.batch = (int)q.size(0);
|
||||
p.q_head = (int)q.size(1);
|
||||
p.q_len = (int)q.size(2);
|
||||
p.head_dim = (int)q.size(3);
|
||||
p.q_stride_b = (int)q.stride(0);
|
||||
p.q_stride_h = (int)q.stride(1);
|
||||
p.q_stride_l = (int)q.stride(2);
|
||||
p.q_stride_d = (int)q.stride(3);
|
||||
}
|
||||
|
||||
// ---- Shared mask packing ----
|
||||
template <typename P>
|
||||
inline void pack_mask(const c10::optional<torch::Tensor>& mask, P& p) {
|
||||
if (p.use_mask) {
|
||||
auto m = mask.value();
|
||||
TORCH_CHECK(m.is_cuda(), "mask must be on CUDA");
|
||||
TORCH_CHECK(m.dtype() == torch::kBool, "mask must be bool");
|
||||
TORCH_CHECK(m.size(0) == p.batch, "mask batch mismatch");
|
||||
TORCH_CHECK(m.size(m.dim() - 1) == p.kv_len, "mask kv_len mismatch");
|
||||
if (m.dim() == 2) {
|
||||
p.mask_b_stride = (int)m.stride(0);
|
||||
p.mask_q_stride = 0;
|
||||
} else if (m.dim() == 3) {
|
||||
TORCH_CHECK(m.size(1) == p.q_len, "mask q_len mismatch");
|
||||
p.mask_b_stride = (int)m.stride(0);
|
||||
p.mask_q_stride = (int)m.stride(1);
|
||||
} else {
|
||||
TORCH_CHECK(false, "mask must be 2D [batch, kv_len] or 3D [batch, q_len, kv_len]");
|
||||
}
|
||||
p.mask = m.data_ptr<bool>();
|
||||
} else {
|
||||
p.mask = nullptr;
|
||||
p.mask_b_stride = 0;
|
||||
p.mask_q_stride = 0;
|
||||
}
|
||||
}
|
||||
|
||||
// ---- attn_pack_params (contiguous KV) ----
|
||||
template<typename T>
|
||||
inline void attn_pack_params(
|
||||
torch::Tensor q,
|
||||
torch::Tensor k,
|
||||
torch::Tensor v,
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
int64_t layout,
|
||||
AttentionParams<T>& p
|
||||
) {
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||
|
||||
TORCH_CHECK(q.is_cuda() && k.is_cuda() && v.is_cuda());
|
||||
TORCH_CHECK(q.dtype() == torch::kBFloat16);
|
||||
TORCH_CHECK(k.dtype() == torch::kBFloat16);
|
||||
TORCH_CHECK(v.dtype() == torch::kBFloat16);
|
||||
TORCH_CHECK(k.sizes() == v.sizes(), "K and V must have identical shapes");
|
||||
TORCH_CHECK(q.dim() == 4 && k.dim() == 4, "Q/K/V must be 4D");
|
||||
|
||||
extract_q_dims_and_strides(q, layout, p);
|
||||
|
||||
if (layout == 1) k = k.transpose(1, 2), v = v.transpose(1, 2);
|
||||
|
||||
p.kv_head = (int)k.size(1);
|
||||
p.kv_len = (int)k.size(2);
|
||||
TORCH_CHECK(k.size(3) == p.head_dim, "K/V head_dim must match Q");
|
||||
|
||||
p.kv_stride_b = (int)k.stride(0);
|
||||
p.kv_stride_h = (int)k.stride(1);
|
||||
p.kv_stride_l = (int)k.stride(2);
|
||||
p.kv_stride_d = (int)k.stride(3);
|
||||
|
||||
p.causal_offset = (int)causal_offset;
|
||||
p.use_mask = mask.has_value() ? 1 : 0;
|
||||
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
|
||||
|
||||
p.q = (const T*)q.data_ptr();
|
||||
p.k = (const T*)k.data_ptr();
|
||||
p.v = (const T*)v.data_ptr();
|
||||
p.o = nullptr;
|
||||
p.o_part = nullptr;
|
||||
p.ml_part = nullptr;
|
||||
|
||||
pack_mask(mask, p);
|
||||
}
|
||||
|
||||
// ---- attn_pack_paged_params ----
|
||||
template<typename T>
|
||||
inline void attn_pack_paged_params(
|
||||
torch::Tensor q,
|
||||
torch::Tensor page_table,
|
||||
torch::Tensor k_cache,
|
||||
torch::Tensor v_cache,
|
||||
int64_t page_size,
|
||||
int64_t kv_len,
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
int64_t layout,
|
||||
PagedAttentionParams<T>& p
|
||||
) {
|
||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||
|
||||
TORCH_CHECK(q.is_cuda() && page_table.is_cuda() && k_cache.is_cuda() && v_cache.is_cuda());
|
||||
TORCH_CHECK(q.dtype() == torch::kBFloat16, "q must be bf16");
|
||||
TORCH_CHECK(k_cache.dtype() == torch::kBFloat16, "k_cache must be bf16");
|
||||
TORCH_CHECK(v_cache.dtype() == torch::kBFloat16, "v_cache must be bf16");
|
||||
TORCH_CHECK(page_table.dtype() == torch::kLong, "page_table must be int64");
|
||||
TORCH_CHECK(k_cache.sizes() == v_cache.sizes(), "k_cache and v_cache must have identical shapes");
|
||||
|
||||
extract_q_dims_and_strides(q, layout, p);
|
||||
|
||||
p.kv_head = (int)k_cache.size(2);
|
||||
p.kv_len = (int)kv_len;
|
||||
p.page_size = (int)page_size;
|
||||
p.max_pages = (int)page_table.size(1);
|
||||
|
||||
TORCH_CHECK(q.size(2) == 1, "Q seq_len must be 1 (decode)");
|
||||
TORCH_CHECK(p.head_dim % 32 == 0, "head_dim must be multiple of 32");
|
||||
TORCH_CHECK(k_cache.size(1) == page_size,
|
||||
"k_cache dim 1 must equal page_size, got ",
|
||||
k_cache.size(1), " vs ", page_size);
|
||||
|
||||
p.causal_offset = (int)causal_offset;
|
||||
p.use_mask = (mask.has_value() && mask.value().defined()) ? 1 : 0;
|
||||
p.scale = (scale > 0.0) ? (float)scale : 1.0f / sqrtf((float)p.head_dim);
|
||||
|
||||
p.page_table = page_table.data_ptr<int64_t>();
|
||||
p.k_cache = (const T*)k_cache.data_ptr();
|
||||
p.v_cache = (const T*)v_cache.data_ptr();
|
||||
p.q = (const T*)q.data_ptr();
|
||||
p.o = nullptr;
|
||||
p.o_part = nullptr;
|
||||
p.ml_part = nullptr;
|
||||
|
||||
pack_mask(mask, p);
|
||||
}
|
||||
@@ -0,0 +1,293 @@
|
||||
#pragma once
|
||||
#include <cfloat>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
// Shared MMA utilities for tensor-core GQA kernels.
|
||||
// mma.sync.m16n8k16 PTX wrappers, ldmatrix helpers, and bf16 packing.
|
||||
|
||||
// mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32
|
||||
__device__ __forceinline__ void mma16816(float* d, const unsigned* a,
|
||||
const unsigned* b, const float* c) {
|
||||
asm volatile(
|
||||
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
|
||||
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};"
|
||||
: "=f"(d[0]), "=f"(d[1]), "=f"(d[2]), "=f"(d[3])
|
||||
: "r"(a[0]), "r"(a[1]), "r"(a[2]), "r"(a[3]), "r"(b[0]), "r"(b[1]),
|
||||
"f"(c[0]), "f"(c[1]), "f"(c[2]), "f"(c[3]));
|
||||
}
|
||||
|
||||
// read two adjacent bf16 from smem as one packed .b32 (elem0 low, elem1 high)
|
||||
__device__ __forceinline__ unsigned ld2(const bf16* p) {
|
||||
return *reinterpret_cast<const unsigned*>(p);
|
||||
}
|
||||
|
||||
// pack two floats into one bf16x2 as .b32
|
||||
__device__ __forceinline__ unsigned pk2(float a, float b) {
|
||||
__nv_bfloat162 v = __floats2bfloat162_rn(a, b);
|
||||
return *reinterpret_cast<unsigned*>(&v);
|
||||
}
|
||||
|
||||
// pack two (non-contiguous) bf16 into one .b32
|
||||
__device__ __forceinline__ unsigned pkb(bf16 a, bf16 b) {
|
||||
__nv_bfloat162 v;
|
||||
v.x = a;
|
||||
v.y = b;
|
||||
return *reinterpret_cast<unsigned*>(&v);
|
||||
}
|
||||
|
||||
// ldmatrix: cooperatively load mma fragments from smem (one instruction per
|
||||
// 16x16 / 16x8 tile) with the exact register layout mma expects — replaces the
|
||||
// scalar per-thread fragment packing, cutting shared-load instructions and bank
|
||||
// conflicts. Each lane supplies the shared address of one 8-wide row.
|
||||
__device__ __forceinline__ void ldmatrix_x4(unsigned* r, const bf16* p) {
|
||||
unsigned a = __cvta_generic_to_shared(p);
|
||||
asm volatile("ldmatrix.sync.aligned.m8n8.x4.shared.b16 {%0,%1,%2,%3}, [%4];"
|
||||
: "=r"(r[0]), "=r"(r[1]), "=r"(r[2]), "=r"(r[3])
|
||||
: "r"(a));
|
||||
}
|
||||
__device__ __forceinline__ void ldmatrix_x2(unsigned* r, const bf16* p) {
|
||||
unsigned a = __cvta_generic_to_shared(p);
|
||||
asm volatile("ldmatrix.sync.aligned.m8n8.x2.shared.b16 {%0,%1}, [%2];"
|
||||
: "=r"(r[0]), "=r"(r[1])
|
||||
: "r"(a));
|
||||
}
|
||||
__device__ __forceinline__ void ldmatrix_x2_trans(unsigned* r, const bf16* p) {
|
||||
unsigned a = __cvta_generic_to_shared(p);
|
||||
asm volatile("ldmatrix.sync.aligned.m8n8.x2.trans.shared.b16 {%0,%1}, [%2];"
|
||||
: "=r"(r[0]), "=r"(r[1])
|
||||
: "r"(a));
|
||||
}
|
||||
|
||||
// XOR swizzle for shared-memory column at 8-bf16 chunk granularity.
|
||||
// Eliminates ldmatrix bank conflicts without LD padding: consecutive rows
|
||||
// land in distinct bank groups. swiz_col(d, r, mask) = ((d>>3)^(r&mask))<<3 | (d&7).
|
||||
// mask must cover log2(HEAD_DIM/8) chunk bits but stay within LD: use 7 for
|
||||
// HEAD_DIM>=64 (8+ chunks), 3 for HEAD_DIM=32 (4 chunks). Default 7 keeps
|
||||
// existing HEAD_DIM>=64 call sites working unchanged.
|
||||
__device__ __forceinline__ int swiz_col(int d, int r, int mask = 7) {
|
||||
return ((d >> 3) ^ (r & mask)) << 3 | (d & 7);
|
||||
}
|
||||
|
||||
// cp.async: copy 16 bytes (8 bf16) from global to shared memory directly,
|
||||
// bypassing registers. Eliminates shared-store bank conflicts and cuts
|
||||
// load-loop instruction count in half (1 cp.async vs 1 LDG + 1 STS).
|
||||
// Requires sm_80+.
|
||||
__device__ __forceinline__ void cp_async_16(bf16* smem_ptr, const void* gmem_ptr) {
|
||||
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
|
||||
asm volatile("cp.async.ca.shared.global [%0], [%1], 16;"
|
||||
:: "r"(smem_addr), "l"(gmem_ptr));
|
||||
}
|
||||
|
||||
// Predicated cp.async: copy 16 bytes when `pred`, otherwise zero-fill the
|
||||
// destination (src-size operand = 0 → no bytes read from src, so an
|
||||
// out-of-bounds src address is never dereferenced). Lets full and partial
|
||||
// tiles share one uniform async load path — no scalar fallback branch.
|
||||
__device__ __forceinline__ void cp_async_16_pred(bf16* smem_ptr,
|
||||
const void* gmem_ptr,
|
||||
bool pred) {
|
||||
unsigned smem_addr = __cvta_generic_to_shared(smem_ptr);
|
||||
int src_size = pred ? 16 : 0;
|
||||
asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;"
|
||||
:: "r"(smem_addr), "l"(gmem_ptr), "r"(src_size));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void cp_async_commit() {
|
||||
asm volatile("cp.async.commit_group;");
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void cp_async_wait_all() {
|
||||
asm volatile("cp.async.wait_all;");
|
||||
}
|
||||
|
||||
// Wait until at most N commit groups are still in flight. Used for
|
||||
// double-buffered pipelining: wait_group<1> lets the next tile's cp.async
|
||||
// continue while ensuring the current tile's data is ready.
|
||||
template <int N>
|
||||
__device__ __forceinline__ void cp_async_wait_group() {
|
||||
asm volatile("cp.async.wait_group %0;" :: "n"(N));
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Q-load: load query rows directly from global memory into mma A-operand
|
||||
// register layout. One call replaces ~15 duplicated lines in each MMA kernel.
|
||||
// stride_row is p.q_stride_h for decode (q_len=1, G heads) or
|
||||
// p.q_stride_l for prefill (multi-q rows).
|
||||
// ---------------------------------------------------------------------------
|
||||
template <int KD>
|
||||
__device__ inline void load_q_mma_frags(
|
||||
const bf16* __restrict__ q,
|
||||
int stride_row,
|
||||
int stride_d,
|
||||
int qra, int qrb,
|
||||
bool va, bool vb,
|
||||
int tid4,
|
||||
unsigned Qa[KD][4])
|
||||
{
|
||||
#pragma unroll
|
||||
for (int kt = 0; kt < KD; kt++) {
|
||||
int c = kt * 16 + tid4 * 2;
|
||||
const unsigned* pau = reinterpret_cast<const unsigned*>(
|
||||
&q[qra * stride_row + c * stride_d]);
|
||||
const unsigned* pbu = reinterpret_cast<const unsigned*>(
|
||||
&q[qrb * stride_row + c * stride_d]);
|
||||
Qa[kt][0] = va ? pau[0] : 0u;
|
||||
Qa[kt][1] = vb ? pbu[0] : 0u;
|
||||
Qa[kt][2] = va ? pau[4] : 0u;
|
||||
Qa[kt][3] = vb ? pbu[4] : 0u;
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Shared MMA compute functions — used by both decode and prefill MMA kernels.
|
||||
// Extracted because S=Q@K^T, online softmax, and P@V are structurally identical
|
||||
// between the two kernels; only the per-row causal/mask bounds differ.
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// S = Q @ K^T (Qa pre-loaded by the caller; scale applied post-mma in the
|
||||
// caller to avoid bf16 precision loss).
|
||||
// LD and SWIZ_MASK are constexpr in the calling kernel — passing them as
|
||||
// runtime ints lets the compiler fold them while keeping the signature clean.
|
||||
template <int KD, int NC8>
|
||||
__device__ inline void mma_compute_scores(
|
||||
const unsigned Qa[KD][4],
|
||||
const bf16* __restrict__ sK,
|
||||
int LD,
|
||||
int SWIZ_MASK,
|
||||
int lane,
|
||||
float Sacc[NC8][4])
|
||||
{
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < NC8; n8++) {
|
||||
Sacc[n8][0] = Sacc[n8][1] = Sacc[n8][2] = Sacc[n8][3] = 0.0f;
|
||||
int krow_l = n8 * 8 + (lane & 7);
|
||||
int kcol_h = (lane & 8) ? 8 : 0;
|
||||
#pragma unroll
|
||||
for (int kt = 0; kt < KD; kt++) {
|
||||
unsigned b[2];
|
||||
ldmatrix_x2(b, &sK[krow_l * LD + swiz_col(kt * 16 + kcol_h, krow_l, SWIZ_MASK)]);
|
||||
mma16816(Sacc[n8], Qa[kt], b, Sacc[n8]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Online softmax + Oacc rescale for one K/V tile.
|
||||
// maxc0/maxc1: per-row KV column bounds (prefill: per-query-row causal limits;
|
||||
// decode: same value for both rows since q_len==1).
|
||||
// qrow0/qrow1: query row indices (for 3D mask indexing; decode passes 0).
|
||||
// mask_b_stride/mask_q_stride: mask layout (2D: mask_q_stride=0; 3D: =kv_len).
|
||||
// Reads Sacc (Q@K^T scores), applies causal/mask, computes P = exp(S - nm),
|
||||
// rescales Oacc by exp(m_old - nm), and updates m/l — all in place.
|
||||
template <int NC8, int DN8>
|
||||
__device__ inline void mma_softmax_tile(
|
||||
int kv0,
|
||||
int maxc0,
|
||||
int maxc1,
|
||||
int qrow0,
|
||||
int qrow1,
|
||||
int mask_b_stride,
|
||||
int mask_q_stride,
|
||||
int mask_batch,
|
||||
const bool* __restrict__ mask,
|
||||
bool has_mask,
|
||||
float Sacc[NC8][4],
|
||||
float Oacc[DN8][4],
|
||||
float& m0, float& m1,
|
||||
float& l0, float& l1,
|
||||
int lane)
|
||||
{
|
||||
int tid4 = lane & 3;
|
||||
|
||||
// Mask out-of-bounds / masked columns: set -FLT_MAX so expf → 0 downstream
|
||||
// without per-element sentinel checks. Compute tile-local row maxima.
|
||||
float rmax0 = -FLT_MAX, rmax1 = -FLT_MAX;
|
||||
int mask_base0 = mask_batch * mask_b_stride + qrow0 * mask_q_stride;
|
||||
int mask_base1 = mask_batch * mask_b_stride + qrow1 * mask_q_stride;
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < NC8; n8++) {
|
||||
int cc = kv0 + n8 * 8 + 2 * tid4;
|
||||
int c1 = cc + 1;
|
||||
bool b0 = (cc >= maxc0) || (has_mask && !mask[mask_base0 + cc]);
|
||||
bool b1 = (c1 >= maxc0) || (has_mask && !mask[mask_base0 + c1]);
|
||||
bool b2 = (cc >= maxc1) || (has_mask && !mask[mask_base1 + cc]);
|
||||
bool b3 = (c1 >= maxc1) || (has_mask && !mask[mask_base1 + c1]);
|
||||
float s0 = b0 ? -FLT_MAX : Sacc[n8][0];
|
||||
float s1 = b1 ? -FLT_MAX : Sacc[n8][1];
|
||||
float s2 = b2 ? -FLT_MAX : Sacc[n8][2];
|
||||
float s3 = b3 ? -FLT_MAX : Sacc[n8][3];
|
||||
Sacc[n8][0] = s0; Sacc[n8][1] = s1;
|
||||
Sacc[n8][2] = s2; Sacc[n8][3] = s3;
|
||||
rmax0 = fmaxf(rmax0, fmaxf(s0, s1));
|
||||
rmax1 = fmaxf(rmax1, fmaxf(s2, s3));
|
||||
}
|
||||
// Warp-reduce row maxima across the 4-lane thread group (xor 1, xor 2).
|
||||
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 1));
|
||||
rmax0 = fmaxf(rmax0, __shfl_xor_sync(0xFFFFFFFF, rmax0, 2));
|
||||
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 1));
|
||||
rmax1 = fmaxf(rmax1, __shfl_xor_sync(0xFFFFFFFF, rmax1, 2));
|
||||
|
||||
// nm = max(running max m, tile-local max rmax) — updated running maximum.
|
||||
float nm0 = fmaxf(m0, rmax0), nm1 = fmaxf(m1, rmax1);
|
||||
// corr rescales Oacc and l by exp(m_old - nm). When all-masked (m == nm ==
|
||||
// -FLT_MAX), exp(0) = 1 — correct, no guard needed.
|
||||
float corr0 = __expf(m0 - nm0);
|
||||
float corr1 = __expf(m1 - nm1);
|
||||
// pn guards only the all-masked-row edge: if nm == -FLT_MAX, exp(S - nm)
|
||||
// gives 1 not 0 for masked entries. Two scalar masks replace 4*NC8
|
||||
// per-element comparisons.
|
||||
float pn0 = (nm0 == -FLT_MAX) ? 0.0f : 1.0f;
|
||||
float pn1 = (nm1 == -FLT_MAX) ? 0.0f : 1.0f;
|
||||
|
||||
// P = exp(S - nm) for each element. Masked entries (Sacc = -FLT_MAX) give
|
||||
// exp(-inf) ≈ 0 naturally; pn zero-fills the all-masked-row edge.
|
||||
float rsum0 = 0.0f, rsum1 = 0.0f;
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < NC8; n8++) {
|
||||
float p0 = pn0 * __expf(Sacc[n8][0] - nm0);
|
||||
float p1 = pn0 * __expf(Sacc[n8][1] - nm0);
|
||||
float p2 = pn1 * __expf(Sacc[n8][2] - nm1);
|
||||
float p3 = pn1 * __expf(Sacc[n8][3] - nm1);
|
||||
Sacc[n8][0] = p0; Sacc[n8][1] = p1;
|
||||
Sacc[n8][2] = p2; Sacc[n8][3] = p3;
|
||||
rsum0 += p0 + p1;
|
||||
rsum1 += p2 + p3;
|
||||
}
|
||||
rsum0 += __shfl_xor_sync(0xFFFFFFFF, rsum0, 1);
|
||||
rsum0 += __shfl_xor_sync(0xFFFFFFFF, rsum0, 2);
|
||||
rsum1 += __shfl_xor_sync(0xFFFFFFFF, rsum1, 1);
|
||||
rsum1 += __shfl_xor_sync(0xFFFFFFFF, rsum1, 2);
|
||||
l0 = l0 * corr0 + rsum0;
|
||||
l1 = l1 * corr1 + rsum1;
|
||||
m0 = nm0; m1 = nm1;
|
||||
|
||||
#pragma unroll
|
||||
for (int j = 0; j < DN8; j++) {
|
||||
Oacc[j][0] *= corr0; Oacc[j][1] *= corr0;
|
||||
Oacc[j][2] *= corr1; Oacc[j][3] *= corr1;
|
||||
}
|
||||
}
|
||||
|
||||
// O += P @ V (Sacc must contain P = attention weights after softmax).
|
||||
template <int DN8, int KT2>
|
||||
__device__ inline void mma_pv_accumulate(
|
||||
float Sacc[][4],
|
||||
const bf16* __restrict__ sV,
|
||||
int LD, int SWIZ_MASK, int lane,
|
||||
float Oacc[DN8][4])
|
||||
{
|
||||
#pragma unroll
|
||||
for (int kt2 = 0; kt2 < KT2; kt2++) {
|
||||
unsigned Pa[4];
|
||||
Pa[0] = pk2(Sacc[kt2 * 2][0], Sacc[kt2 * 2][1]);
|
||||
Pa[1] = pk2(Sacc[kt2 * 2][2], Sacc[kt2 * 2][3]);
|
||||
Pa[2] = pk2(Sacc[kt2 * 2 + 1][0], Sacc[kt2 * 2 + 1][1]);
|
||||
Pa[3] = pk2(Sacc[kt2 * 2 + 1][2], Sacc[kt2 * 2 + 1][3]);
|
||||
int vrow_l = kt2 * 16 + (lane & 15);
|
||||
#pragma unroll
|
||||
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
||||
unsigned b[2];
|
||||
ldmatrix_x2_trans(b, &sV[vrow_l * LD + swiz_col(dn8 * 8, vrow_l, SWIZ_MASK)]);
|
||||
mma16816(Oacc[dn8], Pa, b, Oacc[dn8]);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
#include "attn_paged_decode_split_kv.cuh"
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "attn_paged_decode_split_kv_mma.cuh"
|
||||
#endif
|
||||
|
||||
#include "attn_entry_utils.cuh"
|
||||
|
||||
static void launch_paged_scalar_decode(PagedAttentionParams<bf16>& p) {
|
||||
int group_size = p.q_head / p.kv_head;
|
||||
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||
alloc_split_partials(p);
|
||||
|
||||
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
dim3 grid = dim3(p.batch * p.kv_head, 1, p.num_splits);
|
||||
dim3 block = dim3(32, group_size);
|
||||
paged_attn_decode_split_kv_kernel<<<grid, block, smem>>>(p);
|
||||
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
|
||||
static void launch_paged_mma_decode(PagedAttentionParams<bf16>& p) {
|
||||
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||
alloc_split_partials(p);
|
||||
|
||||
paged_attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES><<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
|
||||
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
#endif
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static void dispatch_paged_decode(PagedAttentionParams<bf16>& p) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
int G = p.q_head / p.kv_head;
|
||||
if (G >= 1 && G <= 16 && p.page_size >= 32) {
|
||||
launch_paged_mma_decode<HEAD_DIM, 32>(p);
|
||||
return;
|
||||
}
|
||||
#endif
|
||||
launch_paged_scalar_decode(p);
|
||||
}
|
||||
|
||||
torch::Tensor attn_paged_decode(
|
||||
torch::Tensor q,
|
||||
torch::Tensor page_table,
|
||||
torch::Tensor k_cache,
|
||||
torch::Tensor v_cache,
|
||||
int64_t page_size,
|
||||
int64_t kv_len,
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
int64_t layout
|
||||
) {
|
||||
PagedAttentionParams<bf16> p;
|
||||
attn_pack_paged_params(q, page_table, k_cache, v_cache,
|
||||
page_size, kv_len, mask, causal_offset, scale, layout, p);
|
||||
|
||||
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
||||
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
||||
p.o = (bf16*)O_view.data_ptr();
|
||||
|
||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p);
|
||||
return O;
|
||||
}
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def("attn_paged_decode", &attn_paged_decode,
|
||||
py::arg("q"),
|
||||
py::arg("page_table"),
|
||||
py::arg("k_cache"),
|
||||
py::arg("v_cache"),
|
||||
py::arg("page_size"),
|
||||
py::arg("kv_len"),
|
||||
py::arg("mask") = py::none(),
|
||||
py::arg("causal_offset") = -1,
|
||||
py::arg("scale") = 0.0,
|
||||
py::arg("layout") = 0,
|
||||
"Paged GQA decode — split-KV with direct page-table access.");
|
||||
}
|
||||
@@ -0,0 +1,147 @@
|
||||
#pragma once
|
||||
#include <cuda_bf16.h>
|
||||
#include <float.h>
|
||||
#include "attn_common.h"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
constexpr int PDC_CHUNK = 64;
|
||||
|
||||
__device__ inline float paged_warp_reduce_sum(float val) {
|
||||
for (int offset = 16; offset > 0; offset >>= 1)
|
||||
val += __shfl_xor_sync(0xFFFFFFFF, val, offset);
|
||||
return val;
|
||||
}
|
||||
|
||||
// Split-KV scalar decode: one warp per query head, grid.z partitions KV.
|
||||
__global__ void paged_attn_decode_split_kv_kernel(PagedAttentionParams<bf16> p) {
|
||||
int batch = blockIdx.x / p.kv_head;
|
||||
int kv_head = blockIdx.x % p.kv_head;
|
||||
int split = blockIdx.z;
|
||||
int group_size = blockDim.y;
|
||||
int q_head = kv_head * group_size + threadIdx.y;
|
||||
int lane = threadIdx.x;
|
||||
int hd_per_thread = p.head_dim / 32;
|
||||
|
||||
// Q: stride-based [batch, q_head, q_len=1, head_dim]
|
||||
float q_reg[8];
|
||||
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
||||
+ lane * hd_per_thread * p.q_stride_d;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
q_reg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]);
|
||||
|
||||
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
|
||||
|
||||
extern __shared__ __align__(16) bf16 k_smem[];
|
||||
|
||||
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
||||
int chunks_per_split = (chunks_total + p.num_splits - 1) / p.num_splits;
|
||||
int ch_begin = split * chunks_per_split;
|
||||
int ch_end = min(chunks_total, ch_begin + chunks_per_split);
|
||||
|
||||
const int mask_base = batch * p.mask_b_stride;
|
||||
|
||||
for (int ci = ch_begin; ci < ch_end; ci++) {
|
||||
int chunk_start = ci * PDC_CHUNK;
|
||||
int this_chunk = min(PDC_CHUNK, p.kv_len - chunk_start);
|
||||
|
||||
int total = this_chunk * p.head_dim;
|
||||
for (int i = threadIdx.y * 32 + lane; i < total; i += blockDim.x * blockDim.y) {
|
||||
int s = i / p.head_dim;
|
||||
int d_dim = i % p.head_dim;
|
||||
int pos = chunk_start + s;
|
||||
int logical_page = pos / p.page_size;
|
||||
int page_offset = pos % p.page_size;
|
||||
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
||||
if (phys_page >= 0) {
|
||||
int64_t off = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
|
||||
+ (int64_t)page_offset * p.kv_head * p.head_dim
|
||||
+ (int64_t)kv_head * p.head_dim
|
||||
+ d_dim;
|
||||
k_smem[i] = p.k_cache[off];
|
||||
} else {
|
||||
k_smem[i] = __float2bfloat16(0.0f);
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
for (int s = 0; s < this_chunk; s++) {
|
||||
float partial = 0.0f;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
partial += q_reg[i] * __bfloat162float(k_smem[s * p.head_dim + lane * hd_per_thread + i]);
|
||||
partial = paged_warp_reduce_sum(partial) * p.scale;
|
||||
|
||||
int kv_idx = chunk_start + s;
|
||||
if (p.use_mask && p.mask && !p.mask[mask_base + kv_idx])
|
||||
partial = -FLT_MAX;
|
||||
if (p.causal_offset >= 0 && kv_idx > p.causal_offset)
|
||||
partial = -FLT_MAX;
|
||||
|
||||
float new_m = fmaxf(m, partial);
|
||||
float alpha = expf(m - new_m);
|
||||
float beta = expf(partial - new_m);
|
||||
d = d * alpha + beta;
|
||||
|
||||
int pos = chunk_start + s;
|
||||
int logical_page = pos / p.page_size;
|
||||
int page_offset = pos % p.page_size;
|
||||
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
||||
if (phys_page >= 0) {
|
||||
int64_t v_base = (int64_t)phys_page * p.page_size * p.kv_head * p.head_dim
|
||||
+ (int64_t)page_offset * p.kv_head * p.head_dim
|
||||
+ (int64_t)kv_head * p.head_dim;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
acc_reg[i] = acc_reg[i] * alpha + __bfloat162float(p.v_cache[v_base + lane * hd_per_thread + i]) * beta;
|
||||
} else {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
acc_reg[i] = acc_reg[i] * alpha + 0.0f * beta;
|
||||
}
|
||||
m = new_m;
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
size_t bh = (size_t)batch * p.q_head + q_head;
|
||||
size_t slot = bh * p.num_splits + split;
|
||||
int d0 = lane * hd_per_thread;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < hd_per_thread; i++)
|
||||
p.o_part[slot * p.head_dim + (d0 + i)] = acc_reg[i];
|
||||
if (lane == 0) {
|
||||
p.ml_part[slot * 2] = m;
|
||||
p.ml_part[slot * 2 + 1] = d;
|
||||
}
|
||||
}
|
||||
|
||||
__global__ void paged_attn_decode_combine_kernel(PagedAttentionParams<bf16> p) {
|
||||
int bh = blockIdx.x;
|
||||
int d = threadIdx.x;
|
||||
if (d >= p.head_dim) return;
|
||||
|
||||
int batch = bh / p.q_head;
|
||||
int q_head = bh % p.q_head;
|
||||
|
||||
size_t split_base = (size_t)bh * p.num_splits;
|
||||
const float* mlp = p.ml_part + split_base * 2;
|
||||
const float* op = p.o_part + split_base * p.head_dim;
|
||||
|
||||
float m = -FLT_MAX, l = 0.0f, acc = 0.0f;
|
||||
for (int s = 0; s < p.num_splits; s++) {
|
||||
float mi = mlp[s * 2];
|
||||
if (mi <= -FLT_MAX) continue;
|
||||
float li = mlp[s * 2 + 1];
|
||||
float nm = fmaxf(m, mi);
|
||||
float corr = __expf(m - nm);
|
||||
float e = __expf(mi - nm);
|
||||
acc = acc * corr + op[s * p.head_dim + d] * e;
|
||||
l = l * corr + li * e;
|
||||
m = nm;
|
||||
}
|
||||
|
||||
float inv = (l > 1e-20f) ? (1.0f / l) : 0.0f;
|
||||
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h + d * p.q_stride_d;
|
||||
p.o[o_off] = __float2bfloat16(acc * inv);
|
||||
}
|
||||
@@ -0,0 +1,170 @@
|
||||
#pragma once
|
||||
#include <cfloat>
|
||||
#include <cuda_bf16.h>
|
||||
#include "attn_common.h"
|
||||
#include "attn_mma_utils.cuh"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
// Paged split-KV tensor-core decode via GQA head-packing.
|
||||
// Identical algorithm to attn_decode_split_kv_mma_kernel but reads K/V
|
||||
// directly from the page pool through a page table, eliminating the gather
|
||||
// copy. Each tile (BC=32) fits within a single page (page_size >= 32), so
|
||||
// the page-table lookup happens once per tile for cp.async.
|
||||
|
||||
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
|
||||
__global__ void paged_attn_decode_split_kv_mma_kernel(PagedAttentionParams<bf16> p) {
|
||||
constexpr int KD = HEAD_DIM / 16;
|
||||
constexpr int NC8 = BC / 8;
|
||||
constexpr int KT2 = BC / 16;
|
||||
constexpr int DN8 = HEAD_DIM / 8;
|
||||
constexpr int LD = HEAD_DIM;
|
||||
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1);
|
||||
constexpr int VEC = 8;
|
||||
constexpr int TOTAL = BC * HEAD_DIM;
|
||||
|
||||
const int lane = threadIdx.x;
|
||||
const int gid = lane >> 2;
|
||||
const int tid4 = lane & 3;
|
||||
|
||||
const int kv_head = blockIdx.x;
|
||||
const int batch = blockIdx.y;
|
||||
const int split = blockIdx.z;
|
||||
const int G = p.q_head / p.kv_head;
|
||||
const int q_head0 = kv_head * G;
|
||||
|
||||
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
|
||||
__shared__ __align__(16) bf16 sV[STAGES * BC * LD];
|
||||
|
||||
// ---- Load Q directly from global into mma A-operand registers ----
|
||||
const int q_base = batch * p.q_stride_b + q_head0 * p.q_stride_h;
|
||||
const int qra = gid;
|
||||
const int qrb = gid + 8;
|
||||
const bool va = qra < G, vb = qrb < G;
|
||||
unsigned Qa[KD][4];
|
||||
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_h, p.q_stride_d,
|
||||
qra, qrb, va, vb, tid4, Qa);
|
||||
|
||||
float Oacc[DN8][4];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < DN8; j++)
|
||||
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||
|
||||
const int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||
const int tiles_per_split = (tiles_total + p.num_splits - 1) / p.num_splits;
|
||||
const int ti_begin = split * tiles_per_split;
|
||||
const int ti_end = min(tiles_total, ti_begin + tiles_per_split);
|
||||
const int has_mask = p.use_mask && p.mask;
|
||||
|
||||
// Paged strides (constant for the block)
|
||||
const int64_t page_stride = (int64_t)p.page_size * p.kv_head * HEAD_DIM;
|
||||
const int64_t pos_stride = (int64_t)p.kv_head * HEAD_DIM;
|
||||
const int64_t head_off = (int64_t)kv_head * HEAD_DIM;
|
||||
|
||||
// ---- Load tile lambda: predicated cp.async, paged addressing ----
|
||||
auto load_tile = [&](int ti, int buf) {
|
||||
int kv0 = ti * BC;
|
||||
bf16* dK = sK + buf * BC * LD;
|
||||
bf16* dV = sV + buf * BC * LD;
|
||||
int logical_page = kv0 / p.page_size;
|
||||
int phys_page = p.page_table[batch * p.max_pages + logical_page];
|
||||
bool page_valid = (phys_page >= 0);
|
||||
#pragma unroll
|
||||
for (int i = lane * VEC; i < TOTAL; i += 32 * VEC) {
|
||||
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
||||
int kc = kv0 + r;
|
||||
bool valid = (kc < p.kv_len) && page_valid;
|
||||
int page_off = kc % p.page_size;
|
||||
int64_t gmem_base = (int64_t)phys_page * page_stride
|
||||
+ (int64_t)page_off * pos_stride
|
||||
+ head_off;
|
||||
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
|
||||
cp_async_16_pred(&dK[off], &p.k_cache[gmem_base + d], valid);
|
||||
cp_async_16_pred(&dV[off], &p.v_cache[gmem_base + d], valid);
|
||||
}
|
||||
cp_async_commit();
|
||||
};
|
||||
|
||||
// ---- Prologue: issue first tile load ----
|
||||
if (ti_begin < ti_end) {
|
||||
load_tile(ti_begin, 0);
|
||||
}
|
||||
|
||||
for (int ti = ti_begin; ti < ti_end; ti++) {
|
||||
constexpr int BUF_MASK = (STAGES > 1) ? (STAGES - 1) : 0;
|
||||
int buf = (ti - ti_begin) & BUF_MASK;
|
||||
|
||||
cp_async_wait_group<0>();
|
||||
__syncwarp();
|
||||
if constexpr (STAGES > 1) {
|
||||
if (ti + 1 < ti_end)
|
||||
load_tile(ti + 1, (ti + 1 - ti_begin) & BUF_MASK);
|
||||
}
|
||||
|
||||
const bf16* bK = sK + buf * BC * LD;
|
||||
const bf16* bV = sV + buf * BC * LD;
|
||||
int kv0 = ti * BC;
|
||||
|
||||
float Sacc[NC8][4];
|
||||
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
|
||||
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < NC8; n8++)
|
||||
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||
|
||||
// Decode: q_len=1, so qrow0=qrow1=0, mask_q_stride irrelevant
|
||||
int maxc = (p.causal_offset >= 0) ? min(p.kv_len, p.causal_offset + 1) : p.kv_len;
|
||||
mma_softmax_tile<NC8, DN8>(kv0, maxc, maxc,
|
||||
0, 0,
|
||||
p.mask_b_stride, 0,
|
||||
batch,
|
||||
p.mask, has_mask,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
|
||||
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
|
||||
__syncwarp();
|
||||
|
||||
if constexpr (STAGES == 1) {
|
||||
if (ti + 1 < ti_end)
|
||||
load_tile(ti + 1, 0);
|
||||
}
|
||||
}
|
||||
|
||||
// ---- write UN-normalised partials for this split ----
|
||||
auto split_slot = [&](int h) -> size_t {
|
||||
size_t bh = (size_t)batch * p.q_head + h;
|
||||
return bh * p.num_splits + split;
|
||||
};
|
||||
#pragma unroll
|
||||
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
||||
int d = dn8 * 8 + 2 * tid4;
|
||||
int r0 = gid, r1 = gid + 8;
|
||||
if (r0 < G) {
|
||||
int h = q_head0 + r0;
|
||||
float* op = p.o_part + split_slot(h) * HEAD_DIM;
|
||||
op[d] = Oacc[dn8][0];
|
||||
op[d + 1] = Oacc[dn8][1];
|
||||
}
|
||||
if (r1 < G) {
|
||||
int h = q_head0 + r1;
|
||||
float* op = p.o_part + split_slot(h) * HEAD_DIM;
|
||||
op[d] = Oacc[dn8][2];
|
||||
op[d + 1] = Oacc[dn8][3];
|
||||
}
|
||||
}
|
||||
if (tid4 == 0) {
|
||||
int r0 = gid, r1 = gid + 8;
|
||||
if (r0 < G) {
|
||||
int h = q_head0 + r0;
|
||||
float* mp = p.ml_part + split_slot(h) * 2;
|
||||
mp[0] = m0; mp[1] = l0;
|
||||
}
|
||||
if (r1 < G) {
|
||||
int h = q_head0 + r1;
|
||||
float* mp = p.ml_part + split_slot(h) * 2;
|
||||
mp[0] = m1; mp[1] = l1;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,64 @@
|
||||
#include "attn_prefill_split_q.cuh"
|
||||
#include "attn_entry_utils.cuh"
|
||||
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "attn_prefill_split_q_mma.cuh"
|
||||
#endif
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static void dispatch_prefill(AttentionParams<bf16>& p) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
constexpr int WARPS = 4, BR = 16;
|
||||
// KV tile: bigger tiles amortize the per-tile cp.async wait + barrier +
|
||||
// loop overhead over more tensor-core work (this kernel is latency-bound,
|
||||
// not compute/bandwidth-bound), so BC=32 wins ~6-8% over BC=16 for
|
||||
// D<=128. D=256 stays at 16: BC=32 double-buffered would need 64KB smem,
|
||||
// over the 48KB static cap.
|
||||
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
||||
dim3 grid((p.q_len + BR * WARPS - 1) / (BR * WARPS), p.q_head, p.batch);
|
||||
dim3 block(WARPS * 32, 1, 1);
|
||||
// Static shared memory — double-buffered K/V only (no sQ: Q goes direct
|
||||
// to registers). 2*BC*LD bf16 each for sK and sV → 4*BC*HEAD_DIM*2 bytes.
|
||||
// Occupancy is smem-capped: D=64→3 blocks/SM (16KB), D=128→1 (32KB),
|
||||
// D=256→1 (32KB, BC=16).
|
||||
attn_prefill_split_q_mma_kernel<HEAD_DIM, WARPS, BC><<<grid, block>>>(p);
|
||||
#else
|
||||
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
||||
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
|
||||
dim3 block(G, ROWS, 1);
|
||||
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC><<<grid, block>>>(p);
|
||||
#endif
|
||||
}
|
||||
|
||||
torch::Tensor attn_prefill(
|
||||
torch::Tensor q,
|
||||
torch::Tensor k,
|
||||
torch::Tensor v,
|
||||
c10::optional<torch::Tensor> mask,
|
||||
int64_t causal_offset,
|
||||
double scale,
|
||||
int64_t layout
|
||||
) {
|
||||
AttentionParams<bf16> p;
|
||||
attn_pack_params(q, k, v, mask, causal_offset, scale, layout, p);
|
||||
TORCH_CHECK(p.head_dim % 16 == 0, "head_dim must be multiple of 16");
|
||||
|
||||
auto O = torch::empty_strided(q.sizes(), q.strides(), q.options());
|
||||
auto O_view = (layout == 1) ? O.transpose(1, 2) : O;
|
||||
p.o = (bf16*)O_view.data_ptr();
|
||||
|
||||
DISPATCH_HEAD_DIM(p.head_dim, dispatch_prefill, p);
|
||||
return O;
|
||||
}
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def("attn_prefill", &attn_prefill,
|
||||
py::arg("q"),
|
||||
py::arg("k"),
|
||||
py::arg("v"),
|
||||
py::arg("mask") = py::none(),
|
||||
py::arg("causal_offset") = -1,
|
||||
py::arg("scale") = 0.0,
|
||||
py::arg("layout") = 0,
|
||||
"GQA prefill (tensor-core mma on sm_80+, scalar fallback)");
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
#pragma once
|
||||
#include <cfloat>
|
||||
#include <cuda_bf16.h>
|
||||
#include "attn_common.h"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
// v9: group-split register blocking. G threads cooperate on one query row,
|
||||
// each owning HEAD_DIM/G dims of qreg[]/acc[]. Small per-thread footprint keeps
|
||||
// occupancy high; the S dot product is reduced across the G-lane group with a
|
||||
// short shuffle chain (log2(G) shuffles) instead of a full 32-lane warp reduce.
|
||||
// Online (per-kv) softmax — cheap because acc[] is only HEAD_DIM/G long.
|
||||
// Templated on <HEAD_DIM, G, ROWS, P_BC>. Block = (G, ROWS). G power-of-two,
|
||||
// G*ROWS a multiple of 32 with groups warp-aligned.
|
||||
|
||||
template <int G>
|
||||
__device__ __forceinline__ float group_reduce_sum(float v, unsigned mask) {
|
||||
#pragma unroll
|
||||
for (int o = G / 2; o > 0; o >>= 1)
|
||||
v += __shfl_xor_sync(mask, v, o);
|
||||
return v;
|
||||
}
|
||||
|
||||
// load 8 contiguous bf16 from (16-byte aligned) smem as one float4, unpack to
|
||||
// 8 floats — cuts shared-load instructions 8x vs scalar bf16 loads.
|
||||
__device__ __forceinline__ void ld8(const bf16* p, float* o) {
|
||||
float4 raw = *reinterpret_cast<const float4*>(p);
|
||||
const __nv_bfloat162* h = reinterpret_cast<const __nv_bfloat162*>(&raw);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < 4; j++) {
|
||||
float2 f = __bfloat1622float2(h[j]);
|
||||
o[2 * j] = f.x;
|
||||
o[2 * j + 1] = f.y;
|
||||
}
|
||||
}
|
||||
|
||||
template <int HEAD_DIM, int G, int ROWS, int P_BC>
|
||||
__global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
||||
constexpr int DPT = HEAD_DIM / G;
|
||||
|
||||
int q_tile = blockIdx.x;
|
||||
int q_head = blockIdx.y;
|
||||
int batch = blockIdx.z;
|
||||
int gpos = threadIdx.x; // 0..G-1 (which d-chunk)
|
||||
int row = threadIdx.y; // 0..ROWS-1
|
||||
int q_row = q_tile * ROWS + row;
|
||||
|
||||
int kv_head = q_head / (p.q_head / p.kv_head);
|
||||
|
||||
__shared__ __align__(16) bf16 sK[P_BC * HEAD_DIM];
|
||||
__shared__ __align__(16) bf16 sV[P_BC * HEAD_DIM];
|
||||
|
||||
// Q: stride-based load [batch, q_head, q_len, head_dim]
|
||||
float qreg[DPT];
|
||||
if (q_row < p.q_len) {
|
||||
int q_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
||||
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < DPT; i++)
|
||||
qreg[i] = __bfloat162float(p.q[q_off + i * p.q_stride_d]) * p.scale;
|
||||
}
|
||||
|
||||
float m = -FLT_MAX, l = 0.0f;
|
||||
float acc[DPT];
|
||||
#pragma unroll
|
||||
for (int i = 0; i < DPT; i++)
|
||||
acc[i] = 0.0f;
|
||||
|
||||
// KV: stride-based base
|
||||
int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||
int mask_batch_base = batch * p.mask_b_stride;
|
||||
int tiles = (p.kv_len + P_BC - 1) / P_BC;
|
||||
int tt = G * ROWS;
|
||||
int lid = row * G + gpos;
|
||||
|
||||
// per-group shuffle mask: only the G lanes of this row's group participate,
|
||||
// so causal masking (differing loop bounds across rows in a warp) is safe.
|
||||
int lane_in_warp = lid & 31;
|
||||
unsigned gmask = (G == 32) ? 0xFFFFFFFFu
|
||||
: (((1u << G) - 1u) << (lane_in_warp & ~(G - 1)));
|
||||
|
||||
for (int ti = 0; ti < tiles; ti++) {
|
||||
int kv0 = ti * P_BC;
|
||||
int tlen = min(P_BC, p.kv_len - kv0);
|
||||
|
||||
// Load K/V into shared memory from strided global
|
||||
for (int i = lid; i < tlen * HEAD_DIM; i += tt) {
|
||||
int s = i / HEAD_DIM;
|
||||
int d_dim = i % HEAD_DIM;
|
||||
int kv_idx = kv0 + s;
|
||||
int g_off = kv_base + kv_idx * p.kv_stride_l + d_dim * p.kv_stride_d;
|
||||
sK[i] = p.k[g_off];
|
||||
sV[i] = p.v[g_off];
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
int lim = tlen;
|
||||
if (p.causal_offset >= 0 && q_row < p.q_len) {
|
||||
int ep = q_row + p.causal_offset + 1;
|
||||
if (kv0 >= ep)
|
||||
lim = 0;
|
||||
else if (kv0 + tlen > ep)
|
||||
lim = ep - kv0;
|
||||
}
|
||||
|
||||
int mask_row_base = mask_batch_base + q_row * p.mask_q_stride;
|
||||
for (int s = 0; s < lim; s++) {
|
||||
const bf16* kr = sK + s * HEAD_DIM + gpos * DPT;
|
||||
float part = 0.0f;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < DPT; i += 8) {
|
||||
float k8[8];
|
||||
ld8(kr + i, k8);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < 8; j++)
|
||||
part = fmaf(qreg[i + j], k8[j], part);
|
||||
}
|
||||
float dot = group_reduce_sum<G>(part, gmask);
|
||||
|
||||
int kv_idx = kv0 + s;
|
||||
if (p.use_mask && p.mask && !p.mask[mask_row_base + kv_idx])
|
||||
dot = -FLT_MAX;
|
||||
|
||||
float nm = fmaxf(m, dot);
|
||||
float al = __expf(m - nm);
|
||||
float be = __expf(dot - nm);
|
||||
l = l * al + be;
|
||||
|
||||
const bf16* vr = sV + s * HEAD_DIM + gpos * DPT;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < DPT; i += 8) {
|
||||
float v8[8];
|
||||
ld8(vr + i, v8);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < 8; j++)
|
||||
acc[i + j] = fmaf(v8[j], be, acc[i + j] * al);
|
||||
}
|
||||
m = nm;
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
if (q_row < p.q_len) {
|
||||
// O: stride-based write
|
||||
int o_off = batch * p.q_stride_b + q_head * p.q_stride_h
|
||||
+ q_row * p.q_stride_l + gpos * DPT * p.q_stride_d;
|
||||
float rl = (l > 1e-10f) ? (1.0f / l) : 0.0f;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < DPT; i++)
|
||||
p.o[o_off + i * p.q_stride_d] = __float2bfloat16(acc[i] * rl);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
#pragma once
|
||||
#include <cfloat>
|
||||
#include <cuda_bf16.h>
|
||||
#include "attn_common.h"
|
||||
#include "attn_mma_utils.cuh"
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
// Tensor-core prefill flash attention (raw mma.sync PTX).
|
||||
// One warp owns BR=16 query rows. S = Q@K^T and O = P@V run on bf16 tensor
|
||||
// cores via mma.sync.m16n8k16 (f32 accumulate). Q fragments are loaded once
|
||||
// straight from global into the mma A-operand layout (no smem staging) and
|
||||
// kept resident in registers across the tile loop. S, O, and the online-softmax
|
||||
// stats (m, l) also live in registers.
|
||||
// Shared memory is statically sized via template parameters — no dynamic
|
||||
// allocation. The mma fragment layout is used directly: the S accumulator
|
||||
// (f32) maps element-for-element onto the P matrix_a (bf16) operand, so
|
||||
// softmax needs no shuffle repack; row reductions fold across the 4-lane
|
||||
// thread group. Templated on <HEAD_DIM, WARPS, BC> with BC a multiple of 16.
|
||||
//
|
||||
// Software pipeline: K/V are double-buffered and loaded via cp.async one tile
|
||||
// ahead, so the next tile streams from global memory while the current tile's
|
||||
// tensor-core math runs — hiding load latency (long_scoreboard). A single
|
||||
// __syncthreads per tile both publishes the freshly loaded tile cross-warp and
|
||||
// (because it runs before the next prefetch) guards the buffer being refilled,
|
||||
// so no second barrier is needed. Predicated cp.async (cp_async_16_pred)
|
||||
// zero-fills rows past kv_len, unifying full and partial tiles on one path.
|
||||
// BC=32 (D<=128) amortizes the per-tile wait+barrier+loop overhead over more
|
||||
// tensor-core work — this kernel is latency-bound (low occupancy from high
|
||||
// register pressure), so fewer, larger tiles beat many tiny ones.
|
||||
//
|
||||
// Optimizations: load Q fragments directly from global in mma A-operand layout
|
||||
// (no sQ staging, no prologue barriers); post-multiply scale in float after
|
||||
// S=Q@K^T to avoid bf16 precision loss; packed bf16x2 output stores;
|
||||
// causal tile skipping (block-level prefetch bound + warp-level compute skip);
|
||||
// XOR swizzle (swiz_col) → eliminates ldmatrix bank conflicts without LD
|
||||
// padding (LD=HEAD_DIM).
|
||||
|
||||
template <int HEAD_DIM, int WARPS, int BC>
|
||||
__global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
||||
constexpr int BR = 16;
|
||||
constexpr int KD = HEAD_DIM / 16; // Q/K k-tiles
|
||||
constexpr int NC8 = BC / 8; // S n-tiles (N=8 each)
|
||||
constexpr int KT2 = BC / 16; // P k-tiles (K=16 each)
|
||||
constexpr int DN8 = HEAD_DIM / 8; // O n-tiles (N=8 each)
|
||||
constexpr int LD = HEAD_DIM; // XOR swizzle (swiz_col) handles bank conflicts
|
||||
constexpr int SWIZ_MASK = (HEAD_DIM >= 64) ? 7 : (HEAD_DIM / 8 - 1); // chunk bits, stay within LD
|
||||
|
||||
const int warp = threadIdx.x / 32;
|
||||
const int lane = threadIdx.x % 32;
|
||||
const int gid = lane >> 2; // 0..7 → rows gid, gid+8
|
||||
const int tid4 = lane & 3; // 0..3
|
||||
const int nthreads = WARPS * 32;
|
||||
|
||||
const int q_head = blockIdx.y;
|
||||
const int batch = blockIdx.z;
|
||||
const int kv_head = q_head / (p.q_head / p.kv_head);
|
||||
const int qrow0 = (blockIdx.x * WARPS + warp) * BR;
|
||||
|
||||
// ---- Static shared memory: double-buffered K/V ----
|
||||
// K/V are double-buffered (STAGES=2): the next tile's cp.async load runs
|
||||
// while the current tile's tensor-core math executes, hiding global-load
|
||||
// latency (FA2-style software pipeline). No dynamic smem / carveout opt-in.
|
||||
constexpr int STAGES = 2;
|
||||
__shared__ __align__(16) bf16 sK[STAGES * BC * LD];
|
||||
__shared__ __align__(16) bf16 sV[STAGES * BC * LD];
|
||||
|
||||
// Load Q fragments straight from global into mma A-operand layout.
|
||||
// stride_row = p.q_stride_l for prefill (multi-q rows across q_len).
|
||||
// See attn_mma_utils.cuh for the shared template.
|
||||
const int q_base = batch * p.q_stride_b + q_head * p.q_stride_h;
|
||||
const int qra = qrow0 + gid;
|
||||
const int qrb = qrow0 + gid + 8;
|
||||
const bool va = qra < p.q_len, vb = qrb < p.q_len;
|
||||
unsigned Qa[KD][4];
|
||||
load_q_mma_frags<KD>(p.q + q_base, p.q_stride_l, p.q_stride_d,
|
||||
qra, qrb, va, vb, tid4, Qa);
|
||||
|
||||
float Oacc[DN8][4];
|
||||
#pragma unroll
|
||||
for (int j = 0; j < DN8; j++)
|
||||
Oacc[j][0] = Oacc[j][1] = Oacc[j][2] = Oacc[j][3] = 0.0f;
|
||||
float m0 = -FLT_MAX, m1 = -FLT_MAX, l0 = 0.0f, l1 = 0.0f;
|
||||
|
||||
// KV: stride-based base
|
||||
const int kv_base = batch * p.kv_stride_b + kv_head * p.kv_stride_h;
|
||||
const int tiles = (p.kv_len + BC - 1) / BC;
|
||||
const int qr0 = qrow0 + gid; // row for c0/c1
|
||||
const int qr1 = qrow0 + gid + 8; // row for c2/c3
|
||||
|
||||
// Causal tile-skip bounds (no-op when causal_offset < 0)
|
||||
const int use_skip = (p.causal_offset >= 0) ? 1 : 0;
|
||||
const int max_kv = qrow0 + BR - 1 + p.causal_offset;
|
||||
const int block_max_kv =
|
||||
blockIdx.x * WARPS * BR + WARPS * BR - 1 + p.causal_offset;
|
||||
const int has_mask = p.use_mask && p.mask;
|
||||
|
||||
// Last active tile: block-level causal bound (all warps in the block share
|
||||
// the K/V load, so the prefetch range is the block max, not per-warp).
|
||||
int t_end = tiles - 1;
|
||||
if (use_skip) {
|
||||
int bt = block_max_kv / BC;
|
||||
if (bt < t_end) t_end = bt;
|
||||
}
|
||||
|
||||
constexpr int VEC = 8; // bf16 per cp.async unit (16 bytes)
|
||||
constexpr int TOTAL = BC * HEAD_DIM;
|
||||
|
||||
// ---- Load tile lambda: predicated cp.async ----
|
||||
// Issue cp.async loads for tile `ti` into shared buffer `buf`. Predicated
|
||||
// loads zero-fill rows past kv_len, so partial tiles need no scalar path.
|
||||
auto load_tile = [&](int ti, int buf) {
|
||||
int kv0 = ti * BC;
|
||||
bf16* dK = sK + buf * BC * LD;
|
||||
bf16* dV = sV + buf * BC * LD;
|
||||
#pragma unroll
|
||||
for (int i = threadIdx.x * VEC; i < TOTAL; i += nthreads * VEC) {
|
||||
int r = i / HEAD_DIM, d = i % HEAD_DIM;
|
||||
int kc = kv0 + r;
|
||||
bool valid = kc < p.kv_len;
|
||||
int off = r * LD + swiz_col(d, r, SWIZ_MASK);
|
||||
int g_off = kv_base + kc * p.kv_stride_l + d * p.kv_stride_d;
|
||||
cp_async_16_pred(&dK[off], &p.k[g_off], valid);
|
||||
cp_async_16_pred(&dV[off], &p.v[g_off], valid);
|
||||
}
|
||||
cp_async_commit();
|
||||
};
|
||||
|
||||
// ---- Prologue: issue first tile load ----
|
||||
load_tile(0, 0);
|
||||
|
||||
for (int ti = 0; ti <= t_end; ti++) {
|
||||
int buf = ti & 1;
|
||||
|
||||
// Wait for the current tile's async copies, then a single barrier: it
|
||||
// both publishes this tile's data cross-warp AND guarantees the prior
|
||||
// compute on the buffer we are about to refill has finished. Issuing
|
||||
// the next tile's load *after* this barrier lets one barrier cover both
|
||||
// hazards (vs two), while the load still overlaps this tile's math.
|
||||
cp_async_wait_group<0>();
|
||||
__syncthreads();
|
||||
if (ti < t_end) load_tile(ti + 1, (ti + 1) & 1);
|
||||
|
||||
const bf16* bK = sK + buf * BC * LD;
|
||||
const bf16* bV = sV + buf * BC * LD;
|
||||
int kv0 = ti * BC;
|
||||
|
||||
// Warp-level causal skip
|
||||
if (!use_skip || kv0 <= max_kv) {
|
||||
|
||||
// S = Q @ K^T + scale + online softmax + O += P @ V
|
||||
float Sacc[NC8][4];
|
||||
mma_compute_scores<KD, NC8>(Qa, bK, LD, SWIZ_MASK, lane, Sacc);
|
||||
|
||||
// post-multiply scale in float (no bf16 precision loss from pre-scaling Q)
|
||||
#pragma unroll
|
||||
for (int n8 = 0; n8 < NC8; n8++)
|
||||
Sacc[n8][0] *= p.scale, Sacc[n8][1] *= p.scale,
|
||||
Sacc[n8][2] *= p.scale, Sacc[n8][3] *= p.scale;
|
||||
|
||||
int maxc0 = (p.causal_offset >= 0) ? min(p.kv_len, qr0 + p.causal_offset + 1)
|
||||
: p.kv_len;
|
||||
int maxc1 = (p.causal_offset >= 0) ? min(p.kv_len, qr1 + p.causal_offset + 1)
|
||||
: p.kv_len;
|
||||
mma_softmax_tile<NC8, DN8>(kv0, maxc0, maxc1,
|
||||
qr0, qr1,
|
||||
p.mask_b_stride, p.mask_q_stride,
|
||||
batch,
|
||||
p.mask, has_mask,
|
||||
Sacc, Oacc, m0, m1, l0, l1, lane);
|
||||
|
||||
mma_pv_accumulate<DN8, KT2>(Sacc, bV, LD, SWIZ_MASK, lane, Oacc);
|
||||
} // if active (warp-level causal skip)
|
||||
}
|
||||
|
||||
// ---- write output ---- (packed bf16x2 stores: one 32-bit STG per pair,
|
||||
// halves store count and removes the uncoalesced scalar-store penalty)
|
||||
float rl0 = (l0 > 1e-20f) ? (1.0f / l0) : 0.0f;
|
||||
float rl1 = (l1 > 1e-20f) ? (1.0f / l1) : 0.0f;
|
||||
// O: stride-based write
|
||||
const int o_base = batch * p.q_stride_b + q_head * p.q_stride_h;
|
||||
#pragma unroll
|
||||
for (int dn8 = 0; dn8 < DN8; dn8++) {
|
||||
int d = dn8 * 8 + 2 * tid4;
|
||||
if (qr0 < p.q_len) {
|
||||
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][0] * rl0,
|
||||
Oacc[dn8][1] * rl0);
|
||||
*reinterpret_cast<__nv_bfloat162*>(&p.o[o_base + qr0 * p.q_stride_l + d * p.q_stride_d]) = v;
|
||||
}
|
||||
if (qr1 < p.q_len) {
|
||||
__nv_bfloat162 v = __floats2bfloat162_rn(Oacc[dn8][2] * rl1,
|
||||
Oacc[dn8][3] * rl1);
|
||||
*reinterpret_cast<__nv_bfloat162*>(&p.o[o_base + qr1 * p.q_stride_l + d * p.q_stride_d]) = v;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,201 @@
|
||||
/*
|
||||
Pure-C test:
|
||||
nvcc -I csrc -arch=sm_89 -O3 \
|
||||
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
|
||||
csrc/tests/attn_decode_test.cu -o test && ./test
|
||||
*/
|
||||
|
||||
#include "test_utils.cuh"
|
||||
#include "../kernels/attn_decode_split_kv.cuh"
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "../kernels/attn_decode_split_kv_mma.cuh"
|
||||
#endif
|
||||
|
||||
// Split-K scratch (torch-free): the production launcher allocates these from
|
||||
// torch; here we pass pre-allocated device buffers so the bench loop doesn't
|
||||
// pay a cudaMalloc per iteration. Size for the maximum split count (32).
|
||||
struct DecodeScratch {
|
||||
float* o_part = nullptr;
|
||||
float* ml_part = nullptr;
|
||||
};
|
||||
|
||||
// Launch the production decode path (tensor-core head-packing MMA on sm_80+,
|
||||
// scalar fallback otherwise), mirroring dispatch_decode() in attn_decode.cu.
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
static bool decode_use_mma(const AttentionParams<bf16>& p) {
|
||||
int G = p.q_head / p.kv_head;
|
||||
return !p.use_mask && G > 1 && G <= 16;
|
||||
}
|
||||
|
||||
template <int HEAD_DIM, int BC, int STAGES = (HEAD_DIM <= 128) ? 2 : 1>
|
||||
static void launch_mma_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||
int tiles_total = (p.kv_len + BC - 1) / BC;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||
p.o_part = sc.o_part;
|
||||
p.ml_part = sc.ml_part;
|
||||
|
||||
attn_decode_split_kv_mma_kernel<HEAD_DIM, BC, STAGES>
|
||||
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
|
||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
#endif
|
||||
|
||||
static void launch_scalar_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||
int gs = p.q_head / p.kv_head;
|
||||
int chunks_total = (p.kv_len + DC_CHUNK - 1) / DC_CHUNK;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||
p.o_part = sc.o_part;
|
||||
p.ml_part = sc.ml_part;
|
||||
|
||||
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
attn_decode_split_kv_kernel<<<dim3(p.batch * p.kv_head, 1, p.num_splits), dim3(32, gs), smem>>>(p);
|
||||
attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static void dispatch_decode_t(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
if (decode_use_mma(p)) { launch_mma_decode<HEAD_DIM, 32>(p, sc); return; }
|
||||
#endif
|
||||
launch_scalar_decode(p, sc);
|
||||
}
|
||||
|
||||
static void dispatch_decode(AttentionParams<bf16>& p, DecodeScratch& sc) {
|
||||
dispatch_by_head_dim(p.head_dim, [&]<int D>() { dispatch_decode_t<D>(p, sc); });
|
||||
}
|
||||
|
||||
// Warmed-up, CUDA-event timed sweep over the production decode MMA path.
|
||||
static void bench() {
|
||||
const int cfgs[][5] = {
|
||||
{1, 32, 4, 512, 128}, // B, Hq, Hk, kv_len, D
|
||||
{1, 32, 4, 1024, 128},
|
||||
{1, 32, 4, 2048, 128},
|
||||
{1, 32, 4, 4096, 128},
|
||||
{16, 32, 4, 2048, 128},
|
||||
{32, 32, 4, 1024, 128},
|
||||
};
|
||||
const int WARMUP = 10, ITERS = 100;
|
||||
printf("\n===== DECODE BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
|
||||
print_bench_header();
|
||||
|
||||
for (int ci = 0; ci < 6; ci++) {
|
||||
int B = cfgs[ci][0], Hq = cfgs[ci][1], Hk = cfgs[ci][2];
|
||||
int sl = cfgs[ci][3], D = cfgs[ci][4];
|
||||
size_t nQ = (size_t)B * Hq * D;
|
||||
size_t nKV = (size_t)B * Hk * sl * D;
|
||||
|
||||
bf16 *dQ, *dK, *dV, *dO;
|
||||
cudaMalloc(&dQ, nQ*2); cudaMalloc(&dK, nKV*2);
|
||||
cudaMalloc(&dV, nKV*2); cudaMalloc(&dO, nQ*2);
|
||||
size_t big = nQ > nKV ? nQ : nKV; bf16* tmp = new bf16[big];
|
||||
for (size_t i = 0; i < nQ; i++) tmp[i] = f2bf(randf());
|
||||
cudaMemcpy(dQ, tmp, nQ*2, cudaMemcpyHostToDevice);
|
||||
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
|
||||
cudaMemcpy(dK, tmp, nKV*2, cudaMemcpyHostToDevice);
|
||||
for (size_t i = 0; i < nKV; i++) tmp[i] = f2bf(randf());
|
||||
cudaMemcpy(dV, tmp, nKV*2, cudaMemcpyHostToDevice);
|
||||
delete[] tmp;
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
p.batch = B; p.q_head = Hq; p.kv_head = Hk; p.q_len = 1; p.kv_len = sl;
|
||||
p.head_dim = D; p.use_mask = 0; p.causal_offset = -1;
|
||||
p.scale = 1.0f / sqrtf((float)D);
|
||||
set_default_strides(p);
|
||||
p.q = dQ; p.k = dK; p.v = dV; p.mask = nullptr; p.o = dO;
|
||||
|
||||
DecodeScratch sc;
|
||||
cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float));
|
||||
cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float));
|
||||
|
||||
auto launch = [&]() { dispatch_decode(p, sc); };
|
||||
double flops = 4.0 * B * Hq * (double)sl * D;
|
||||
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
|
||||
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
|
||||
|
||||
char cfg[64];
|
||||
snprintf(cfg, sizeof(cfg),
|
||||
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
|
||||
B, Hq, Hk, 1, sl, D, 0);
|
||||
print_bench_row(cfg, r);
|
||||
|
||||
cudaFree(dQ); cudaFree(dK); cudaFree(dV); cudaFree(dO);
|
||||
cudaFree(sc.o_part); cudaFree(sc.ml_part);
|
||||
}
|
||||
}
|
||||
|
||||
int main() {
|
||||
const int configs[][5] = {
|
||||
{1, 2, 1, 64, 32}, // B,Hq,Hk,seq_len,D
|
||||
{1, 32, 4, 512, 128},
|
||||
{1, 32, 4, 1024, 128},
|
||||
};
|
||||
int n_cfgs = sizeof(configs) / sizeof(configs[0]);
|
||||
|
||||
for (int ci = 0; ci < n_cfgs; ci++) {
|
||||
int B = configs[ci][0], Hq = configs[ci][1], Hk = configs[ci][2];
|
||||
int sl = configs[ci][3], D = configs[ci][4], gs = Hq / Hk;
|
||||
printf("=== B=%d Hq=%d Hk=%d seq=%d D=%d gs=%d ===\n", B,Hq,Hk,sl,D,gs);
|
||||
|
||||
size_t nQ = B*Hq*1*D, nKV = B*Hk*sl*D;
|
||||
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
||||
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||
|
||||
bool* hMask=new bool[B*sl];
|
||||
for (int i=0;i<B*sl;i++) hMask[i]=true;
|
||||
|
||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||
bool* dMask;
|
||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||
cudaMalloc(&dMask,B*sl);
|
||||
|
||||
tmp=new bf16[max(nQ,nKV)];
|
||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
cudaMemcpy(dMask,hMask,B*sl,cudaMemcpyHostToDevice);
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=1; p.kv_len=sl; p.head_dim=D;
|
||||
p.use_mask=0; p.causal_offset=-1;
|
||||
p.scale=1.0f/sqrtf((float)D);
|
||||
set_default_strides(p);
|
||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
// Split-K scratch (max 32 splits), sized for the production MMA path.
|
||||
DecodeScratch sc;
|
||||
cudaMalloc(&sc.o_part, (size_t)B*Hq*32*D*sizeof(float));
|
||||
cudaMalloc(&sc.ml_part, (size_t)B*Hq*32*2*sizeof(float));
|
||||
|
||||
double t0=now_ms();
|
||||
dispatch_decode(p, sc);
|
||||
cudaDeviceSynchronize();
|
||||
double kms=now_ms()-t0;
|
||||
cudaError_t err=cudaGetLastError();
|
||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||
|
||||
bf16* hOut=new bf16[nQ];
|
||||
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
||||
|
||||
float* ref=new float[nQ];
|
||||
cpu_attention_ref(hQ, hK, hV, hMask, ref, B, Hq, Hk, 1, sl, D, -1);
|
||||
|
||||
float max_err=0;
|
||||
for (size_t i=0;i<nQ;i++){
|
||||
float d=fabsf(bf2f(hOut[i])-ref[i]);
|
||||
if(d>max_err) max_err=d;
|
||||
}
|
||||
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
|
||||
|
||||
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);cudaFree(dMask);
|
||||
cudaFree(sc.o_part);cudaFree(sc.ml_part);
|
||||
delete[]hQ;delete[]hK;delete[]hV;delete[]hMask;delete[]hOut;delete[]ref;delete[]tmp;
|
||||
}
|
||||
printf("All tests passed!\n");
|
||||
bench();
|
||||
return 0;
|
||||
}
|
||||
@@ -0,0 +1,332 @@
|
||||
// Compile:
|
||||
// nvcc -I csrc -arch=sm_89 -O3 --use_fast_math --ptxas-options=-O3 \
|
||||
// --extra-device-vectorization csrc/tests/attn_paged_decode_test.cu \
|
||||
// -o /tmp/test_paged && /tmp/test_paged
|
||||
|
||||
#include <cstring>
|
||||
#include "test_utils.cuh"
|
||||
#include "../kernels/attn_paged_decode_split_kv.cuh"
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "../kernels/attn_paged_decode_split_kv_mma.cuh"
|
||||
#endif
|
||||
|
||||
// Copy contiguous K/V from page pool (reference gather)
|
||||
static void gather_kv_cpu(
|
||||
const bf16* h_k_pool, const bf16* h_v_pool,
|
||||
const int64_t* h_pt, int B, int Hkv, int kv_len,
|
||||
int page_size, int head_dim,
|
||||
bf16* h_k, bf16* h_v)
|
||||
{
|
||||
int max_pages = (kv_len + page_size - 1) / page_size;
|
||||
size_t page_stride = (size_t)page_size * Hkv * head_dim;
|
||||
for (int b = 0; b < B; b++) {
|
||||
for (int pos = 0; pos < kv_len; pos++) {
|
||||
int log_pg = pos / page_size;
|
||||
int pg_off = pos % page_size;
|
||||
int phys = (int)h_pt[b * max_pages + log_pg];
|
||||
for (int h = 0; h < Hkv; h++) {
|
||||
size_t src_base = (size_t)phys * page_stride
|
||||
+ (size_t)pg_off * Hkv * head_dim
|
||||
+ h * head_dim;
|
||||
size_t dst_base = ((size_t)b * Hkv + h) * kv_len * head_dim + (size_t)pos * head_dim;
|
||||
memcpy(h_k + dst_base, h_k_pool + src_base, head_dim * sizeof(bf16));
|
||||
memcpy(h_v + dst_base, h_v_pool + src_base, head_dim * sizeof(bf16));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static void launch_paged_decode(PagedAttentionParams<bf16, float>& p) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
int G_check = p.q_head / p.kv_head;
|
||||
bool use_mma = !p.use_mask && G_check >= 1 && G_check <= 16 && p.page_size >= 32;
|
||||
if (use_mma) {
|
||||
constexpr int STAGES = (HEAD_DIM <= 128) ? 2 : 1;
|
||||
int tiles_total = (p.kv_len + 32 - 1) / 32;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, tiles_total);
|
||||
paged_attn_decode_split_kv_mma_kernel<HEAD_DIM, 32, STAGES>
|
||||
<<<dim3(p.kv_head, p.batch, p.num_splits), 32>>>(p);
|
||||
} else
|
||||
#endif
|
||||
{
|
||||
int group_sz = p.q_head / p.kv_head;
|
||||
int chunks_total = (p.kv_len + PDC_CHUNK - 1) / PDC_CHUNK;
|
||||
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
|
||||
size_t smem = PDC_CHUNK * p.head_dim * sizeof(bf16);
|
||||
paged_attn_decode_split_kv_kernel<<<
|
||||
dim3(p.batch * p.kv_head, 1, p.num_splits),
|
||||
dim3(32, group_sz), smem>>>(p);
|
||||
}
|
||||
paged_attn_decode_combine_kernel<<<p.batch * p.q_head, p.head_dim>>>(p);
|
||||
}
|
||||
|
||||
template <int HEAD_DIM>
|
||||
static int run_test(int B, int Hq, int Hkv, int kv_len, int page_size, int seed) {
|
||||
printf("B=%d Hq=%d Hkv=%d kv_len=%d page_sz=%d head_dim=%d ... ", B, Hq, Hkv, kv_len, page_size, HEAD_DIM);
|
||||
fflush(stdout);
|
||||
|
||||
int max_pages = (kv_len + page_size - 1) / page_size;
|
||||
int n_phys_pages = B * max_pages;
|
||||
|
||||
size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
|
||||
size_t sz_o = sz_q;
|
||||
size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16);
|
||||
size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t);
|
||||
int max_splits = 32;
|
||||
size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float);
|
||||
size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float);
|
||||
|
||||
bf16 *d_q, *d_o_paged, *d_o_ref;
|
||||
bf16 *d_k_pool, *d_v_pool;
|
||||
int64_t* d_pt;
|
||||
float *d_op, *d_ml;
|
||||
|
||||
cudaMalloc(&d_q, sz_q);
|
||||
cudaMalloc(&d_o_paged, sz_o);
|
||||
cudaMalloc(&d_o_ref, sz_o);
|
||||
cudaMalloc(&d_k_pool, sz_kv);
|
||||
cudaMalloc(&d_v_pool, sz_kv);
|
||||
cudaMalloc(&d_pt, sz_pt);
|
||||
cudaMalloc(&d_op, sz_op);
|
||||
cudaMalloc(&d_ml, sz_ml);
|
||||
|
||||
srand(seed);
|
||||
auto rnd = [&]() { return (rand() / (float)RAND_MAX) * 2.0f - 1.0f; };
|
||||
|
||||
bf16* h_q = (bf16*)malloc(sz_q);
|
||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
|
||||
h_q[i] = __float2bfloat16(rnd());
|
||||
cudaMemcpy(d_q, h_q, sz_q, cudaMemcpyHostToDevice);
|
||||
|
||||
bf16* h_k_pool = (bf16*)malloc(sz_kv);
|
||||
bf16* h_v_pool = (bf16*)malloc(sz_kv);
|
||||
size_t ps = (size_t)page_size * Hkv * HEAD_DIM;
|
||||
for (int pg = 0; pg < n_phys_pages; pg++) {
|
||||
for (int off = 0; off < page_size; off++) {
|
||||
for (int h = 0; h < Hkv; h++) {
|
||||
for (int d = 0; d < HEAD_DIM; d++) {
|
||||
float v = sinf((float)(pg * 7919 + off * 1049 + h * 331 + d));
|
||||
size_t idx = (size_t)pg * ps + (size_t)off * Hkv * HEAD_DIM + h * HEAD_DIM + d;
|
||||
h_k_pool[idx] = __float2bfloat16(v);
|
||||
h_v_pool[idx] = __float2bfloat16(v * 0.3f);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
cudaMemcpy(d_k_pool, h_k_pool, sz_kv, cudaMemcpyHostToDevice);
|
||||
cudaMemcpy(d_v_pool, h_v_pool, sz_kv, cudaMemcpyHostToDevice);
|
||||
|
||||
int64_t* h_pt = (int64_t*)malloc(sz_pt);
|
||||
int next_pg = 0;
|
||||
for (int b = 0; b < B; b++)
|
||||
for (int p = 0; p < max_pages; p++)
|
||||
h_pt[b * max_pages + p] = next_pg++;
|
||||
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
|
||||
|
||||
bf16* h_k_cont = (bf16*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(bf16));
|
||||
bf16* h_v_cont = (bf16*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(bf16));
|
||||
gather_kv_cpu(h_k_pool, h_v_pool, h_pt, B, Hkv, kv_len, page_size, HEAD_DIM, h_k_cont, h_v_cont);
|
||||
|
||||
float* h_q_f = (float*)malloc((size_t)B * Hq * HEAD_DIM * sizeof(float));
|
||||
float* h_k_f = (float*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(float));
|
||||
float* h_v_f = (float*)malloc((size_t)B * kv_len * Hkv * HEAD_DIM * sizeof(float));
|
||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++) h_q_f[i] = bf2f(h_q[i]);
|
||||
for (int i = 0; i < B * kv_len * Hkv * HEAD_DIM; i++) {
|
||||
h_k_f[i] = bf2f(h_k_cont[i]);
|
||||
h_v_f[i] = bf2f(h_v_cont[i]);
|
||||
}
|
||||
|
||||
float* h_o_ref = (float*)calloc(B * Hq * HEAD_DIM, sizeof(float));
|
||||
cpu_attention_ref(h_q_f, h_k_f, h_v_f, nullptr, h_o_ref, B, Hq, Hkv, 1, kv_len, HEAD_DIM, -1);
|
||||
|
||||
float scale_val = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
PagedAttentionParams<bf16, float> p;
|
||||
p.batch = B; p.q_head = Hq; p.kv_head = Hkv; p.q_len = 1;
|
||||
p.kv_len = kv_len; p.head_dim = HEAD_DIM;
|
||||
p.use_mask = 0; p.causal_offset = -1;
|
||||
set_default_paged_strides(p);
|
||||
p.num_splits = 1; p.scale = scale_val;
|
||||
p.page_size = page_size; p.max_pages = max_pages;
|
||||
p.page_table = d_pt;
|
||||
p.k_cache = d_k_pool; p.v_cache = d_v_pool;
|
||||
p.q = d_q; p.mask = nullptr; p.o = d_o_paged;
|
||||
p.o_part = d_op; p.ml_part = d_ml;
|
||||
|
||||
launch_paged_decode<HEAD_DIM>(p);
|
||||
cudaDeviceSynchronize();
|
||||
|
||||
bf16* h_o_bf16 = (bf16*)malloc(sz_o);
|
||||
cudaMemcpy(h_o_bf16, d_o_paged, sz_o, cudaMemcpyDeviceToHost);
|
||||
float* h_o_paged = (float*)malloc(B * Hq * HEAD_DIM * sizeof(float));
|
||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++)
|
||||
h_o_paged[i] = __bfloat162float(h_o_bf16[i]);
|
||||
|
||||
float max_err = 0.0f;
|
||||
int bad_idx = -1;
|
||||
for (int i = 0; i < B * Hq * HEAD_DIM; i++) {
|
||||
float e = fabsf(h_o_paged[i] - h_o_ref[i]);
|
||||
if (e > max_err) { max_err = e; bad_idx = i; }
|
||||
}
|
||||
|
||||
bool pass = max_err < 0.02f;
|
||||
|
||||
if (pass) {
|
||||
printf("PASS (max_abs_err=%.4e)\n", max_err);
|
||||
} else {
|
||||
int b = bad_idx / (Hq * HEAD_DIM);
|
||||
int h = (bad_idx / HEAD_DIM) % Hq;
|
||||
int d = bad_idx % HEAD_DIM;
|
||||
printf("FAIL (max_abs_err=%.4e at [%d,%d,%d]: ref=%.4f got=%.4f)\n",
|
||||
max_err, b, h, d, h_o_ref[bad_idx], h_o_paged[bad_idx]);
|
||||
printf(" ref[0..7]:");
|
||||
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
|
||||
printf(" %.4f", h_o_ref[i]);
|
||||
printf("\n got[0..7]:");
|
||||
for (int i = 0; i < 8 && i < HEAD_DIM; i++)
|
||||
printf(" %.4f", h_o_paged[i]);
|
||||
printf("\n");
|
||||
}
|
||||
|
||||
free(h_q); free(h_k_pool); free(h_v_pool); free(h_pt);
|
||||
free(h_k_cont); free(h_v_cont);
|
||||
free(h_q_f); free(h_k_f); free(h_v_f);
|
||||
free(h_o_ref); free(h_o_bf16); free(h_o_paged);
|
||||
cudaFree(d_q); cudaFree(d_o_paged); cudaFree(d_o_ref);
|
||||
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
|
||||
cudaFree(d_op); cudaFree(d_ml);
|
||||
|
||||
return pass ? 0 : 1;
|
||||
}
|
||||
|
||||
struct TestCase {
|
||||
int head_dim;
|
||||
int B, Hq, Hkv, kv_len, page_size, seed;
|
||||
};
|
||||
|
||||
static const TestCase TESTS[] = {
|
||||
{128, 1, 1, 1, 8, 128, 1},
|
||||
{128, 1, 4, 4, 128, 128, 2},
|
||||
{128, 2, 4, 4, 256, 128, 3},
|
||||
{128, 1, 4, 1, 64, 64, 4},
|
||||
{128, 1, 8, 2, 64, 128, 5},
|
||||
{128, 2, 16, 4, 128, 128, 6},
|
||||
{64, 1, 4, 2, 32, 128, 7},
|
||||
{256, 1, 2, 1, 16, 128, 8},
|
||||
{32, 1, 4, 2, 32, 64, 9},
|
||||
{128, 3, 8, 2, 256, 128, 10},
|
||||
{128, 2, 32, 8, 512, 128, 11},
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
{128, 1, 16, 2, 256, 128, 12},
|
||||
{128, 2, 32, 4, 512, 128, 13},
|
||||
#endif
|
||||
};
|
||||
|
||||
static int dispatch_test(const TestCase& tc) {
|
||||
bool matched = false;
|
||||
int r = 0;
|
||||
dispatch_by_head_dim(tc.head_dim, [&]<int D>() {
|
||||
matched = true;
|
||||
r = run_test<D>(tc.B, tc.Hq, tc.Hkv, tc.kv_len, tc.page_size, tc.seed);
|
||||
});
|
||||
return matched ? r : 1;
|
||||
}
|
||||
|
||||
// Warmed-up, CUDA-event timed sweep over paged decode configs.
|
||||
// Bytes = K + V read through page table (B*Hk*kv*D each), bf16.
|
||||
template <int HEAD_DIM>
|
||||
static void bench_config(int B, int Hq, int Hkv, int kv_len, int page_size) {
|
||||
int max_pages = (kv_len + page_size - 1) / page_size;
|
||||
int n_phys_pages = B * max_pages;
|
||||
|
||||
size_t sz_q = (size_t)B * Hq * 1 * HEAD_DIM * sizeof(bf16);
|
||||
size_t sz_kv = (size_t)n_phys_pages * page_size * Hkv * HEAD_DIM * sizeof(bf16);
|
||||
size_t sz_pt = (size_t)B * max_pages * sizeof(int64_t);
|
||||
int max_splits = 32;
|
||||
size_t sz_op = (size_t)B * Hq * max_splits * HEAD_DIM * sizeof(float);
|
||||
size_t sz_ml = (size_t)B * Hq * max_splits * 2 * sizeof(float);
|
||||
|
||||
bf16 *d_q, *d_o, *d_k_pool, *d_v_pool;
|
||||
int64_t* d_pt;
|
||||
float *d_op, *d_ml;
|
||||
cudaMalloc(&d_q, sz_q); cudaMalloc(&d_o, sz_q);
|
||||
cudaMalloc(&d_k_pool, sz_kv); cudaMalloc(&d_v_pool, sz_kv);
|
||||
cudaMalloc(&d_pt, sz_pt);
|
||||
cudaMalloc(&d_op, sz_op); cudaMalloc(&d_ml, sz_ml);
|
||||
|
||||
bf16* tmp = (bf16*)malloc(sz_kv > sz_q ? sz_kv : sz_q);
|
||||
for (size_t i = 0; i < sz_q / sizeof(bf16); i++) tmp[i] = f2bf(randf());
|
||||
cudaMemcpy(d_q, tmp, sz_q, cudaMemcpyHostToDevice);
|
||||
for (size_t i = 0; i < sz_kv / sizeof(bf16); i++) tmp[i] = f2bf(randf());
|
||||
cudaMemcpy(d_k_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
|
||||
cudaMemcpy(d_v_pool, tmp, sz_kv, cudaMemcpyHostToDevice);
|
||||
|
||||
int64_t* h_pt = (int64_t*)malloc(sz_pt);
|
||||
int next_pg = 0;
|
||||
for (int b = 0; b < B; b++)
|
||||
for (int p = 0; p < max_pages; p++)
|
||||
h_pt[b * max_pages + p] = next_pg++;
|
||||
cudaMemcpy(d_pt, h_pt, sz_pt, cudaMemcpyHostToDevice);
|
||||
free(h_pt);
|
||||
|
||||
float scale_val = 1.0f / sqrtf((float)HEAD_DIM);
|
||||
PagedAttentionParams<bf16, float> pa;
|
||||
pa.batch = B; pa.q_head = Hq; pa.kv_head = Hkv; pa.q_len = 1;
|
||||
pa.kv_len = kv_len; pa.head_dim = HEAD_DIM;
|
||||
pa.use_mask = 0; pa.causal_offset = -1;
|
||||
set_default_paged_strides(pa);
|
||||
pa.num_splits = 1; pa.scale = scale_val;
|
||||
pa.page_size = page_size; pa.max_pages = max_pages;
|
||||
pa.page_table = d_pt;
|
||||
pa.k_cache = d_k_pool; pa.v_cache = d_v_pool;
|
||||
pa.q = d_q; pa.mask = nullptr; pa.o = d_o;
|
||||
pa.o_part = d_op; pa.ml_part = d_ml;
|
||||
|
||||
const int WARMUP = 10, ITERS = 100;
|
||||
auto launch = [&]() { launch_paged_decode<HEAD_DIM>(pa); };
|
||||
double flops = 4.0 * B * Hq * (double)kv_len * HEAD_DIM;
|
||||
size_t nKV = (size_t)B * Hkv * kv_len * HEAD_DIM;
|
||||
double bytes = 2.0 * (2.0 * nKV * sizeof(bf16));
|
||||
BenchResult r = bench_kernel(launch, WARMUP, ITERS, flops, bytes);
|
||||
|
||||
char cfg[64];
|
||||
snprintf(cfg, sizeof(cfg),
|
||||
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d page=%3d",
|
||||
B, Hq, Hkv, 1, kv_len, HEAD_DIM, page_size);
|
||||
print_bench_row(cfg, r);
|
||||
|
||||
free(tmp);
|
||||
cudaFree(d_q); cudaFree(d_o);
|
||||
cudaFree(d_k_pool); cudaFree(d_v_pool); cudaFree(d_pt);
|
||||
cudaFree(d_op); cudaFree(d_ml);
|
||||
}
|
||||
|
||||
static void bench() {
|
||||
printf("\n===== PAGED DECODE BENCH =====\n");
|
||||
print_bench_header();
|
||||
bench_config<128>(1, 32, 4, 512, 128);
|
||||
bench_config<128>(1, 32, 4, 1024, 128);
|
||||
bench_config<128>(1, 32, 4, 2048, 128);
|
||||
bench_config<128>(1, 32, 4, 4096, 128);
|
||||
bench_config<128>(16, 32, 4, 2048, 128);
|
||||
bench_config<128>(32, 32, 4, 1024, 128);
|
||||
}
|
||||
|
||||
int main() {
|
||||
int n = sizeof(TESTS) / sizeof(TESTS[0]);
|
||||
int fail = 0;
|
||||
printf("=== Paged Decode vs CPU reference (%d cases) ===\n\n", n);
|
||||
|
||||
for (int i = 0; i < n; i++) {
|
||||
fail += dispatch_test(TESTS[i]);
|
||||
if (fail) break;
|
||||
}
|
||||
|
||||
if (fail) {
|
||||
printf("\nFAILED (%d/%d tests failed)\n", fail, n);
|
||||
return fail;
|
||||
}
|
||||
printf("\nAll %d tests passed!\n", n);
|
||||
bench();
|
||||
return 0;
|
||||
}
|
||||
@@ -0,0 +1,178 @@
|
||||
/*
|
||||
Pure-C test:
|
||||
nvcc -I csrc -arch=sm_89 -O3 \
|
||||
--use_fast_math --ptxas-options=-O3 --extra-device-vectorization \
|
||||
csrc/tests/attn_prefill_test.cu -o test && ./test
|
||||
*/
|
||||
|
||||
#include "test_utils.cuh"
|
||||
#include "../kernels/attn_prefill_split_q.cuh"
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
#include "../kernels/attn_prefill_split_q_mma.cuh"
|
||||
#endif
|
||||
|
||||
// Launch the production prefill path (tensor-core MMA on sm_80+, else the
|
||||
// scalar fallback), mirroring dispatch_prefill() in attn_prefill.cu.
|
||||
template <int HEAD_DIM>
|
||||
static void launch_prefill(AttentionParams<bf16>& p) {
|
||||
#ifndef ASTRAI_NO_MMA
|
||||
constexpr int WARPS = 4, BR = 16;
|
||||
constexpr int BC = (HEAD_DIM <= 128) ? 32 : 16;
|
||||
dim3 grid((p.q_len + BR * WARPS - 1) / (BR * WARPS), p.q_head, p.batch);
|
||||
dim3 block(WARPS * 32, 1, 1);
|
||||
attn_prefill_split_q_mma_kernel<HEAD_DIM, WARPS, BC><<<grid, block>>>(p);
|
||||
#else
|
||||
constexpr int G = 8, ROWS = 32, P_BC = 32;
|
||||
dim3 grid((p.q_len + ROWS - 1) / ROWS, p.q_head, p.batch);
|
||||
dim3 block(G, ROWS, 1);
|
||||
attn_prefill_split_q_kernel_t<HEAD_DIM, G, ROWS, P_BC><<<grid, block>>>(p);
|
||||
#endif
|
||||
}
|
||||
|
||||
static void dispatch_prefill(AttentionParams<bf16>& p) {
|
||||
switch (p.head_dim) {
|
||||
case 64: launch_prefill<64>(p); break;
|
||||
case 128: launch_prefill<128>(p); break;
|
||||
default: printf("bench: unsupported D=%d\n", p.head_dim);
|
||||
}
|
||||
}
|
||||
|
||||
// Warmed-up, CUDA-event timed throughput sweep over the production MMA path.
|
||||
// Reports per-call latency and effective tensor-core TFLOP/s (2 matmuls:
|
||||
// QK^T and P@V, each 2*B*Hq*ql*kl*D flops; halved for causal).
|
||||
static void bench() {
|
||||
const int cfgs[][7] = {
|
||||
{1,32,4,512,512,128,0},
|
||||
{1,32,4,1024,1024,128,0},
|
||||
{1,32,4,2048,2048,128,0},
|
||||
{1,32,4,2048,2048,128,1},
|
||||
{4,32,4,2048,2048,128,1},
|
||||
{1,32,4,4096,4096,128,1},
|
||||
};
|
||||
int n = sizeof(cfgs)/sizeof(cfgs[0]);
|
||||
const int WARMUP = 10, ITERS = 50;
|
||||
printf("\n===== PREFILL BENCH (warmup=%d iters=%d) =====\n", WARMUP, ITERS);
|
||||
printf("%-46s | %10s | %10s | %10s\n",
|
||||
"config", "latency", "bandwidth", "throughput");
|
||||
printf("---------------------------------------------------------------"
|
||||
"----------------------------\n");
|
||||
|
||||
for (int ci = 0; ci < n; ci++) {
|
||||
int B=cfgs[ci][0], Hq=cfgs[ci][1], Hk=cfgs[ci][2];
|
||||
int ql=cfgs[ci][3], kl=cfgs[ci][4], D=cfgs[ci][5], causal=cfgs[ci][6];
|
||||
size_t nQ=(size_t)B*Hq*ql*D, nKV=(size_t)B*Hk*kl*D;
|
||||
|
||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||
size_t big = nQ>nKV?nQ:nKV; tmp=new bf16[big];
|
||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(randf());
|
||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
|
||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(randf());
|
||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
||||
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||
set_default_strides(p);
|
||||
p.scale=1.0f/sqrtf((float)D);
|
||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
for (int i=0;i<WARMUP;i++) dispatch_prefill(p);
|
||||
cudaDeviceSynchronize();
|
||||
cudaError_t err=cudaGetLastError();
|
||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return;}
|
||||
|
||||
cudaEvent_t s,e; cudaEventCreate(&s); cudaEventCreate(&e);
|
||||
cudaEventRecord(s);
|
||||
for (int i=0;i<ITERS;i++) dispatch_prefill(p);
|
||||
cudaEventRecord(e); cudaEventSynchronize(e);
|
||||
float ms=0; cudaEventElapsedTime(&ms,s,e); ms/=ITERS;
|
||||
|
||||
double flops = 4.0*B*Hq*(double)ql*kl*D;
|
||||
if (causal) flops *= 0.5;
|
||||
double tflops = flops/(ms*1e-3)/1e12;
|
||||
// HBM traffic: Q + O (B*Hq*ql*D each) + K + V (B*Hk*kl*D each), bf16.
|
||||
double bytes = 2.0 * (2.0*nQ + 2.0*nKV);
|
||||
double gbps = bytes/(ms*1e-3)/1e9;
|
||||
|
||||
char cfg[64];
|
||||
snprintf(cfg, sizeof(cfg),
|
||||
"B=%2d Hq=%2d Hk=%d q=%4d kv=%4d D=%3d causal=%d",
|
||||
B,Hq,Hk,ql,kl,D,causal);
|
||||
printf("%-46s | %7.4f ms | %7.1f GB/s | %6.2f TFLOP/s\n",
|
||||
cfg, ms, gbps, tflops);
|
||||
|
||||
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
||||
delete[]tmp; cudaEventDestroy(s); cudaEventDestroy(e);
|
||||
}
|
||||
}
|
||||
|
||||
int main() {
|
||||
const int configs[][7] = {
|
||||
{1,2,1,64,128,64,0}, // tiny: B,Hq,Hk,q,kv,D,causal
|
||||
{1,32,4,512,512,128,0}, // standard
|
||||
{1,32,4,128,256,128,0}, // medium
|
||||
{1,4,2,256,256,128,1}, // causal
|
||||
};
|
||||
int n_configs = sizeof(configs) / sizeof(configs[0]);
|
||||
|
||||
for (int ci = 0; ci < n_configs; ci++) {
|
||||
int B=configs[ci][0], Hq=configs[ci][1], Hk=configs[ci][2];
|
||||
int ql=configs[ci][3], kl=configs[ci][4], D=configs[ci][5];
|
||||
int causal=configs[ci][6];
|
||||
printf("=== B=%d Hq=%d Hk=%d q=%d kv=%d D=%d causal=%d ===\n",
|
||||
B,Hq,Hk,ql,kl,D,causal);
|
||||
|
||||
size_t nQ = B*Hq*ql*D, nKV = B*Hk*kl*D;
|
||||
float *hQ=new float[nQ], *hK=new float[nKV], *hV=new float[nKV];
|
||||
for (size_t i=0;i<nQ;i++) hQ[i]=randf();
|
||||
for (size_t i=0;i<nKV;i++){hK[i]=randf();hV[i]=randf();}
|
||||
|
||||
bf16 *dQ,*dK,*dV,*dO,*tmp;
|
||||
cudaMalloc(&dQ,nQ*2); cudaMalloc(&dK,nKV*2);
|
||||
cudaMalloc(&dV,nKV*2); cudaMalloc(&dO,nQ*2);
|
||||
tmp=new bf16[max(nQ,nKV)];
|
||||
for (size_t i=0;i<nQ;i++) tmp[i]=f2bf(hQ[i]);
|
||||
cudaMemcpy(dQ,tmp,nQ*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hK[i]);
|
||||
cudaMemcpy(dK,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
for (size_t i=0;i<nKV;i++) tmp[i]=f2bf(hV[i]);
|
||||
cudaMemcpy(dV,tmp,nKV*2,cudaMemcpyHostToDevice);
|
||||
|
||||
AttentionParams<bf16> p;
|
||||
p.batch=B; p.q_head=Hq; p.kv_head=Hk; p.q_len=ql; p.kv_len=kl; p.head_dim=D;
|
||||
p.use_mask=0; p.causal_offset=causal?0:-1;
|
||||
set_default_strides(p);
|
||||
p.scale=1.0f/sqrtf((float)D);
|
||||
p.q=dQ; p.k=dK; p.v=dV; p.mask=nullptr; p.o=dO;
|
||||
|
||||
double t0=now_ms();
|
||||
dispatch_prefill(p);
|
||||
cudaDeviceSynchronize();
|
||||
double kms=now_ms()-t0;
|
||||
cudaError_t err=cudaGetLastError();
|
||||
if (err!=cudaSuccess){printf("CUDA err: %s\n",cudaGetErrorString(err));return 1;}
|
||||
|
||||
bf16* hOut=new bf16[nQ];
|
||||
cudaMemcpy(hOut,dO,nQ*2,cudaMemcpyDeviceToHost);
|
||||
|
||||
float* ref=new float[nQ];
|
||||
cpu_attention_ref(hQ, hK, hV, nullptr, ref, B, Hq, Hk, ql, kl, D, causal ? 0 : -1);
|
||||
|
||||
float max_err=0;
|
||||
for (size_t i=0;i<nQ;i++) {
|
||||
float d=fabsf(bf2f(hOut[i])-ref[i]);
|
||||
if(d>max_err) max_err=d;
|
||||
}
|
||||
printf("kernel: %.3f ms max_err: %.6e\n\n",kms,max_err);
|
||||
|
||||
cudaFree(dQ);cudaFree(dK);cudaFree(dV);cudaFree(dO);
|
||||
delete[]hQ;delete[]hK;delete[]hV;delete[]hOut;delete[]ref;delete[]tmp;
|
||||
}
|
||||
printf("All tests passed!\n");
|
||||
bench();
|
||||
return 0;
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
#pragma once
|
||||
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cmath>
|
||||
#include <chrono>
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
using bf16 = __nv_bfloat16;
|
||||
|
||||
inline bf16 f2bf(float x) { return __float2bfloat16(x); }
|
||||
inline float bf2f(bf16 x) { return __bfloat162float(x); }
|
||||
|
||||
inline float randf() { return (float)rand() / (float)RAND_MAX - 0.5f; }
|
||||
|
||||
inline double now_ms() {
|
||||
using namespace std::chrono;
|
||||
return duration_cast<milliseconds>(steady_clock::now().time_since_epoch()).count();
|
||||
}
|
||||
|
||||
inline int compute_num_splits(int base_blocks, int tiles_total) {
|
||||
int sm_count = 0;
|
||||
cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, 0);
|
||||
int n = (2 * sm_count + base_blocks - 1) / base_blocks;
|
||||
if (n > tiles_total) n = tiles_total;
|
||||
if (n > 32) n = 32;
|
||||
if (n < 1) n = 1;
|
||||
return n;
|
||||
}
|
||||
|
||||
#define CUDA_CHECK(call) \
|
||||
do { \
|
||||
cudaError_t _e = (call); \
|
||||
if (_e != cudaSuccess) { \
|
||||
printf("CUDA error %s at %s:%d\n", cudaGetErrorString(_e), __FILE__, __LINE__); \
|
||||
exit(1); \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
struct BenchResult {
|
||||
float ms;
|
||||
double gbps;
|
||||
double tflops;
|
||||
};
|
||||
|
||||
template <typename Fn>
|
||||
BenchResult bench_kernel(Fn launch, int warmup, int iters,
|
||||
double flops, double bytes) {
|
||||
for (int i = 0; i < warmup; i++) launch();
|
||||
cudaDeviceSynchronize();
|
||||
cudaError_t err = cudaGetLastError();
|
||||
if (err != cudaSuccess) {
|
||||
printf("CUDA error before bench: %s\n", cudaGetErrorString(err));
|
||||
return {0, 0, 0};
|
||||
}
|
||||
|
||||
cudaEvent_t s, e;
|
||||
cudaEventCreate(&s); cudaEventCreate(&e);
|
||||
cudaEventRecord(s);
|
||||
for (int i = 0; i < iters; i++) launch();
|
||||
cudaEventRecord(e); cudaEventSynchronize(e);
|
||||
float ms = 0; cudaEventElapsedTime(&ms, s, e); ms /= iters;
|
||||
cudaEventDestroy(s); cudaEventDestroy(e);
|
||||
|
||||
return {ms, bytes / (ms * 1e-3) / 1e9, flops / (ms * 1e-3) / 1e12};
|
||||
}
|
||||
|
||||
inline void print_bench_header() {
|
||||
printf("%-46s | %10s | %10s | %10s\n",
|
||||
"config", "latency", "bandwidth", "throughput");
|
||||
printf("---------------------------------------------------------------"
|
||||
"----------------------------\n");
|
||||
}
|
||||
|
||||
inline void print_bench_row(const char* cfg, const BenchResult& r) {
|
||||
printf("%-46s | %7.4f ms | %7.1f GB/s | %6.2f TFLOP/s\n",
|
||||
cfg, r.ms, r.gbps, r.tflops);
|
||||
}
|
||||
|
||||
template <int... Ds>
|
||||
struct _HeadSwitch;
|
||||
|
||||
template <int D>
|
||||
struct _HeadSwitch<D> {
|
||||
template <typename Fn>
|
||||
static void call(int hd, Fn&& fn) { if (hd == D) fn.template operator()<D>(); }
|
||||
};
|
||||
|
||||
template <int D, int... Rest>
|
||||
struct _HeadSwitch<D, Rest...> {
|
||||
template <typename Fn>
|
||||
static void call(int hd, Fn&& fn) {
|
||||
if (hd == D) fn.template operator()<D>();
|
||||
else _HeadSwitch<Rest...>::call(hd, fn);
|
||||
}
|
||||
};
|
||||
|
||||
// Default set: 32, 64, 128, 256
|
||||
template <typename Fn>
|
||||
void dispatch_by_head_dim(int head_dim, Fn&& fn) {
|
||||
_HeadSwitch<32, 64, 128, 256>::call(head_dim, fn);
|
||||
}
|
||||
|
||||
// Set default strides for contiguous b h l d layout on AttentionParams.
|
||||
template<typename P>
|
||||
inline void set_default_strides(P& p) {
|
||||
p.q_stride_b = p.q_head * p.q_len * p.head_dim;
|
||||
p.q_stride_h = p.q_len * p.head_dim;
|
||||
p.q_stride_l = p.head_dim;
|
||||
p.q_stride_d = 1;
|
||||
p.kv_stride_b = p.kv_head * p.kv_len * p.head_dim;
|
||||
p.kv_stride_h = p.kv_len * p.head_dim;
|
||||
p.kv_stride_l = p.head_dim;
|
||||
p.kv_stride_d = 1;
|
||||
p.mask_b_stride = p.kv_len;
|
||||
p.mask_q_stride = 0;
|
||||
}
|
||||
|
||||
// Set default Q strides for contiguous b h l d layout on PagedAttentionParams.
|
||||
template<typename P>
|
||||
inline void set_default_paged_strides(P& p) {
|
||||
p.q_stride_b = p.q_head * p.q_len * p.head_dim;
|
||||
p.q_stride_h = p.q_len * p.head_dim;
|
||||
p.q_stride_l = p.head_dim;
|
||||
p.q_stride_d = 1;
|
||||
p.mask_b_stride = p.kv_len;
|
||||
p.mask_q_stride = 0;
|
||||
}
|
||||
|
||||
// Generic CPU reference for multi-query / grouped-query attention.
|
||||
// Tensor shapes (all float*):
|
||||
// Q : [B, Hq, q_len, D]
|
||||
// K : [B, Hk, kv_len, D]
|
||||
// V : [B, Hk, kv_len, D]
|
||||
// O : [B, Hq, q_len, D]
|
||||
// mask: if q_len == 1, shape is [B, kv_len]; otherwise mask is not supported.
|
||||
// causal_offset: -1 = non-causal; >=0 = absolute position of first Q token.
|
||||
static void cpu_attention_ref(
|
||||
const float* Q, const float* K, const float* V, const bool* mask,
|
||||
float* O, int B, int Hq, int Hk, int q_len, int kv_len, int D,
|
||||
int causal_offset
|
||||
) {
|
||||
float scale = 1.0f / sqrtf((float)D);
|
||||
int n_rep = Hq / Hk;
|
||||
for (int b = 0; b < B; b++) {
|
||||
for (int h = 0; h < Hq; h++) {
|
||||
int kv_h = h / n_rep;
|
||||
for (int qi = 0; qi < q_len; qi++) {
|
||||
float mv = -INFINITY, sv = 0.0f;
|
||||
float accum[256] = {0.0f};
|
||||
int lim = kv_len;
|
||||
if (causal_offset >= 0) {
|
||||
int c = qi + causal_offset + 1;
|
||||
lim = (c < kv_len) ? c : kv_len;
|
||||
}
|
||||
for (int kj = 0; kj < lim; kj++) {
|
||||
if (mask != nullptr && q_len == 1) {
|
||||
if (!mask[b * kv_len + kj]) continue;
|
||||
}
|
||||
float dot = 0.0f;
|
||||
size_t q_idx = ((size_t)b * Hq + h) * q_len + qi;
|
||||
size_t kv_idx = ((size_t)b * Hk + kv_h) * kv_len + kj;
|
||||
for (int d = 0; d < D; d++)
|
||||
dot += Q[q_idx * D + d] * K[kv_idx * D + d];
|
||||
dot *= scale;
|
||||
float nm = fmaxf(mv, dot);
|
||||
float a = expf(mv - nm);
|
||||
float b_exp = expf(dot - nm);
|
||||
sv = sv * a + b_exp;
|
||||
for (int d = 0; d < D; d++)
|
||||
accum[d] = accum[d] * a + V[kv_idx * D + d] * b_exp;
|
||||
mv = nm;
|
||||
}
|
||||
float inv = 1.0f / sv;
|
||||
size_t o_idx = ((size_t)b * Hq + h) * q_len + qi;
|
||||
for (int d = 0; d < D; d++)
|
||||
O[o_idx * D + d] = accum[d] * inv;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
+3
-3
@@ -9,8 +9,8 @@ readme = "README.md"
|
||||
requires-python = ">=3.12"
|
||||
dependencies = [
|
||||
"h5py==3.15.1",
|
||||
"numpy==2.3.2",
|
||||
"torch==2.7.1",
|
||||
"numpy==2.4.4",
|
||||
"torch==2.11.0",
|
||||
"tokenizers==0.21.4",
|
||||
"tqdm==4.67.1",
|
||||
"safetensors==0.5.3",
|
||||
@@ -37,7 +37,7 @@ dev = ["pytest==9.0.2", "ruff"]
|
||||
where = ["."]
|
||||
|
||||
[tool.pip]
|
||||
extra-index-url = "https://download.pytorch.org/whl/cu126"
|
||||
extra-index-url = "https://download.pytorch.org/whl/cu128"
|
||||
|
||||
[tool.setuptools.dynamic]
|
||||
version = { attr = "astrai.__version__" }
|
||||
|
||||
@@ -5,7 +5,7 @@ from huggingface_hub import snapshot_download
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[2]
|
||||
DEFAULT_LOCAL_DIR = Path(PROJECT_ROOT, "params")
|
||||
DEFAULT_REPO_ID = "ViperEk/KHAOSZ"
|
||||
DEFAULT_REPO_ID = "ViperEkura/AstrAI-V1-instruct"
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(
|
||||
|
||||
@@ -26,11 +26,9 @@ def batch_generate():
|
||||
|
||||
prompts = [
|
||||
tokenizer.apply_chat_template(
|
||||
[
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": q},
|
||||
],
|
||||
[{"role": "user", "content": q}],
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
)
|
||||
for q in inputs
|
||||
]
|
||||
|
||||
+73
-16
@@ -1,3 +1,4 @@
|
||||
from argparse import ArgumentParser
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
@@ -7,15 +8,69 @@ from astrai.model import AutoModel
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[2]
|
||||
PARAMETER_ROOT = Path(PROJECT_ROOT, "params")
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = ArgumentParser(description="Interactive streaming chat")
|
||||
parser.add_argument(
|
||||
"--model_path",
|
||||
type=Path,
|
||||
default=PROJECT_ROOT / "params",
|
||||
help="Path to model weights (params/ or checkpoint/epoch_N_step_M/)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--temperature",
|
||||
type=float,
|
||||
default=0.8,
|
||||
help="Sampling temperature (default: 0.8)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--top_p",
|
||||
type=float,
|
||||
default=0.95,
|
||||
help="Top-p sampling threshold",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--top_k",
|
||||
type=int,
|
||||
default=50,
|
||||
help="Top-k sampling threshold",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_tokens",
|
||||
type=int,
|
||||
default=2048,
|
||||
help="Maximum tokens to generate",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--frequency_penalty",
|
||||
type=float,
|
||||
default=0.5,
|
||||
help="Penalty per occurrence for repeated tokens (0.0 disables, "
|
||||
"range -2.0~2.0, typical 0.3-1.0)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--rep_window",
|
||||
type=int,
|
||||
default=64,
|
||||
help="Number of recent prompt tokens to include in penalty history",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--system_prompt",
|
||||
type=str,
|
||||
default="",
|
||||
help="Optional system prompt (default: empty, model not SFT-trained on system role)",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def chat():
|
||||
model = AutoModel.from_pretrained(PARAMETER_ROOT)
|
||||
tokenizer = AutoTokenizer.from_pretrained(PARAMETER_ROOT)
|
||||
model.to(device="cuda", dtype=torch.bfloat16)
|
||||
args = parse_args()
|
||||
model_path = args.model_path
|
||||
|
||||
messages = [{"role": "system", "content": "You are a helpful assistant."}]
|
||||
model = AutoModel.from_pretrained(model_path)
|
||||
tokenizer = AutoTokenizer.from_pretrained(model_path)
|
||||
model.to(device="cuda", dtype=torch.bfloat16)
|
||||
engine = InferenceEngine(model=model, tokenizer=tokenizer)
|
||||
|
||||
while True:
|
||||
@@ -23,27 +78,29 @@ def chat():
|
||||
if query == "!exit":
|
||||
break
|
||||
|
||||
# Add user message
|
||||
messages.append({"role": "user", "content": query})
|
||||
msgs = []
|
||||
if args.system_prompt:
|
||||
msgs.append({"role": "system", "content": args.system_prompt})
|
||||
msgs.append({"role": "user", "content": query})
|
||||
prompt = tokenizer.apply_chat_template(
|
||||
msgs, tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
|
||||
# Generate response
|
||||
full_response = ""
|
||||
prompt = tokenizer.apply_chat_template(messages, tokenize=False)
|
||||
|
||||
for token in engine.generate(
|
||||
prompt=prompt,
|
||||
stream=True,
|
||||
max_tokens=2048,
|
||||
temperature=0.8,
|
||||
top_p=0.95,
|
||||
top_k=50,
|
||||
max_tokens=args.max_tokens,
|
||||
temperature=args.temperature,
|
||||
top_p=args.top_p,
|
||||
top_k=args.top_k,
|
||||
frequency_penalty=args.frequency_penalty,
|
||||
rep_window=args.rep_window,
|
||||
):
|
||||
print(token, end="", flush=True)
|
||||
full_response += token
|
||||
|
||||
print()
|
||||
# Add assistant response to messages
|
||||
messages.append({"role": "assistant", "content": full_response.strip()})
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -0,0 +1,321 @@
|
||||
"""SVD effective rank & weight statistics analysis for model checkpoints."""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import safetensors.torch
|
||||
import torch
|
||||
|
||||
|
||||
def effective_rank_metrics(w: torch.Tensor) -> dict:
|
||||
if w.ndim == 1:
|
||||
return {"shape": tuple(w.shape), "is_1d": True}
|
||||
|
||||
w = w.float()
|
||||
s = torch.linalg.svdvals(w)
|
||||
s_sq = s**2
|
||||
total = s_sq.sum()
|
||||
cumsum = torch.cumsum(s_sq, dim=0) / total
|
||||
|
||||
min_dim = min(w.shape[0], w.shape[1])
|
||||
er_90 = (cumsum < 0.90).sum().item() + 1
|
||||
er_95 = (cumsum < 0.95).sum().item() + 1
|
||||
er_99 = (cumsum < 0.99).sum().item() + 1
|
||||
|
||||
p = s_sq / total
|
||||
p = p[p > 1e-30]
|
||||
entropy = -(p * torch.log(p)).sum()
|
||||
entropic_rank = torch.exp(entropy).item()
|
||||
|
||||
return {
|
||||
"shape": tuple(w.shape),
|
||||
"min_dim": min_dim,
|
||||
"er_90": er_90,
|
||||
"er_95": er_95,
|
||||
"er_99": er_99,
|
||||
"er_99_norm": er_99 / min_dim,
|
||||
"er_95_norm": er_95 / min_dim,
|
||||
"entropic_rank": entropic_rank,
|
||||
"entropic_rank_norm": entropic_rank / min_dim,
|
||||
"top1_ratio": s[0].item() / s.sum().item(),
|
||||
"top5_ratio": s[:5].sum().item() / s.sum().item(),
|
||||
"decay_ratio": s[-1].item() / s[0].item(),
|
||||
"condition_number": s[0].item() / s[-1].item(),
|
||||
"mean": w.mean().item(),
|
||||
"std": w.std().item(),
|
||||
"min": w.min().item(),
|
||||
"max": w.max().item(),
|
||||
}
|
||||
|
||||
|
||||
def format_header(headers: list[str], widths: list[int]) -> str:
|
||||
return "".join(h.ljust(w) for h, w in zip(headers, widths))
|
||||
|
||||
|
||||
def format_row(values: list[str], widths: list[int]) -> str:
|
||||
return "".join(v.ljust(w) for v, w in zip(values, widths))
|
||||
|
||||
|
||||
def group_by_component(results: dict[str, dict]) -> dict[str, list[dict]]:
|
||||
groups: dict[str, list[dict]] = {}
|
||||
for key, r in results.items():
|
||||
parts = key.split(".")
|
||||
if parts[0] == "layers" and len(parts) >= 3:
|
||||
sub = parts[2:]
|
||||
if sub[0] == "attention":
|
||||
comp = f"attn.{sub[1]}"
|
||||
elif sub[0] == "mlp":
|
||||
comp = f"mlp.{sub[1]}"
|
||||
elif sub[0] == "input_norm":
|
||||
comp = "input_norm"
|
||||
elif sub[0] == "post_attention_norm":
|
||||
comp = "post_attn_norm"
|
||||
else:
|
||||
comp = ".".join(sub)
|
||||
else:
|
||||
comp = key
|
||||
groups.setdefault(comp, []).append(r)
|
||||
return groups
|
||||
|
||||
|
||||
def print_component_summary(results: dict[str, dict], title: str):
|
||||
groups = group_by_component(results)
|
||||
matrix_groups = {
|
||||
k: [v for v in vs if not v.get("is_1d")]
|
||||
for k, vs in groups.items()
|
||||
if any(not v.get("is_1d") for v in vs)
|
||||
}
|
||||
|
||||
widths = [20, 12, 12, 12, 12, 12]
|
||||
print(f"\n{title}")
|
||||
print(
|
||||
format_header(
|
||||
["Component", "N", "ER@99%", "EntRank%", "Top1 σ(%)", "Cond. Num"], widths
|
||||
)
|
||||
)
|
||||
print("-" * sum(widths))
|
||||
|
||||
for name in sorted(matrix_groups.keys()):
|
||||
items = matrix_groups[name]
|
||||
n = len(items)
|
||||
print(
|
||||
format_row(
|
||||
[
|
||||
name,
|
||||
str(n),
|
||||
f"{sum(r['er_99_norm'] for r in items) / n:.4f}",
|
||||
f"{sum(r['entropic_rank_norm'] for r in items) / n:.4f}",
|
||||
f"{sum(r['top1_ratio'] for r in items) / n:.4f}",
|
||||
f"{sum(r['condition_number'] for r in items) / n:.1f}",
|
||||
],
|
||||
widths,
|
||||
)
|
||||
)
|
||||
|
||||
all_er = [
|
||||
r["er_99_norm"]
|
||||
for vs in matrix_groups.values()
|
||||
for r in vs
|
||||
if not r.get("is_1d")
|
||||
]
|
||||
if all_er:
|
||||
m = sum(all_er) / len(all_er)
|
||||
print(f"\n Overall Mean ER@99: {m:.4f} ({m * 100:.1f}% of dimension)")
|
||||
if m > 0.85:
|
||||
print(" → HIGH utilization: model near capacity → need more params")
|
||||
elif m > 0.5:
|
||||
print(" → MODERATE utilization: some headroom left")
|
||||
else:
|
||||
print(" → LOW utilization: significant unused capacity")
|
||||
|
||||
|
||||
def print_layer_grid(results: dict[str, dict]):
|
||||
comps = [
|
||||
"attn.q_proj",
|
||||
"attn.k_proj",
|
||||
"attn.v_proj",
|
||||
"attn.o_proj",
|
||||
"mlp.up",
|
||||
"mlp.gate",
|
||||
"mlp.down",
|
||||
]
|
||||
widths = [6] + [10] * len(comps)
|
||||
metric = "er_99_norm"
|
||||
|
||||
print(f"\n--- Per-Layer Effective Rank (99% energy) ---")
|
||||
print(format_header(["Layer"] + comps, widths))
|
||||
print("-" * sum(widths))
|
||||
|
||||
layer_data: dict[int, dict[str, dict]] = {}
|
||||
for key, r in results.items():
|
||||
parts = key.split(".")
|
||||
if parts[0] != "layers":
|
||||
continue
|
||||
li = int(parts[1])
|
||||
sub = parts[2:]
|
||||
if sub[0] == "attention":
|
||||
cname = f"attn.{sub[1]}"
|
||||
elif sub[0] == "mlp":
|
||||
cname = f"mlp.{sub[1]}"
|
||||
else:
|
||||
continue
|
||||
layer_data.setdefault(li, {})[cname] = r
|
||||
|
||||
for li in sorted(layer_data):
|
||||
values = [str(li)]
|
||||
for c in comps:
|
||||
v = layer_data[li].get(c, {}).get(metric, 0)
|
||||
values.append(f"{v:.4f}")
|
||||
print(format_row(values, widths))
|
||||
|
||||
|
||||
def print_weight_stats(results: dict[str, dict]):
|
||||
groups = group_by_component(results)
|
||||
widths = [20, 12, 12, 12, 12]
|
||||
print(f"\n--- Weight Value Statistics ---")
|
||||
print(format_header(["Component", "Mean", "Std", "Min", "Max"], widths))
|
||||
print("-" * sum(widths))
|
||||
|
||||
for name in sorted(groups.keys()):
|
||||
items = groups[name]
|
||||
means = [r.get("mean", 0) for r in items]
|
||||
stds = [r.get("std", 0) for r in items]
|
||||
mins = [r.get("min", 0) for r in items]
|
||||
maxs = [r.get("max", 0) for r in items]
|
||||
g_mean = sum(means) / len(means)
|
||||
g_std = sum(stds) / len(stds)
|
||||
g_min = min(mins)
|
||||
g_max = max(maxs)
|
||||
print(
|
||||
format_row(
|
||||
[
|
||||
name,
|
||||
f"{g_mean:.6f}",
|
||||
f"{g_std:.6f}",
|
||||
f"{g_min:.6f}",
|
||||
f"{g_max:.6f}",
|
||||
],
|
||||
widths,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def print_params_summary(results: dict[str, dict]):
|
||||
total_2d = sum(
|
||||
r["shape"][0] * r["shape"][1] for r in results.values() if not r.get("is_1d")
|
||||
)
|
||||
total_1d = sum(r["shape"][0] for r in results.values() if r.get("is_1d"))
|
||||
print(f"\n Total 2D params: {total_2d:,}")
|
||||
print(f" Total 1D params: {total_1d:,}")
|
||||
print(f" Total params: {total_2d + total_1d:,}")
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="SVD effective rank & weight statistics of a model checkpoint."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ckpt_dir",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to checkpoint directory (containing model.safetensors + config.json).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--compare",
|
||||
type=str,
|
||||
nargs="*",
|
||||
help="Additional checkpoint directories to compare against.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--no_svd",
|
||||
action="store_true",
|
||||
help="Skip SVD analysis, only show weight statistics (mean/std/min/max).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Save results as JSON to this path.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
all_results = {}
|
||||
|
||||
def analyze_one(ckpt_dir: str, label: str):
|
||||
ckpt_dir = Path(ckpt_dir)
|
||||
weights_path = ckpt_dir / "model.safetensors"
|
||||
if not weights_path.exists():
|
||||
print(f"ERROR: {weights_path} not found")
|
||||
return {}
|
||||
|
||||
meta = {}
|
||||
meta_path = ckpt_dir / "meta.json"
|
||||
if meta_path.exists():
|
||||
with open(meta_path) as f:
|
||||
meta = json.load(f)
|
||||
|
||||
print(f"\n{'=' * 70}")
|
||||
print(f" {label}: {ckpt_dir}")
|
||||
if meta:
|
||||
print(
|
||||
f" Iteration: {meta.get('iteration', '?')}, "
|
||||
f"Strategy: {meta.get('strategy', '?')}, "
|
||||
f"nprocs={meta.get('nprocs', '?')}"
|
||||
)
|
||||
print(f"{'=' * 70}")
|
||||
|
||||
print(f"Loading weights...")
|
||||
sd = safetensors.torch.load_file(str(weights_path))
|
||||
print(f" {len(sd)} keys loaded")
|
||||
|
||||
weight_keys = [
|
||||
k
|
||||
for k in sd
|
||||
if ".weight" in k and "rotary_embedding" not in k and "freqs_cis" not in k
|
||||
]
|
||||
|
||||
results = {}
|
||||
if not args.no_svd:
|
||||
print(f"Computing SVD on {len(weight_keys)} tensors...")
|
||||
for i, k in enumerate(sorted(weight_keys)):
|
||||
print(f" [{i + 1}/{len(weight_keys)}] {k:<60s}", end="\r")
|
||||
results[k] = effective_rank_metrics(sd[k])
|
||||
print()
|
||||
else:
|
||||
print(f"Computing stats on {len(weight_keys)} tensors (no SVD)...")
|
||||
for i, k in enumerate(sorted(weight_keys)):
|
||||
t = sd[k]
|
||||
results[k] = {
|
||||
"shape": tuple(t.shape),
|
||||
"is_1d": t.ndim == 1,
|
||||
"mean": t.float().mean().item(),
|
||||
"std": t.float().std().item(),
|
||||
"min": t.float().min().item(),
|
||||
"max": t.float().max().item(),
|
||||
}
|
||||
|
||||
print_params_summary(results)
|
||||
if not args.no_svd:
|
||||
print_component_summary(
|
||||
results, "\n=== SVD Effective Rank by Component ==="
|
||||
)
|
||||
print_layer_grid(results)
|
||||
print_weight_stats(results)
|
||||
all_results[label] = results
|
||||
return results
|
||||
|
||||
analyze_one(args.ckpt_dir, "Primary")
|
||||
|
||||
if args.compare:
|
||||
for cdir in args.compare:
|
||||
analyze_one(cdir, f"Compare_{cdir}")
|
||||
|
||||
if args.output:
|
||||
with open(args.output, "w", encoding="utf-8") as f:
|
||||
json.dump(all_results, f, indent=2)
|
||||
print(f"\nResults saved to {args.output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,406 @@
|
||||
"""HumanEval benchmark — functional pipeline design.
|
||||
|
||||
Pipeline:
|
||||
load -> generate -> extract -> test -> score -> report
|
||||
|
||||
Each stage is a pure function (except GPU/CPU-bound I/O stages).
|
||||
Config is a single dataclass; side effects are isolated at pipeline boundaries.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import subprocess
|
||||
import sys
|
||||
from dataclasses import dataclass
|
||||
from math import prod
|
||||
from typing import Dict, Iterator, List, Optional, Sequence, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import tqdm
|
||||
from datasets import load_dataset
|
||||
|
||||
from astrai.inference import InferenceEngine
|
||||
from astrai.model import AutoModel
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Config
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
HUMANEVAL_HF_DATASET = "openai/openai_humaneval"
|
||||
|
||||
STOP_SEQUENCES = [
|
||||
"\nclass ",
|
||||
"\ndef ",
|
||||
"\n# ",
|
||||
"\nif __name__",
|
||||
"\nprint(",
|
||||
"\n\n\n",
|
||||
]
|
||||
|
||||
|
||||
@dataclass
|
||||
class EvalConfig:
|
||||
param_path: str = "./params"
|
||||
data_path: str = "./humaneval/HumanEval.jsonl"
|
||||
output: Optional[str] = None
|
||||
|
||||
test_only: Optional[str] = None
|
||||
generate_only: bool = False
|
||||
|
||||
num_samples: int = 200
|
||||
max_tokens: int = 512
|
||||
temperature: float = 0.8
|
||||
top_p: float = 0.95
|
||||
top_k: int = 50
|
||||
batch_size: int = 32
|
||||
test_timeout: float = 3.0
|
||||
test_workers: int = 8
|
||||
k_values: Tuple[int, ...] = (1, 10, 100)
|
||||
problem_indices: Optional[List[int]] = None
|
||||
|
||||
|
||||
def download(path: str):
|
||||
if os.path.exists(path):
|
||||
return
|
||||
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
|
||||
print(f"Downloading HumanEval from HuggingFace ({HUMANEVAL_HF_DATASET}) ...")
|
||||
ds = load_dataset(HUMANEVAL_HF_DATASET, split="test")
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
for item in ds:
|
||||
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
||||
print(f" saved {len(ds)} problems to {path}")
|
||||
|
||||
|
||||
def load_jsonl(path: str) -> List[dict]:
|
||||
rows = []
|
||||
with open(path, encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if line:
|
||||
rows.append(json.loads(line))
|
||||
return rows
|
||||
|
||||
|
||||
def save_json(path: str, data):
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
json.dump(data, f, indent=2, ensure_ascii=False)
|
||||
|
||||
|
||||
def create_engine(param_path: str, batch_size: int) -> InferenceEngine:
|
||||
model = AutoModel.from_pretrained(param_path)
|
||||
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||
model.to(device="cuda", dtype=torch.bfloat16)
|
||||
return InferenceEngine(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=batch_size,
|
||||
)
|
||||
|
||||
|
||||
def trim_stop(text: str) -> str:
|
||||
for stop in STOP_SEQUENCES:
|
||||
idx = text.find(stop)
|
||||
if idx != -1:
|
||||
text = text[:idx]
|
||||
return text
|
||||
|
||||
|
||||
def extract_body(code: str, entry_point: str) -> Optional[str]:
|
||||
pattern = rf"def\s+{re.escape(entry_point)}\b[^:]*:"
|
||||
match = re.search(pattern, code)
|
||||
if not match:
|
||||
return code
|
||||
|
||||
lines = code[match.end() :].split("\n")
|
||||
body_lines = []
|
||||
started = False
|
||||
|
||||
for line in lines:
|
||||
stripped = line.rstrip()
|
||||
if not stripped and not started:
|
||||
continue
|
||||
if not stripped and started:
|
||||
body_lines.append("")
|
||||
continue
|
||||
if not started:
|
||||
started = True
|
||||
if stripped.lstrip() == stripped and started:
|
||||
break
|
||||
body_lines.append(stripped)
|
||||
|
||||
body = "\n".join(body_lines)
|
||||
return body if body.strip() else None
|
||||
|
||||
|
||||
def deduplicate(seq: Sequence[str]) -> List[str]:
|
||||
seen = set()
|
||||
return [x for x in seq if not (x in seen or seen.add(x))]
|
||||
|
||||
|
||||
def generate_batch(
|
||||
engine: InferenceEngine,
|
||||
prompt: str,
|
||||
n: int,
|
||||
batch_size: int,
|
||||
max_tokens: int,
|
||||
temperature: float,
|
||||
top_p: float,
|
||||
top_k: int,
|
||||
) -> List[str]:
|
||||
completions = []
|
||||
remaining = n
|
||||
while remaining > 0:
|
||||
current = min(batch_size, remaining)
|
||||
outputs = engine.generate(
|
||||
prompt=[prompt] * current,
|
||||
stream=False,
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
)
|
||||
completions.extend(outputs if isinstance(outputs, list) else [outputs])
|
||||
remaining -= current
|
||||
return deduplicate(completions)
|
||||
|
||||
|
||||
def extract_completions(
|
||||
raw: Sequence[str],
|
||||
entry_point: str,
|
||||
) -> List[str]:
|
||||
bodies = []
|
||||
for r in raw:
|
||||
t = trim_stop(r)
|
||||
body = extract_body(t, entry_point)
|
||||
if body:
|
||||
bodies.append(body)
|
||||
return bodies
|
||||
|
||||
|
||||
def generate_all(
|
||||
engine: InferenceEngine,
|
||||
problems: Sequence[dict],
|
||||
cfg: EvalConfig,
|
||||
) -> List[dict]:
|
||||
results = []
|
||||
for problem in tqdm.tqdm(problems, desc="Generating", unit="problem"):
|
||||
raw = generate_batch(
|
||||
engine,
|
||||
problem["prompt"],
|
||||
cfg.num_samples,
|
||||
cfg.batch_size,
|
||||
cfg.max_tokens,
|
||||
cfg.temperature,
|
||||
cfg.top_p,
|
||||
cfg.top_k,
|
||||
)
|
||||
bodies = extract_completions(raw, problem["entry_point"])
|
||||
results.append(
|
||||
dict(
|
||||
task_id=problem["task_id"],
|
||||
entry_point=problem["entry_point"],
|
||||
prompt=problem["prompt"],
|
||||
test=problem["test"],
|
||||
completions=bodies,
|
||||
)
|
||||
)
|
||||
return results
|
||||
|
||||
|
||||
def execute_one(args: tuple) -> bool:
|
||||
full_code, entry_point, timeout = args
|
||||
try:
|
||||
r = subprocess.run(
|
||||
[sys.executable, "-c", full_code],
|
||||
capture_output=True,
|
||||
timeout=timeout,
|
||||
)
|
||||
return r.returncode == 0
|
||||
except subprocess.TimeoutExpired:
|
||||
return False
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def test_one(item: dict, cfg: EvalConfig, pool=None) -> Tuple[str, int, int]:
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
|
||||
task_id = item["task_id"]
|
||||
completions = item["completions"]
|
||||
codes = [
|
||||
(
|
||||
item["prompt"] + c + "\n" + item["test"],
|
||||
item["entry_point"],
|
||||
cfg.test_timeout,
|
||||
)
|
||||
for c in completions
|
||||
]
|
||||
n = len(codes)
|
||||
|
||||
def _run(p):
|
||||
return sum(1 for ok in p.map(execute_one, codes) if ok)
|
||||
|
||||
if pool is not None:
|
||||
passed = _run(pool)
|
||||
else:
|
||||
with ProcessPoolExecutor(max_workers=cfg.test_workers) as p:
|
||||
passed = _run(p)
|
||||
|
||||
return task_id, n, passed
|
||||
|
||||
|
||||
def test_all(
|
||||
items: Sequence[dict],
|
||||
cfg: EvalConfig,
|
||||
) -> Iterator[Tuple[str, int, int]]:
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
|
||||
pool = ProcessPoolExecutor(max_workers=cfg.test_workers)
|
||||
try:
|
||||
for item in tqdm.tqdm(items, desc="Testing", unit="problem"):
|
||||
yield test_one(item, cfg, pool)
|
||||
finally:
|
||||
pool.shutdown(wait=True)
|
||||
|
||||
|
||||
def pass_at_k(n: int, c: int, k: int) -> float:
|
||||
if n - c < k:
|
||||
return 1.0
|
||||
return 1.0 - float(prod(1.0 - k / np.arange(n - c + 1, n + 1)))
|
||||
|
||||
|
||||
def score_results(
|
||||
results: Iterator[Tuple[str, int, int]],
|
||||
k_values: Tuple[int, ...],
|
||||
) -> Dict:
|
||||
"""Score pass@k for each problem.
|
||||
|
||||
k values are filtered per-problem: if a problem has n < k samples
|
||||
(e.g. after deduplication), pass@k is not computed for that problem.
|
||||
The summary averages only over problems where the k was computed.
|
||||
"""
|
||||
scores = {k: [] for k in k_values}
|
||||
output = {}
|
||||
for task_id, n, passed in results:
|
||||
entry = {"task_id": task_id, "n": n, "passed": passed}
|
||||
for k in k_values:
|
||||
if k <= n:
|
||||
pk = round(pass_at_k(n, passed, k), 4)
|
||||
entry[f"pass@{k}"] = pk
|
||||
scores[k].append(pk)
|
||||
else:
|
||||
entry[f"pass@{k}"] = None
|
||||
output[task_id] = entry
|
||||
|
||||
summary = {}
|
||||
for k in k_values:
|
||||
vals = scores[k]
|
||||
if vals:
|
||||
summary[f"pass@{k}"] = round(float(np.mean(vals)), 4)
|
||||
else:
|
||||
summary[f"pass@{k}"] = None
|
||||
output["_summary"] = summary
|
||||
return output
|
||||
|
||||
|
||||
def run_pipeline(cfg: EvalConfig) -> Dict:
|
||||
if cfg.test_only:
|
||||
with open(cfg.test_only, encoding="utf-8") as f:
|
||||
generated = json.load(f)
|
||||
else:
|
||||
download(cfg.data_path)
|
||||
|
||||
problems = load_jsonl(cfg.data_path)
|
||||
if cfg.problem_indices:
|
||||
problems = [problems[i] for i in cfg.problem_indices if i < len(problems)]
|
||||
|
||||
engine = create_engine(cfg.param_path, cfg.batch_size)
|
||||
|
||||
try:
|
||||
generated = generate_all(engine, problems, cfg)
|
||||
finally:
|
||||
engine.shutdown()
|
||||
|
||||
if cfg.output:
|
||||
mid = cfg.output.replace(".json", "_completions.json")
|
||||
save_json(mid, generated)
|
||||
print(f"Completions saved to {mid}")
|
||||
|
||||
if cfg.generate_only:
|
||||
return {}
|
||||
|
||||
results = test_all(generated, cfg)
|
||||
scored = score_results(results, cfg.k_values)
|
||||
return scored
|
||||
|
||||
|
||||
def parse_args(argv: Optional[List[str]] = None) -> EvalConfig:
|
||||
p = argparse.ArgumentParser(description="HumanEval benchmark")
|
||||
p.add_argument("--param_path", type=str, default="./params")
|
||||
p.add_argument("--data_path", type=str, default="./humaneval/HumanEval.jsonl")
|
||||
p.add_argument("--output", type=str, default=None)
|
||||
p.add_argument(
|
||||
"--test_only",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Skip generation, test existing completions JSON",
|
||||
)
|
||||
p.add_argument(
|
||||
"--generate_only", action="store_true", help="Only generate, skip testing"
|
||||
)
|
||||
p.add_argument("--num_samples", type=int, default=200)
|
||||
p.add_argument("--max_tokens", type=int, default=512)
|
||||
p.add_argument("--temperature", type=float, default=0.8)
|
||||
p.add_argument("--top_p", type=float, default=0.95)
|
||||
p.add_argument("--top_k", type=int, default=50)
|
||||
p.add_argument("--batch_size", type=int, default=32)
|
||||
p.add_argument("--test_workers", type=int, default=8)
|
||||
p.add_argument("--test_timeout", type=float, default=3.0)
|
||||
p.add_argument("--problems", type=int, nargs="+", default=None)
|
||||
args = p.parse_args(argv)
|
||||
|
||||
return EvalConfig(
|
||||
param_path=args.param_path,
|
||||
data_path=args.data_path,
|
||||
output=args.output,
|
||||
test_only=args.test_only,
|
||||
generate_only=args.generate_only,
|
||||
num_samples=args.num_samples,
|
||||
max_tokens=args.max_tokens,
|
||||
temperature=args.temperature,
|
||||
top_p=args.top_p,
|
||||
top_k=args.top_k,
|
||||
batch_size=args.batch_size,
|
||||
test_workers=args.test_workers,
|
||||
test_timeout=args.test_timeout,
|
||||
problem_indices=args.problems,
|
||||
)
|
||||
|
||||
|
||||
def report(scored: Dict):
|
||||
summary = scored.pop("_summary", {})
|
||||
print(f"\n{'=' * 60}")
|
||||
for k, v in summary.items():
|
||||
if v is not None:
|
||||
print(f" {k}: {v:.2%}")
|
||||
else:
|
||||
print(f" {k}: N/A")
|
||||
print(f"{'=' * 60}")
|
||||
scored["_summary"] = summary
|
||||
|
||||
|
||||
def main():
|
||||
cfg = parse_args()
|
||||
scored = run_pipeline(cfg)
|
||||
report(scored)
|
||||
if cfg.output:
|
||||
save_json(cfg.output, scored)
|
||||
print(f"Results saved to {cfg.output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,506 @@
|
||||
"""IFD (Instruction Following Difficulty) data quality scoring.
|
||||
|
||||
IFD = conditional_NLL / unconditional_NLL
|
||||
|
||||
- Messages format: plain text concatenation (no chat template)
|
||||
- Plain format: raw instr_key + resp_key fields
|
||||
|
||||
v2 changelog:
|
||||
- Same token set: unconditional pass prefixes resp with a plain-text sentinel
|
||||
(default ``\\n``; use ``--sentinel_text ""`` for bos/pad fallback).
|
||||
Both branches predict the identical N resp tokens.
|
||||
Single-token answers (rl=1) are now supported.
|
||||
- ctx_len tracked in output
|
||||
- skip_reason for None samples (no more silent None)
|
||||
- --per_token for per-token IFD breakdown
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import statistics
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import tqdm
|
||||
|
||||
from astrai.model import AutoModel
|
||||
from astrai.preprocessing.packing import plan_bfd
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
|
||||
def _pack_bins(pairs, max_len):
|
||||
"""BFD bin packing: pack (c+r) into bins of max total length.
|
||||
|
||||
Reuses :func:`plan_bfd` so the BFD heuristic stays single-sourced.
|
||||
"""
|
||||
# Treat each pair as a single sequence of length len(c)+len(r) for
|
||||
# planning purposes; plan_bfd works on pure lengths.
|
||||
fake_sequences = [[0] * (len(c) + len(r)) for c, r in pairs]
|
||||
plan = plan_bfd(fake_sequences, max_len)
|
||||
return [
|
||||
[(i, pairs[i][0], pairs[i][1]) for i in bin_indices] for bin_indices in plan
|
||||
]
|
||||
|
||||
|
||||
def _resolve_sentinel_ids(tokenizer, sentinel_text):
|
||||
"""Tokenize the sentinel text for the unconditional pass prefix.
|
||||
|
||||
Falls back to bos/pad_token_id when sentinel_text is empty or
|
||||
cannot be encoded.
|
||||
"""
|
||||
if sentinel_text:
|
||||
ids = tokenizer.encode(sentinel_text, add_special_tokens=False)
|
||||
if ids:
|
||||
return ids
|
||||
for attr in ("bos_token_id", "pad_token_id", "eos_token_id"):
|
||||
tid = getattr(tokenizer, attr, None)
|
||||
if tid is not None:
|
||||
return [tid]
|
||||
return [0]
|
||||
|
||||
|
||||
def _collect_input_files(input_path: str) -> list:
|
||||
"""Resolve *input_path* to a list of JSONL/JSON files."""
|
||||
if os.path.isdir(input_path):
|
||||
files = []
|
||||
for ext in ("*.jsonl", "*.json"):
|
||||
files.extend(
|
||||
sorted(glob.glob(os.path.join(input_path, "**", ext), recursive=True))
|
||||
)
|
||||
return files
|
||||
return sorted(glob.glob(input_path))
|
||||
|
||||
|
||||
def _load_items(filepath: str) -> list:
|
||||
"""Load JSONL or JSON (array / single dict) into a list of dicts."""
|
||||
with open(filepath, "r", encoding="utf-8") as f:
|
||||
if filepath.lower().endswith(".json"):
|
||||
data = json.load(f)
|
||||
if isinstance(data, dict):
|
||||
return [data]
|
||||
return data
|
||||
return [json.loads(line) for line in f if line.strip()]
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def _score_batch(
|
||||
pairs, model, device, max_len=2048, sentinel_ids=None, per_token=False
|
||||
):
|
||||
"""BFD-packed IFD with text-sentinel-anchored unconditional pass.
|
||||
|
||||
Conditional: (ctx + resp[0..i-1]) → resp[i], i = 0..N-1
|
||||
Unconditional: (<sentinel> + resp[0..i-1]) → resp[i], i = 0..N-1
|
||||
|
||||
Both branches predict the identical N response tokens. A short
|
||||
plain-text sentinel gives the unconditional pass a prefix so that
|
||||
every response token can be predicted. Single-token answers (rl=1)
|
||||
are supported.
|
||||
"""
|
||||
if not pairs:
|
||||
return []
|
||||
|
||||
if sentinel_ids is None:
|
||||
sentinel_ids = [0]
|
||||
|
||||
bins = _pack_bins(pairs, max_len)
|
||||
result = [None] * len(pairs)
|
||||
|
||||
# ---- conditional pass (packed, per-document position IDs) ----
|
||||
for bin_items in bins:
|
||||
seq_ids = []
|
||||
global_pos = []
|
||||
doc_ids = []
|
||||
doc_offsets = []
|
||||
|
||||
for di, (orig_idx, c, r) in enumerate(bin_items):
|
||||
ctx_len = len(c)
|
||||
start = len(seq_ids)
|
||||
item_len = len(c) + len(r)
|
||||
seq_ids.extend(c)
|
||||
seq_ids.extend(r)
|
||||
end = len(seq_ids)
|
||||
global_pos.extend(range(item_len))
|
||||
doc_ids.extend([di] * item_len)
|
||||
doc_offsets.append((start, end, orig_idx, ctx_len))
|
||||
|
||||
full_ids = torch.tensor([seq_ids], device=device, dtype=torch.long)
|
||||
pos_ids = torch.tensor([global_pos], device=device, dtype=torch.long)
|
||||
seq_len = len(seq_ids)
|
||||
causal = torch.tril(
|
||||
torch.ones(seq_len, seq_len, dtype=torch.bool, device=device)
|
||||
)
|
||||
doc_t = torch.tensor([doc_ids], device=device)
|
||||
doc_mask = doc_t.unsqueeze(-1) == doc_t.unsqueeze(-2)
|
||||
attn_mask = (causal & doc_mask[0]).unsqueeze(0).unsqueeze(0)
|
||||
logits_full = model(full_ids, position_ids=pos_ids, input_mask=attn_mask)[
|
||||
"logits"
|
||||
][0]
|
||||
|
||||
for start, end, orig_idx, ctx_len in doc_offsets:
|
||||
rl = end - start - ctx_len
|
||||
resp_start = start + ctx_len - 1
|
||||
resp_logits = logits_full[resp_start : end - 1]
|
||||
resp_targets = torch.tensor(
|
||||
seq_ids[start + ctx_len : end], device=device, dtype=torch.long
|
||||
)
|
||||
cond_losses = F.cross_entropy(
|
||||
resp_logits, resp_targets, reduction="none"
|
||||
).cpu()
|
||||
result[orig_idx] = {
|
||||
"_cond_losses": cond_losses,
|
||||
"_rl": rl,
|
||||
"_ctx_len": ctx_len,
|
||||
}
|
||||
|
||||
# ---- unconditional pass (sentinel-prefixed, batched 2D) ----
|
||||
valid_items = [
|
||||
(
|
||||
i,
|
||||
result[i]["_rl"],
|
||||
result[i]["_ctx_len"],
|
||||
result[i]["_cond_losses"],
|
||||
pairs[i][1],
|
||||
)
|
||||
for i in range(len(pairs))
|
||||
if result[i] is not None and "_cond_losses" in result[i]
|
||||
]
|
||||
if not valid_items:
|
||||
return result
|
||||
|
||||
valid_items.sort(key=lambda x: -x[1])
|
||||
prefix_len = len(sentinel_ids)
|
||||
max_rl = prefix_len + max(rl for _, rl, _, _, _ in valid_items)
|
||||
bsz = len(valid_items)
|
||||
|
||||
u_batch = torch.zeros(bsz, max_rl, dtype=torch.long, device=device)
|
||||
for ri, (_, rl, _, _, r_ids) in enumerate(valid_items):
|
||||
u_batch[ri, :prefix_len] = torch.tensor(sentinel_ids, dtype=torch.long)
|
||||
u_batch[ri, prefix_len : prefix_len + rl] = torch.tensor(
|
||||
r_ids, dtype=torch.long
|
||||
)
|
||||
|
||||
logits_resp = model(u_batch)["logits"]
|
||||
|
||||
for ri, (orig_idx, rl, ctx_len, cond_losses, _) in enumerate(valid_items):
|
||||
unp_logits = logits_resp[ri, prefix_len - 1 : prefix_len - 1 + rl]
|
||||
unp_targets = u_batch[ri, prefix_len : prefix_len + rl]
|
||||
uncond_losses = F.cross_entropy(unp_logits, unp_targets, reduction="none").cpu()
|
||||
|
||||
L_cond = cond_losses.mean().item()
|
||||
L_uncond = uncond_losses.mean().item()
|
||||
ifd = L_cond / L_uncond if L_uncond > 0 else None
|
||||
|
||||
out = {
|
||||
"L_cond": round(L_cond, 6),
|
||||
"L_uncond": round(L_uncond, 6),
|
||||
"ifd": round(ifd, 6) if ifd is not None else None,
|
||||
"ctx_len": ctx_len,
|
||||
"resp_len": rl,
|
||||
}
|
||||
if per_token:
|
||||
per = [
|
||||
(round(c.item() / u.item(), 6) if u.item() > 0 else None)
|
||||
for c, u in zip(cond_losses, uncond_losses)
|
||||
]
|
||||
out["ifd_per_token"] = per
|
||||
result[orig_idx] = out
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def _trim(context_ids, resp_ids, max_len):
|
||||
"""Truncate to fit max_len, keeping response intact if possible."""
|
||||
if len(resp_ids) > max_len // 2:
|
||||
resp_ids = resp_ids[: max_len // 2]
|
||||
full_ids = context_ids + resp_ids
|
||||
if len(full_ids) <= max_len:
|
||||
return context_ids, resp_ids
|
||||
overflow = len(full_ids) - max_len
|
||||
if overflow >= len(context_ids):
|
||||
return [], resp_ids[:max_len]
|
||||
return context_ids[overflow:], resp_ids
|
||||
|
||||
|
||||
def process_file(
|
||||
model,
|
||||
tokenizer,
|
||||
input_file,
|
||||
output_file,
|
||||
instr_key,
|
||||
resp_key,
|
||||
max_len=2048,
|
||||
data_format="plain",
|
||||
batch_size=1,
|
||||
device=None,
|
||||
sentinel_ids=None,
|
||||
per_token=False,
|
||||
max_samples=None,
|
||||
):
|
||||
"""Score a single file, write per-sample JSONL, return summary stats."""
|
||||
if device is None:
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
if sentinel_ids is None:
|
||||
sentinel_ids = _resolve_sentinel_ids(tokenizer, "\n")
|
||||
|
||||
data = _load_items(input_file)
|
||||
|
||||
if max_samples and len(data) > max_samples:
|
||||
import random
|
||||
|
||||
data = random.sample(data, max_samples)
|
||||
|
||||
results = []
|
||||
all_ifds = []
|
||||
buffer = []
|
||||
|
||||
label = os.path.splitext(os.path.basename(input_file))[0]
|
||||
|
||||
for item in tqdm.tqdm(data, desc=f" {label}", unit="sample", leave=False):
|
||||
if data_format == "messages":
|
||||
turns = []
|
||||
for i, msg in enumerate(item.get("messages", [])):
|
||||
if msg.get("role") != "assistant":
|
||||
continue
|
||||
ctx_text = "\n\n".join(m["content"] for m in item["messages"][:i])
|
||||
ctx_ids = tokenizer.encode(ctx_text)
|
||||
resp_ids = tokenizer.encode(msg["content"], add_special_tokens=False)
|
||||
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len)
|
||||
if ctx_ids and resp_ids:
|
||||
turns.append((ctx_ids, resp_ids))
|
||||
if not turns:
|
||||
results.append(
|
||||
{
|
||||
**item,
|
||||
"ifd": None,
|
||||
"skip_reason": "no valid assistant turns",
|
||||
"ifd_turns": [],
|
||||
}
|
||||
)
|
||||
continue
|
||||
buffer.append((item, turns, "messages"))
|
||||
else:
|
||||
ctx_ids = tokenizer.encode(item[instr_key], add_special_tokens=False)
|
||||
resp_ids = tokenizer.encode(item[resp_key], add_special_tokens=False)
|
||||
ctx_ids, resp_ids = _trim(ctx_ids, resp_ids, max_len)
|
||||
if not ctx_ids or not resp_ids:
|
||||
results.append(
|
||||
{
|
||||
**item,
|
||||
"ifd": None,
|
||||
"ifd_detail": {"skip_reason": "empty ctx or resp"},
|
||||
}
|
||||
)
|
||||
continue
|
||||
buffer.append((item, [(ctx_ids, resp_ids)], "plain"))
|
||||
|
||||
if len(buffer) >= batch_size:
|
||||
_flush_buffer(
|
||||
buffer,
|
||||
results,
|
||||
all_ifds,
|
||||
model,
|
||||
device,
|
||||
max_len,
|
||||
sentinel_ids,
|
||||
per_token,
|
||||
)
|
||||
|
||||
if buffer:
|
||||
_flush_buffer(
|
||||
buffer, results, all_ifds, model, device, max_len, sentinel_ids, per_token
|
||||
)
|
||||
|
||||
with open(output_file, "w", encoding="utf-8") as f:
|
||||
for item in results:
|
||||
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
||||
|
||||
valid_ifd = [v for v in all_ifds if v is not None]
|
||||
stats = {
|
||||
"samples": len(data),
|
||||
"valid_ifd": len(valid_ifd),
|
||||
"skipped": len(data) - len(valid_ifd),
|
||||
}
|
||||
if valid_ifd:
|
||||
stats["mean_ifd"] = statistics.mean(valid_ifd)
|
||||
stats["median_ifd"] = statistics.median(valid_ifd)
|
||||
if len(valid_ifd) > 1:
|
||||
stats["stdev_ifd"] = statistics.stdev(valid_ifd)
|
||||
stats["min_ifd"] = min(valid_ifd)
|
||||
stats["max_ifd"] = max(valid_ifd)
|
||||
|
||||
print(f"\n{'=' * 50}")
|
||||
print(f" [{label}]")
|
||||
print(f"{'=' * 50}")
|
||||
print(f" Samples: {len(data)}")
|
||||
print(f" Valid IFD: {len(valid_ifd)}")
|
||||
print(f" Skipped: {len(data) - len(valid_ifd)}")
|
||||
print(f" Mean IFD: {statistics.mean(valid_ifd):.4f}")
|
||||
print(f" Median IFD: {statistics.median(valid_ifd):.4f}")
|
||||
if len(valid_ifd) > 1:
|
||||
print(f" Stdev IFD: {statistics.stdev(valid_ifd):.4f}")
|
||||
print(f" Min IFD: {min(valid_ifd):.4f}")
|
||||
print(f" Max IFD: {max(valid_ifd):.4f}")
|
||||
print(f"{'=' * 50}")
|
||||
print(f" Results saved to {output_file}")
|
||||
return stats
|
||||
|
||||
|
||||
def _flush_buffer(
|
||||
buffer, results, all_ifds, model, device, max_len, sentinel_ids, per_token
|
||||
):
|
||||
all_pairs = []
|
||||
indices = []
|
||||
for item, turns, fmt in buffer:
|
||||
start = len(all_pairs)
|
||||
all_pairs.extend(turns)
|
||||
indices.append((item, turns, fmt, start, len(all_pairs)))
|
||||
|
||||
raw = _score_batch(
|
||||
all_pairs,
|
||||
model,
|
||||
device,
|
||||
max_len,
|
||||
sentinel_ids=sentinel_ids,
|
||||
per_token=per_token,
|
||||
)
|
||||
|
||||
for item, turns, fmt, start, end in indices:
|
||||
turn_scores = raw[start:end]
|
||||
if fmt == "messages":
|
||||
valid = [
|
||||
s for s in turn_scores if s is not None and s.get("ifd") is not None
|
||||
]
|
||||
if not valid:
|
||||
results.append({**item, "ifd": None, "ifd_turns": turn_scores})
|
||||
else:
|
||||
avg = sum(s["ifd"] for s in valid) / len(valid)
|
||||
all_ifds.append(avg)
|
||||
results.append(
|
||||
{
|
||||
**item,
|
||||
"ifd": avg,
|
||||
"ifd_detail": valid[0] if len(valid) == 1 else None,
|
||||
"ifd_turns": turn_scores,
|
||||
}
|
||||
)
|
||||
else:
|
||||
score = turn_scores[0]
|
||||
all_ifds.append(score.get("ifd"))
|
||||
results.append({**item, "ifd": score.get("ifd"), "ifd_detail": score})
|
||||
|
||||
buffer.clear()
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Compute IFD scores for instruction-response data"
|
||||
)
|
||||
parser.add_argument("--param_path", type=str, required=True, help="Model directory")
|
||||
parser.add_argument(
|
||||
"--input_path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Input file, glob pattern, or directory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Directory for output files (summary.json + per-file JSONL).",
|
||||
)
|
||||
parser.add_argument("--max_len", type=int, default=2048, help="Max token length")
|
||||
parser.add_argument(
|
||||
"--format",
|
||||
type=str,
|
||||
default="plain",
|
||||
choices=["plain", "messages"],
|
||||
help="Input format",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--instr_key", type=str, default="instruction", help="Key for instruction field"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--resp_key", type=str, default="response", help="Key for response field"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size", type=int, default=8, help="Batch size for model forward passes"
|
||||
)
|
||||
parser.add_argument("--device", type=str, default=None, help="Device (e.g. cuda:0)")
|
||||
parser.add_argument(
|
||||
"--dtype",
|
||||
type=str,
|
||||
default="bfloat16" if torch.cuda.is_available() else "float32",
|
||||
help="Torch dtype",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sentinel_text",
|
||||
type=str,
|
||||
default="\n",
|
||||
help='Plain-text prefix for unconditional pass (default: "\\n"). Use "" for bos/pad fallback.',
|
||||
)
|
||||
parser.add_argument(
|
||||
"--per_token",
|
||||
action="store_true",
|
||||
help="Include per-token IFD breakdown in output",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_samples",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Maximum number of samples per file (random subsample). Default: all.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.device is None:
|
||||
args.device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
dtype = getattr(torch, args.dtype)
|
||||
|
||||
print(f"Loading model from {args.param_path} ...")
|
||||
model = AutoModel.from_pretrained(args.param_path)
|
||||
tokenizer = AutoTokenizer.from_pretrained(args.param_path)
|
||||
model.to(device=args.device, dtype=dtype)
|
||||
model.eval()
|
||||
|
||||
sentinel_ids = _resolve_sentinel_ids(tokenizer, args.sentinel_text)
|
||||
|
||||
input_files = _collect_input_files(args.input_path)
|
||||
if not input_files:
|
||||
print(f"No input files found at {args.input_path}")
|
||||
return
|
||||
|
||||
print(f"Found {len(input_files)} file(s) to evaluate")
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
all_stats = {}
|
||||
for filepath in input_files:
|
||||
label = os.path.splitext(os.path.basename(filepath))[0]
|
||||
output_file = os.path.join(args.output_dir, f"{label}_ifd.jsonl")
|
||||
|
||||
stats = process_file(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
input_file=filepath,
|
||||
output_file=output_file,
|
||||
instr_key=args.instr_key,
|
||||
resp_key=args.resp_key,
|
||||
max_len=args.max_len,
|
||||
data_format=args.format,
|
||||
batch_size=args.batch_size,
|
||||
device=args.device,
|
||||
sentinel_ids=sentinel_ids,
|
||||
per_token=args.per_token,
|
||||
max_samples=args.max_samples,
|
||||
)
|
||||
all_stats[label] = stats
|
||||
|
||||
summary_path = os.path.join(args.output_dir, "summary.json")
|
||||
with open(summary_path, "w", encoding="utf-8") as f:
|
||||
json.dump(all_stats, f, ensure_ascii=False, indent=2)
|
||||
print(f"\nSummary saved to {summary_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,602 @@
|
||||
"""IFEval instruction-following evaluation benchmark.
|
||||
|
||||
Evaluates model responses against regex-based constraint verifiers.
|
||||
Supports all IFEval constraint types except language detection.
|
||||
|
||||
Usage::
|
||||
|
||||
python scripts/eval/evaluate_ifeval.py --param_path ./params \
|
||||
--data_path ifeval.jsonl --output results.json \
|
||||
--temperature 0.1 --max_tokens 512
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from typing import Callable, Dict, List, Optional
|
||||
|
||||
import torch
|
||||
import tqdm
|
||||
from datasets import load_dataset
|
||||
|
||||
from astrai.inference import InferenceEngine
|
||||
from astrai.model import AutoModel
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
IFEVAL_HF_DATASET = "google/IFEval"
|
||||
CONSTRAINT_VERIFIERS: Dict[str, Callable[[str, dict], bool]] = {}
|
||||
|
||||
|
||||
def register(instruction_id: str):
|
||||
def decorator(fn):
|
||||
CONSTRAINT_VERIFIERS[instruction_id] = fn
|
||||
return fn
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
@register("keywords:existence")
|
||||
def check_keyword_existence(response: str, kwargs: dict) -> bool:
|
||||
for kw in kwargs["keywords"]:
|
||||
if not re.search(re.escape(kw), response, re.IGNORECASE):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
@register("keywords:frequency")
|
||||
def check_keyword_frequency(response: str, kwargs: dict) -> bool:
|
||||
keyword = kwargs["keyword"]
|
||||
frequency = kwargs.get("frequency", 1)
|
||||
relation = kwargs.get("relation", "at least")
|
||||
count = len(re.findall(re.escape(keyword), response, re.IGNORECASE))
|
||||
if relation == "less than":
|
||||
return count < frequency
|
||||
return count >= frequency
|
||||
|
||||
|
||||
@register("keywords:forbidden_words")
|
||||
def check_forbidden_words(response: str, kwargs: dict) -> bool:
|
||||
for word in kwargs["forbidden_words"]:
|
||||
if re.search(r"\b" + re.escape(word) + r"\b", response, re.IGNORECASE):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
@register("keywords:letter_frequency")
|
||||
def check_letter_frequency(response: str, kwargs: dict) -> bool:
|
||||
letter = kwargs["letter"].lower()
|
||||
frequency = kwargs.get("let_frequency", 1)
|
||||
relation = kwargs.get("let_relation", "at least")
|
||||
count = response.lower().count(letter)
|
||||
if relation == "less than":
|
||||
return count < frequency
|
||||
return count >= frequency
|
||||
|
||||
|
||||
@register("detectable_content:number_placeholders")
|
||||
def check_placeholders(response: str, kwargs: dict) -> bool:
|
||||
num = kwargs.get("num_placeholders", 1)
|
||||
placeholders = re.findall(r"\[.*?\]", response)
|
||||
return len(placeholders) >= num
|
||||
|
||||
|
||||
@register("detectable_content:postscript")
|
||||
def check_postscript(response: str, kwargs: dict) -> bool:
|
||||
marker = kwargs.get("postscript_marker", "P.S.")
|
||||
response_lower = response.lower()
|
||||
if marker == "P.P.S":
|
||||
return bool(re.search(r"p\.\s?p\.\s?s", response_lower))
|
||||
elif marker == "P.S.":
|
||||
return bool(re.search(r"p\.\s?s\.", response_lower))
|
||||
else:
|
||||
return bool(re.search(re.escape(marker.lower()), response_lower))
|
||||
|
||||
|
||||
@register("detectable_format:number_bullet_lists")
|
||||
def check_bullet_lists(response: str, kwargs: dict) -> bool:
|
||||
num = kwargs.get("num_bullets", 1)
|
||||
bullets = re.findall(r"^\s*\*[^\*].*$", response, re.MULTILINE)
|
||||
dashes = re.findall(r"^\s*-.*$", response, re.MULTILINE)
|
||||
return len(bullets) + len(dashes) == num
|
||||
|
||||
|
||||
@register("detectable_format:number_highlighted_sections")
|
||||
def check_highlighted_sections(response: str, kwargs: dict) -> bool:
|
||||
num = kwargs.get("num_highlights", 1)
|
||||
highlights = re.findall(r"\*[^\n\*]+\*", response)
|
||||
count = 0
|
||||
for h in highlights:
|
||||
if h.strip("*").strip():
|
||||
count += 1
|
||||
return count >= num
|
||||
|
||||
|
||||
@register("detectable_format:multiple_sections")
|
||||
def check_multiple_sections(response: str, kwargs: dict) -> bool:
|
||||
splitter = kwargs.get("section_spliter", "Section")
|
||||
num = kwargs.get("num_sections", 1)
|
||||
pattern = r"\s?" + re.escape(splitter) + r"\s?\d+\s?"
|
||||
sections = re.split(pattern, response)
|
||||
return len(sections) - 1 >= num
|
||||
|
||||
|
||||
@register("detectable_format:title")
|
||||
def check_title(response: str, kwargs: dict) -> bool:
|
||||
titles = re.findall(r"<<[^>\n]+>>", response)
|
||||
for title in titles:
|
||||
if title.strip("<>").strip():
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
@register("detectable_format:json_format")
|
||||
def check_json_format(response: str, kwargs: dict) -> bool:
|
||||
value = response.strip()
|
||||
for prefix in ("```json", "```Json", "```JSON", "```"):
|
||||
if value.lower().startswith(prefix.lower()):
|
||||
value = value[len(prefix) :].strip()
|
||||
if value.endswith("```"):
|
||||
value = value[:-3].strip()
|
||||
try:
|
||||
json.loads(value)
|
||||
return True
|
||||
except (ValueError, json.JSONDecodeError):
|
||||
return False
|
||||
|
||||
|
||||
@register("detectable_format:general_punctuation")
|
||||
def check_general_punctuation(response: str, kwargs: dict) -> bool:
|
||||
punctuation_blacklist = kwargs.get("punctuation_blacklist", [])
|
||||
for punct in punctuation_blacklist:
|
||||
if punct in response:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
@register("detectable_format:number_highlighted_words")
|
||||
def check_highlighted_words(response: str, kwargs: dict) -> bool:
|
||||
num = kwargs.get("num_highlights", 1)
|
||||
highlights = re.findall(r"\*[^\s\*][^\*]*[^\s\*]\*", response)
|
||||
return len(highlights) >= num
|
||||
|
||||
|
||||
@register("startend:end_checker")
|
||||
def check_end_checker(response: str, kwargs: dict) -> bool:
|
||||
end_phrase = kwargs["end_phrase"]
|
||||
return (
|
||||
response.strip()
|
||||
.rstrip('"')
|
||||
.rstrip()
|
||||
.lower()
|
||||
.endswith(end_phrase.strip().lower())
|
||||
)
|
||||
|
||||
|
||||
@register("startend:quotation")
|
||||
def check_quotation(response: str, kwargs: dict) -> bool:
|
||||
value = response.strip()
|
||||
return value.startswith('"') and value.endswith('"')
|
||||
|
||||
|
||||
@register("startend:start_checker")
|
||||
def check_start_checker(response: str, kwargs: dict) -> bool:
|
||||
starter = kwargs["starter"]
|
||||
return bool(re.search(r"^\s*" + re.escape(starter), response, re.MULTILINE))
|
||||
|
||||
|
||||
@register("change_case:english_capital")
|
||||
def check_english_capital(response: str, kwargs: dict) -> bool:
|
||||
return response.isupper()
|
||||
|
||||
|
||||
@register("change_case:english_lowercase")
|
||||
def check_english_lowercase(response: str, kwargs: dict) -> bool:
|
||||
return response.islower()
|
||||
|
||||
|
||||
@register("change_case:capital_word_frequency")
|
||||
def check_capital_word_frequency(response: str, kwargs: dict) -> bool:
|
||||
frequency = kwargs.get("capital_frequency", 1)
|
||||
relation = kwargs.get("capital_relation", "at least")
|
||||
capital_words = re.findall(r"\b[A-Z]{2,}\b", response)
|
||||
count = len(capital_words)
|
||||
if relation == "less than":
|
||||
return count < frequency
|
||||
return count >= frequency
|
||||
|
||||
|
||||
@register("punctuation:no_comma")
|
||||
def check_no_comma(response: str, kwargs: dict) -> bool:
|
||||
return "," not in response
|
||||
|
||||
|
||||
def count_words(text: str) -> int:
|
||||
return len(re.findall(r"\b\w+\b", text))
|
||||
|
||||
|
||||
def count_sentences(text: str) -> int:
|
||||
text = text.strip()
|
||||
if not text:
|
||||
return 0
|
||||
sentences = re.split(r"(?<=[.!?])\s+", text)
|
||||
return len([s for s in sentences if s.strip()])
|
||||
|
||||
|
||||
@register("length_constraints:number_words")
|
||||
def check_number_words(response: str, kwargs: dict) -> bool:
|
||||
num = kwargs.get("num_words", 100)
|
||||
relation = kwargs.get("relation", "at least")
|
||||
cnt = count_words(response)
|
||||
if relation == "less than":
|
||||
return cnt < num
|
||||
return cnt >= num
|
||||
|
||||
|
||||
@register("length_constraints:number_sentences")
|
||||
def check_number_sentences(response: str, kwargs: dict) -> bool:
|
||||
num = kwargs.get("num_sentences", 5)
|
||||
relation = kwargs.get("relation", "at least")
|
||||
cnt = count_sentences(response)
|
||||
if relation == "less than":
|
||||
return cnt < num
|
||||
return cnt >= num
|
||||
|
||||
|
||||
@register("length_constraints:number_paragraphs")
|
||||
def check_number_paragraphs(response: str, kwargs: dict) -> bool:
|
||||
num = kwargs.get("num_paragraphs", 1)
|
||||
if "***" in response:
|
||||
paragraphs = re.split(r"\s?\*\*\*\s?", response)
|
||||
else:
|
||||
paragraphs = re.split(r"\n\n+", response)
|
||||
actual = len([p for p in paragraphs if p.strip()])
|
||||
return actual == num
|
||||
|
||||
|
||||
@register("length_constraints:nth_paragraph_first_word")
|
||||
def check_nth_paragraph_first_word(response: str, kwargs: dict) -> bool:
|
||||
num_paragraphs = kwargs.get("num_paragraphs", 1)
|
||||
nth = kwargs.get("nth_paragraph", 1)
|
||||
first_word = kwargs.get("first_word", "").lower()
|
||||
|
||||
paragraphs = re.split(r"\n\n+", response)
|
||||
paragraphs = [p.strip() for p in paragraphs if p.strip()]
|
||||
|
||||
if len(paragraphs) != num_paragraphs:
|
||||
return False
|
||||
if nth > len(paragraphs):
|
||||
return False
|
||||
|
||||
target = paragraphs[nth - 1]
|
||||
words = target.split()
|
||||
if not words:
|
||||
return False
|
||||
|
||||
word = words[0].strip().lstrip("'\"").rstrip(".,!?:;\"'")
|
||||
return word.lower() == first_word
|
||||
|
||||
|
||||
@register("length_constraints:nth_word_checker")
|
||||
def check_nth_word(response: str, kwargs: dict) -> bool:
|
||||
nth = kwargs.get("nth_word", 1)
|
||||
target = kwargs.get("target_word", "").lower()
|
||||
words = re.findall(r"\b\w+\b", response)
|
||||
if nth > len(words):
|
||||
return False
|
||||
return words[nth - 1].lower() == target
|
||||
|
||||
|
||||
@register("combination:repeat_prompt")
|
||||
def check_repeat_prompt(response: str, kwargs: dict) -> bool:
|
||||
prompt = kwargs["prompt_to_repeat"]
|
||||
return response.strip().lower().startswith(prompt.strip().lower())
|
||||
|
||||
|
||||
@register("combination:two_responses")
|
||||
def check_two_responses(response: str, kwargs: dict) -> bool:
|
||||
parts = response.split("******")
|
||||
valid = [p for p in parts if p.strip()]
|
||||
if len(valid) != 2:
|
||||
return False
|
||||
return valid[0].strip() != valid[1].strip()
|
||||
|
||||
|
||||
def download_ifeval(data_path: str):
|
||||
if os.path.exists(data_path):
|
||||
return
|
||||
os.makedirs(os.path.dirname(data_path) or ".", exist_ok=True)
|
||||
print(f"Downloading IFEval from HuggingFace ({IFEVAL_HF_DATASET}) ...")
|
||||
ds = load_dataset(IFEVAL_HF_DATASET, split="train")
|
||||
with open(data_path, "w", encoding="utf-8") as f:
|
||||
for item in ds:
|
||||
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
||||
print(f" saved {len(ds)} items to {data_path}")
|
||||
|
||||
|
||||
def load_problems(data_path: str) -> List[dict]:
|
||||
problems = []
|
||||
with open(data_path, "r", encoding="utf-8") as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if line:
|
||||
problems.append(json.loads(line))
|
||||
return problems
|
||||
|
||||
|
||||
def verify_response(response: str, instruction_id: str, kwargs: dict) -> Optional[bool]:
|
||||
verifier = CONSTRAINT_VERIFIERS.get(instruction_id)
|
||||
if verifier is None:
|
||||
return None
|
||||
try:
|
||||
return verifier(response, kwargs)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def generate_one(
|
||||
engine: InferenceEngine,
|
||||
tokenizer: AutoTokenizer,
|
||||
prompt: str,
|
||||
max_tokens: int,
|
||||
temperature: float,
|
||||
top_p: float,
|
||||
top_k: int,
|
||||
) -> str:
|
||||
formatted = tokenizer.apply_chat_template(
|
||||
[{"role": "user", "content": prompt}],
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
)
|
||||
output = engine.generate(
|
||||
prompt=formatted,
|
||||
stream=False,
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
)
|
||||
if isinstance(output, list):
|
||||
return output[0]
|
||||
return output
|
||||
|
||||
|
||||
def evaluate(
|
||||
engine: InferenceEngine,
|
||||
tokenizer: AutoTokenizer,
|
||||
problems: List[dict],
|
||||
max_tokens: int,
|
||||
temperature: float,
|
||||
top_p: float,
|
||||
top_k: int,
|
||||
num_samples: int = 1,
|
||||
) -> Dict:
|
||||
results = {}
|
||||
constraint_stats: Dict[str, Dict[str, int]] = {}
|
||||
total_constraints = 0
|
||||
total_passed = 0
|
||||
|
||||
for problem in tqdm.tqdm(problems, desc="IFEval", unit="problem"):
|
||||
key = problem["key"]
|
||||
prompt = problem["prompt"]
|
||||
instruction_ids = problem["instruction_id_list"]
|
||||
kwargs_list = problem["kwargs"]
|
||||
|
||||
samples = []
|
||||
for _ in range(num_samples):
|
||||
response = generate_one(
|
||||
engine, tokenizer, prompt, max_tokens, temperature, top_p, top_k
|
||||
)
|
||||
samples.append(response)
|
||||
|
||||
constraint_results = []
|
||||
passed = 0
|
||||
verified = 0
|
||||
|
||||
for idx, instruction_id in enumerate(instruction_ids):
|
||||
kwargs = kwargs_list[idx] if idx < len(kwargs_list) else {}
|
||||
best_pass = False
|
||||
for response in samples:
|
||||
result = verify_response(response, instruction_id, kwargs)
|
||||
if result is None:
|
||||
continue
|
||||
if result:
|
||||
best_pass = True
|
||||
break
|
||||
|
||||
verifier_exists = instruction_id in CONSTRAINT_VERIFIERS
|
||||
if verifier_exists:
|
||||
verified += 1
|
||||
if best_pass:
|
||||
passed += 1
|
||||
|
||||
constraint_results.append(
|
||||
{
|
||||
"instruction_id": instruction_id,
|
||||
"passed": best_pass,
|
||||
"supported": verifier_exists,
|
||||
"kwargs": kwargs,
|
||||
}
|
||||
)
|
||||
|
||||
if verifier_exists:
|
||||
if instruction_id not in constraint_stats:
|
||||
constraint_stats[instruction_id] = {
|
||||
"total": 0,
|
||||
"passed": 0,
|
||||
}
|
||||
constraint_stats[instruction_id]["total"] += 1
|
||||
if best_pass:
|
||||
constraint_stats[instruction_id]["passed"] += 1
|
||||
|
||||
total_constraints += verified
|
||||
total_passed += passed
|
||||
|
||||
accuracy = passed / verified if verified > 0 else None
|
||||
results[str(key)] = {
|
||||
"key": key,
|
||||
"prompt": prompt,
|
||||
"response": samples[0],
|
||||
"num_samples": num_samples,
|
||||
"num_constraints": len(instruction_ids),
|
||||
"num_verified": verified,
|
||||
"num_passed": passed,
|
||||
"accuracy": round(accuracy, 4) if accuracy is not None else None,
|
||||
"constraints": constraint_results,
|
||||
}
|
||||
|
||||
overall_accuracy = (
|
||||
round(total_passed / total_constraints, 4) if total_constraints > 0 else 0.0
|
||||
)
|
||||
|
||||
type_summary = {}
|
||||
for inst_id, stats in sorted(constraint_stats.items()):
|
||||
type_summary[inst_id] = {
|
||||
"total": stats["total"],
|
||||
"passed": stats["passed"],
|
||||
"accuracy": round(stats["passed"] / stats["total"], 4)
|
||||
if stats["total"] > 0
|
||||
else 0.0,
|
||||
}
|
||||
|
||||
unsupported_count = sum(
|
||||
1
|
||||
for p in problems
|
||||
for iid in p["instruction_id_list"]
|
||||
if iid not in CONSTRAINT_VERIFIERS
|
||||
)
|
||||
|
||||
results["_summary"] = {
|
||||
"total_problems": len(problems),
|
||||
"total_constraints": total_constraints,
|
||||
"total_passed": total_passed,
|
||||
"overall_accuracy": overall_accuracy,
|
||||
"unsupported_constraints": unsupported_count,
|
||||
"supported_types": sorted(CONSTRAINT_VERIFIERS.keys()),
|
||||
"per_type_accuracy": type_summary,
|
||||
}
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="IFEval benchmark")
|
||||
parser.add_argument(
|
||||
"--param_path", type=str, default="./params", help="Model directory"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--data_path",
|
||||
type=str,
|
||||
default="./ifeval/input_data.jsonl",
|
||||
help="IFEval JSONL file (auto-download if missing)",
|
||||
)
|
||||
parser.add_argument("--output", type=str, default=None, help="Output JSON path")
|
||||
parser.add_argument(
|
||||
"--max_tokens", type=int, default=512, help="Max generation tokens"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--temperature",
|
||||
type=float,
|
||||
default=0.1,
|
||||
help="Sampling temperature",
|
||||
)
|
||||
parser.add_argument("--top_p", type=float, default=0.95, help="Top-p sampling")
|
||||
parser.add_argument("--top_k", type=int, default=50, help="Top-k sampling")
|
||||
parser.add_argument(
|
||||
"--num_samples",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of samples per problem (best-of-n scoring)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size", type=int, default=1, help="Inference batch size"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--limit",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Limit to first N problems (for quick testing)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dump_responses",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Path to dump raw model responses (JSONL)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
download_ifeval(args.data_path)
|
||||
problems = load_problems(args.data_path)
|
||||
if args.limit:
|
||||
problems = problems[: args.limit]
|
||||
|
||||
print(f"Loaded {len(problems)} problems")
|
||||
print(f"Supported constraint types: {len(CONSTRAINT_VERIFIERS)}")
|
||||
|
||||
model = AutoModel.from_pretrained(args.param_path)
|
||||
tokenizer = AutoTokenizer.from_pretrained(args.param_path)
|
||||
model.to(device="cuda", dtype=torch.bfloat16)
|
||||
model.eval()
|
||||
|
||||
engine = InferenceEngine(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
max_batch_size=args.batch_size,
|
||||
)
|
||||
|
||||
results = evaluate(
|
||||
engine=engine,
|
||||
tokenizer=tokenizer,
|
||||
problems=problems,
|
||||
max_tokens=args.max_tokens,
|
||||
temperature=args.temperature,
|
||||
top_p=args.top_p,
|
||||
top_k=args.top_k,
|
||||
num_samples=args.num_samples,
|
||||
)
|
||||
|
||||
summary = results.pop("_summary")
|
||||
print(f"\n{'=' * 60}")
|
||||
print(f" Problems: {summary['total_problems']}")
|
||||
print(f" Constraints: {summary['total_constraints']}")
|
||||
print(f" Passed: {summary['total_passed']}")
|
||||
print(f" Accuracy: {summary['overall_accuracy']:.2%}")
|
||||
print(f" Unsupported: {summary['unsupported_constraints']}")
|
||||
print(f"{'=' * 60}")
|
||||
|
||||
print("\nPer-type accuracy:")
|
||||
for inst_id, stats in sorted(summary["per_type_accuracy"].items()):
|
||||
print(
|
||||
f" {inst_id:50s} {stats['accuracy']:.2%} "
|
||||
f"({stats['passed']}/{stats['total']})"
|
||||
)
|
||||
|
||||
if args.output:
|
||||
results["_summary"] = summary
|
||||
with open(args.output, "w", encoding="utf-8") as f:
|
||||
json.dump(results, f, indent=2, ensure_ascii=False)
|
||||
print(f"\nResults saved to {args.output}")
|
||||
|
||||
if args.dump_responses:
|
||||
with open(args.dump_responses, "w", encoding="utf-8") as f:
|
||||
for k, v in results.items():
|
||||
if k.startswith("_"):
|
||||
continue
|
||||
f.write(
|
||||
json.dumps(
|
||||
{
|
||||
"key": v["key"],
|
||||
"prompt": v["prompt"],
|
||||
"response": v["response"],
|
||||
},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
+ "\n"
|
||||
)
|
||||
print(f"Responses dumped to {args.dump_responses}")
|
||||
|
||||
engine.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -4,18 +4,18 @@ import argparse
|
||||
import csv
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
import urllib.request
|
||||
import zipfile
|
||||
import random
|
||||
from collections import defaultdict
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import tqdm
|
||||
from datasets import load_dataset
|
||||
|
||||
from astrai.model import AutoModel
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
MMLU_URL = "https://github.com/hendrycks/test/archive/refs/heads/master.zip"
|
||||
MMLU_HF_DATASET = "cais/mmlu"
|
||||
MMLU_SUBJECTS = [
|
||||
"abstract_algebra",
|
||||
"anatomy",
|
||||
@@ -77,24 +77,40 @@ MMLU_SUBJECTS = [
|
||||
]
|
||||
|
||||
|
||||
def _download_and_extract(url: str, data_dir: str):
|
||||
zip_path = os.path.join(data_dir, "mmlu.zip")
|
||||
os.makedirs(data_dir, exist_ok=True)
|
||||
print(f"Downloading MMLU data from {url}...")
|
||||
urllib.request.urlretrieve(url, zip_path)
|
||||
print("Extracting...")
|
||||
with zipfile.ZipFile(zip_path, "r") as zf:
|
||||
zf.extractall(data_dir)
|
||||
os.remove(zip_path)
|
||||
def _write_subject_csv(data_dir: str, split: str, subject: str, rows: list[dict]):
|
||||
split_dir = os.path.join(data_dir, split)
|
||||
os.makedirs(split_dir, exist_ok=True)
|
||||
path = os.path.join(split_dir, f"{subject}_{split}.csv")
|
||||
with open(path, "w", encoding="utf-8", newline="") as f:
|
||||
writer = csv.writer(f)
|
||||
for row in rows:
|
||||
writer.writerow(row)
|
||||
|
||||
|
||||
def download_mmlu(data_dir: str):
|
||||
_download_and_extract(MMLU_URL, data_dir)
|
||||
src = os.path.join(data_dir, "test-master", "data")
|
||||
if os.path.exists(src):
|
||||
for item in os.listdir(src):
|
||||
os.rename(os.path.join(src, item), os.path.join(data_dir, item))
|
||||
shutil.rmtree(os.path.join(data_dir, "test-master"))
|
||||
print(f"Downloading MMLU from HuggingFace ({MMLU_HF_DATASET}) ...")
|
||||
letters = ("A", "B", "C", "D")
|
||||
split_map = {"dev": "dev", "val": "validation", "test": "test"}
|
||||
for local_split, hf_split in split_map.items():
|
||||
ds = load_dataset(MMLU_HF_DATASET, "all", split=hf_split)
|
||||
grouped: dict[str, list[dict]] = defaultdict(list)
|
||||
for item in tqdm.tqdm(ds, desc=f" {local_split}", leave=False):
|
||||
subject = item["subject"]
|
||||
choices = item["choices"]
|
||||
ans_letter = letters[item["answer"]]
|
||||
grouped[subject].append(
|
||||
[
|
||||
item["question"],
|
||||
f"A){choices[0]}",
|
||||
f"B){choices[1]}",
|
||||
f"C){choices[2]}",
|
||||
f"D){choices[3]}",
|
||||
ans_letter,
|
||||
]
|
||||
)
|
||||
for subject, rows in grouped.items():
|
||||
_write_subject_csv(data_dir, local_split, subject, rows)
|
||||
print(f" {local_split}: {len(ds)} items, {len(grouped)} subjects")
|
||||
print(f"MMLU data saved to {data_dir}")
|
||||
|
||||
|
||||
@@ -125,17 +141,12 @@ def load_csv(path: str) -> list[dict]:
|
||||
return data
|
||||
|
||||
|
||||
def build_prompt(
|
||||
question: str, choices: dict, subject: str, n_shot: int, dev_data: list[dict]
|
||||
) -> str:
|
||||
prompt = ""
|
||||
if n_shot > 0 and dev_data:
|
||||
prompt = f"The following are multiple choice questions (with answers) about {subject}.\n\n"
|
||||
for item in dev_data[:n_shot]:
|
||||
prompt += f"Question: {item['question']}\n"
|
||||
for k in ("A", "B", "C", "D"):
|
||||
prompt += f"{k}. {item[k]}\n"
|
||||
prompt += f"Answer: {item['answer']}\n\n"
|
||||
def build_prompt(question: str, choices: dict, subject: str) -> str:
|
||||
"""Build the raw question prompt (without few-shot examples).
|
||||
|
||||
Few-shot examples are handled by ``apply_chat`` to avoid duplication.
|
||||
"""
|
||||
prompt = f"The following are multiple choice questions (with answers) about {subject}.\n\n"
|
||||
prompt += f"Question: {question}\n"
|
||||
for k in ("A", "B", "C", "D"):
|
||||
prompt += f"{k}. {choices[k]}\n"
|
||||
@@ -143,10 +154,35 @@ def build_prompt(
|
||||
return prompt
|
||||
|
||||
|
||||
def apply_chat(
|
||||
tokenizer,
|
||||
raw_prompt: str,
|
||||
n_shot: int,
|
||||
dev_data: list[dict] | None,
|
||||
subject: str = "",
|
||||
) -> str:
|
||||
"""Wrap raw MMLU prompt in the model's chat template format.
|
||||
|
||||
For few-shot, prepend example Q&A pairs as user/assistant exchanges.
|
||||
Few-shot examples use the same subject preamble as the test question to
|
||||
keep the format consistent.
|
||||
"""
|
||||
messages = []
|
||||
if n_shot > 0 and dev_data:
|
||||
for item in dev_data[:n_shot]:
|
||||
q = build_prompt(item["question"], item, subject)
|
||||
messages.append({"role": "user", "content": q})
|
||||
messages.append({"role": "assistant", "content": item["answer"]})
|
||||
messages.append({"role": "user", "content": raw_prompt})
|
||||
return tokenizer.apply_chat_template(
|
||||
messages, tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
|
||||
|
||||
def choice_logprob(
|
||||
model, tokenizer, context_ids: list[int], choice_letter: str, device: str
|
||||
) -> float:
|
||||
choice_text = f" {choice_letter}"
|
||||
choice_text = choice_letter
|
||||
choice_ids = tokenizer.encode(choice_text, add_special_tokens=False)
|
||||
input_ids = context_ids + choice_ids
|
||||
max_len = model.config.max_len
|
||||
@@ -170,6 +206,25 @@ def choice_logprob(
|
||||
return score
|
||||
|
||||
|
||||
def _permute_choices(item: dict, rng: random.Random) -> tuple[dict, str]:
|
||||
"""Shuffle the option order of a question.
|
||||
|
||||
Returns ``(permuted_item, new_answer_letter)``. The question text and
|
||||
the *content* of each choice are unchanged; only which letter (A/B/C/D)
|
||||
maps to which content is shuffled. This neutralises the model's
|
||||
positional bias (e.g. always picking B).
|
||||
"""
|
||||
letters = ("A", "B", "C", "D")
|
||||
contents = [item[k] for k in letters]
|
||||
perm = list(letters)
|
||||
rng.shuffle(perm)
|
||||
permuted = {"question": item["question"]}
|
||||
for new_letter, orig_letter in zip(letters, perm):
|
||||
permuted[new_letter] = item[orig_letter]
|
||||
new_answer = letters[perm.index(item["answer"])]
|
||||
return permuted, new_answer
|
||||
|
||||
|
||||
def evaluate_subject(
|
||||
model,
|
||||
tokenizer,
|
||||
@@ -178,17 +233,24 @@ def evaluate_subject(
|
||||
dev_data: list[dict] | None,
|
||||
device: str,
|
||||
n_shot: int,
|
||||
seed: int = 0,
|
||||
) -> tuple[float, int, int]:
|
||||
rng = random.Random(seed) if seed >= 0 else None
|
||||
correct = 0
|
||||
total = 0
|
||||
for item in tqdm.tqdm(test_data, desc=f"{subject:40s}", leave=False):
|
||||
prompt = build_prompt(item["question"], item, subject, n_shot, dev_data or [])
|
||||
context_ids = tokenizer.encode(prompt)
|
||||
if rng is not None:
|
||||
permuted, answer = _permute_choices(item, rng)
|
||||
else:
|
||||
permuted, answer = item, item["answer"]
|
||||
raw_prompt = build_prompt(permuted["question"], permuted, subject)
|
||||
context = apply_chat(tokenizer, raw_prompt, n_shot, dev_data or [], subject)
|
||||
context_ids = tokenizer.encode(context)
|
||||
scores = {
|
||||
c: choice_logprob(model, tokenizer, context_ids, c, device)
|
||||
for c in ("A", "B", "C", "D")
|
||||
}
|
||||
if max(scores, key=scores.get) == item["answer"]:
|
||||
if max(scores, key=scores.get) == answer:
|
||||
correct += 1
|
||||
total += 1
|
||||
return correct / total, correct, total
|
||||
@@ -223,6 +285,12 @@ def main():
|
||||
default="bfloat16" if torch.cuda.is_available() else "float32",
|
||||
help="Torch dtype",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--seed",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Seed for option permutation (0 to enable, -1 to disable)",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.download or not os.path.exists(args.data_dir):
|
||||
@@ -233,6 +301,7 @@ def main():
|
||||
device = args.device
|
||||
dtype = getattr(torch, args.dtype)
|
||||
model.to(device=device, dtype=dtype)
|
||||
model.eval()
|
||||
|
||||
subjects = args.subjects or MMLU_SUBJECTS
|
||||
results = {}
|
||||
@@ -253,7 +322,14 @@ def main():
|
||||
test_data = load_csv(test_path)
|
||||
|
||||
acc, corr, tot = evaluate_subject(
|
||||
model, tokenizer, subject, test_data, dev_data, device, args.n_shot
|
||||
model,
|
||||
tokenizer,
|
||||
subject,
|
||||
test_data,
|
||||
dev_data,
|
||||
device,
|
||||
args.n_shot,
|
||||
seed=args.seed,
|
||||
)
|
||||
results[subject] = {"accuracy": round(acc, 4), "correct": corr, "total": tot}
|
||||
total_correct += corr
|
||||
@@ -0,0 +1,464 @@
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import statistics
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import tqdm
|
||||
|
||||
from astrai.model import AutoModel
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
|
||||
def _collect_input_files(input_path: str) -> List[str]:
|
||||
"""Resolve *input_path* to a list of JSONL/JSON files."""
|
||||
if os.path.isdir(input_path):
|
||||
files = []
|
||||
for ext in ("*.jsonl", "*.json"):
|
||||
files.extend(
|
||||
sorted(glob.glob(os.path.join(input_path, "**", ext), recursive=True))
|
||||
)
|
||||
return files
|
||||
return sorted(glob.glob(input_path))
|
||||
|
||||
|
||||
def _load_items(filepath: str) -> List[dict]:
|
||||
"""Load JSONL or JSON (array / single dict) into a list of dicts."""
|
||||
with open(filepath, "r", encoding="utf-8") as f:
|
||||
if filepath.lower().endswith(".json"):
|
||||
data = json.load(f)
|
||||
if isinstance(data, dict):
|
||||
return [data]
|
||||
return data
|
||||
return [json.loads(line) for line in f if line.strip()]
|
||||
|
||||
|
||||
def _encode_batch(
|
||||
tokenizer: AutoTokenizer, texts: List[str], max_length: int
|
||||
) -> Tuple[List[List[int]], List[List[int]]]:
|
||||
"""Encode *texts* and return (token_ids, attention_masks).
|
||||
|
||||
Each sequence is left-aligned and padded to the batch max length.
|
||||
"""
|
||||
encoded = [tokenizer.encode(t)[:max_length] for t in texts]
|
||||
if not encoded:
|
||||
return [], []
|
||||
max_len = max(len(seq) for seq in encoded)
|
||||
padded_ids = []
|
||||
masks = []
|
||||
for seq in encoded:
|
||||
pad_len = max_len - len(seq)
|
||||
padded_ids.append(seq + [tokenizer.pad_id] * pad_len)
|
||||
masks.append([1] * len(seq) + [0] * pad_len)
|
||||
return padded_ids, masks
|
||||
|
||||
|
||||
def _compute_batch(
|
||||
model,
|
||||
input_ids: torch.Tensor,
|
||||
attention_mask: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Forward pass and return (log_probs, valid_mask) of shape [B, S-1].
|
||||
|
||||
log_probs[i, j] = log P(token j+1 | tokens 0..j)
|
||||
"""
|
||||
output = model(input_ids, input_mask=attention_mask)
|
||||
logits = output["logits"][:, :-1, :] # [B, S-1, V]
|
||||
targets = input_ids[:, 1:] # [B, S-1]
|
||||
valid = attention_mask[:, 1:].float() # [B, S-1]
|
||||
|
||||
log_probs = F.log_softmax(logits.float(), dim=-1) # [B, S-1, V]
|
||||
token_log_probs = log_probs.gather(2, targets.unsqueeze(-1)).squeeze(-1) # [B, S-1]
|
||||
|
||||
return token_log_probs, valid
|
||||
|
||||
|
||||
def _token_type(token_id: int, stop_ids: frozenset, decode_fn) -> str:
|
||||
"""Classify a token into a coarse type for analysis.
|
||||
|
||||
*stop_ids* is a pre-built set of special token IDs.
|
||||
*decode_fn* is ``tokenizer.decode`` (or a wrapper) for single-token
|
||||
decoding.
|
||||
"""
|
||||
if token_id in stop_ids:
|
||||
return "special"
|
||||
decoded = decode_fn([token_id], skip_special_tokens=True)
|
||||
if any("\u4e00" <= ch <= "\u9fff" for ch in decoded):
|
||||
return "cjk"
|
||||
if any(ord(ch) > 127 for ch in decoded):
|
||||
return "non_ascii"
|
||||
return "ascii"
|
||||
|
||||
|
||||
def _percentiles(values: List[float]) -> Dict[str, float]:
|
||||
"""Compute common percentiles from a list of floats.
|
||||
|
||||
Uses linear interpolation between closest ranks (same convention
|
||||
as NumPy's default).
|
||||
"""
|
||||
if not values:
|
||||
return {}
|
||||
sorted_vals = sorted(values)
|
||||
n = len(sorted_vals)
|
||||
|
||||
def _pct(p: float) -> float:
|
||||
if n == 1:
|
||||
return sorted_vals[0]
|
||||
k = p * (n - 1)
|
||||
f = int(k)
|
||||
c = min(f + 1, n - 1)
|
||||
return sorted_vals[f] + (sorted_vals[c] - sorted_vals[f]) * (k - f)
|
||||
|
||||
return {
|
||||
"p50": _pct(0.50),
|
||||
"p90": _pct(0.90),
|
||||
"p95": _pct(0.95),
|
||||
"p99": _pct(0.99),
|
||||
}
|
||||
|
||||
|
||||
class LossAccumulator:
|
||||
"""Accumulate per-token losses with optional streaming mode.
|
||||
|
||||
When *stream* is True (token_level=False), losses are not kept
|
||||
in memory individually — only a running sum/count and a histogram
|
||||
(for approximate percentiles) are maintained. When *stream* is
|
||||
False, all losses are retained for exact statistics and per-record
|
||||
output.
|
||||
"""
|
||||
|
||||
_HIST_BINS = 1000
|
||||
_HIST_MAX = 20.0 # clamp losses above this for histogram
|
||||
|
||||
def __init__(self, stream: bool):
|
||||
self.stream = stream
|
||||
self.losses: List[float] = [] if not stream else []
|
||||
self.total: float = 0.0
|
||||
self.count: int = 0
|
||||
self.hist = torch.zeros(self._HIST_BINS, dtype=torch.long)
|
||||
# per-type losses (only populated when not streaming)
|
||||
self.by_type: Dict[str, List[float]] = {}
|
||||
self.type_total: Dict[str, float] = {}
|
||||
self.type_count: Dict[str, int] = {}
|
||||
|
||||
def add(self, losses: List[float]):
|
||||
self.total += sum(losses)
|
||||
self.count += len(losses)
|
||||
if self.stream:
|
||||
clamped = [min(max(l, 0.0), self._HIST_MAX) for l in losses]
|
||||
idx = torch.tensor(clamped) / self._HIST_MAX * (self._HIST_BINS - 1)
|
||||
self.hist += torch.bincount(
|
||||
idx.long().clamp(0, self._HIST_BINS - 1),
|
||||
minlength=self._HIST_BINS,
|
||||
)
|
||||
else:
|
||||
self.losses.extend(losses)
|
||||
|
||||
def add_typed(self, ttype: str, losses: List[float]):
|
||||
if not self.stream:
|
||||
self.by_type.setdefault(ttype, []).extend(losses)
|
||||
self.type_total[ttype] = self.type_total.get(ttype, 0.0) + sum(losses)
|
||||
self.type_count[ttype] = self.type_count.get(ttype, 0) + len(losses)
|
||||
|
||||
def stats(self) -> Dict:
|
||||
result: Dict = {}
|
||||
if self.count == 0:
|
||||
return result
|
||||
mean_loss = self.total / self.count
|
||||
result["overall"] = {
|
||||
"num_tokens": self.count,
|
||||
"mean_loss": mean_loss,
|
||||
"ppl": float(torch.exp(torch.tensor(mean_loss))),
|
||||
}
|
||||
if self.stream:
|
||||
result["overall"].update(self._hist_percentiles())
|
||||
else:
|
||||
result["overall"]["median_loss"] = statistics.median(self.losses)
|
||||
result["overall"].update(_percentiles(self.losses))
|
||||
|
||||
if self.type_count:
|
||||
result["by_token_type"] = {}
|
||||
for ttype in sorted(self.type_count.keys()):
|
||||
cnt = self.type_count[ttype]
|
||||
tmean = self.type_total[ttype] / cnt
|
||||
entry: Dict = {
|
||||
"num_tokens": cnt,
|
||||
"mean_loss": tmean,
|
||||
"ppl": float(torch.exp(torch.tensor(tmean))),
|
||||
}
|
||||
if not self.stream and ttype in self.by_type:
|
||||
entry["median_loss"] = statistics.median(self.by_type[ttype])
|
||||
entry.update(_percentiles(self.by_type[ttype]))
|
||||
result["by_token_type"][ttype] = entry
|
||||
return result
|
||||
|
||||
def _hist_percentiles(self) -> Dict[str, float]:
|
||||
"""Approximate percentiles from the histogram."""
|
||||
total = self.hist.sum().item()
|
||||
if total == 0:
|
||||
return {}
|
||||
cum = torch.cumsum(self.hist.float(), dim=0)
|
||||
result = {}
|
||||
for label, p in [("p50", 0.5), ("p90", 0.9), ("p95", 0.95), ("p99", 0.99)]:
|
||||
target = p * total
|
||||
idx = int(torch.searchsorted(cum, target).item())
|
||||
idx = min(idx, self._HIST_BINS - 1)
|
||||
result[label] = (idx + 0.5) / self._HIST_BINS * self._HIST_MAX
|
||||
return result
|
||||
|
||||
|
||||
def process_file(
|
||||
model,
|
||||
tokenizer: AutoTokenizer,
|
||||
items: List[dict],
|
||||
text_key: str,
|
||||
batch_size: int,
|
||||
max_length: int,
|
||||
token_level: bool,
|
||||
max_samples: Optional[int],
|
||||
output_file: Optional[str],
|
||||
label: str,
|
||||
device: str = "cuda",
|
||||
) -> Dict:
|
||||
"""Evaluate a single dataset (list of items), return summary stats.
|
||||
|
||||
If *token_level* is True and *output_file* is set, per-record token_ids
|
||||
and log_probs are written as JSONL alongside the summary.
|
||||
"""
|
||||
if max_samples and len(items) > max_samples:
|
||||
import random
|
||||
|
||||
items = random.sample(items, max_samples)
|
||||
|
||||
texts = [item[text_key] for item in items if text_key in item]
|
||||
print(f" [{label}] {len(texts)} samples, text_key='{text_key}'")
|
||||
|
||||
acc = LossAccumulator(stream=not token_level)
|
||||
per_sample: List[dict] = []
|
||||
|
||||
if token_level:
|
||||
stop_ids = frozenset(tokenizer.stop_ids)
|
||||
decode_fn = tokenizer.decode
|
||||
|
||||
num_batches = (len(texts) + batch_size - 1) // batch_size
|
||||
for i in tqdm.tqdm(
|
||||
range(0, len(texts), batch_size),
|
||||
total=num_batches,
|
||||
desc=f" {label}",
|
||||
leave=False,
|
||||
):
|
||||
batch_texts = texts[i : i + batch_size]
|
||||
padded_ids, masks = _encode_batch(tokenizer, batch_texts, max_length)
|
||||
|
||||
input_ids = torch.tensor(padded_ids, device=device, dtype=torch.long)
|
||||
attention_mask = torch.tensor(masks, device=device, dtype=torch.bool)
|
||||
|
||||
token_log_probs, valid = _compute_batch(model, input_ids, attention_mask)
|
||||
|
||||
for b in range(len(batch_texts)):
|
||||
seq_len = int(valid[b].sum().item())
|
||||
lps = token_log_probs[b, :seq_len].tolist()
|
||||
losses = [-lp for lp in lps]
|
||||
acc.add(losses)
|
||||
|
||||
if token_level:
|
||||
# log_probs correspond to positions 1..seq_len (predicted
|
||||
# from position 0..seq_len-1), so token_ids must skip BOS
|
||||
# at position 0 to stay aligned with log_probs.
|
||||
ids = padded_ids[b][1 : seq_len + 1]
|
||||
per_sample.append(
|
||||
{
|
||||
"text": batch_texts[b][:200],
|
||||
"token_ids": ids,
|
||||
"log_probs": [round(lp, 4) for lp in lps],
|
||||
"ppl": float(torch.exp(torch.tensor(statistics.mean(losses))))
|
||||
if losses
|
||||
else None,
|
||||
}
|
||||
)
|
||||
typed_losses: Dict[str, List[float]] = {}
|
||||
for tid, loss in zip(ids, losses):
|
||||
ttype = _token_type(tid, stop_ids, decode_fn)
|
||||
typed_losses.setdefault(ttype, []).append(loss)
|
||||
for ttype, tl in typed_losses.items():
|
||||
acc.add_typed(ttype, tl)
|
||||
|
||||
stats = acc.stats()
|
||||
|
||||
if token_level and output_file:
|
||||
with open(output_file, "w", encoding="utf-8") as f:
|
||||
for item in per_sample:
|
||||
f.write(json.dumps(item, ensure_ascii=False) + "\n")
|
||||
|
||||
return stats
|
||||
|
||||
|
||||
def print_stats(label: str, stats: Dict):
|
||||
"""Pretty-print summary statistics."""
|
||||
print(f"\n{'=' * 60}")
|
||||
print(f" {label}")
|
||||
print(f"{'=' * 60}")
|
||||
ov = stats.get("overall", {})
|
||||
if ov:
|
||||
print(f" tokens: {ov['num_tokens']:,}")
|
||||
print(f" mean loss: {ov['mean_loss']:.4f}")
|
||||
if "median_loss" in ov:
|
||||
print(f" median loss: {ov['median_loss']:.4f}")
|
||||
print(f" ppl: {ov['ppl']:.2f}")
|
||||
if "p50" in ov:
|
||||
print(
|
||||
f" p50/p90/p95/p99: "
|
||||
f"{ov['p50']:.2f} / {ov['p90']:.2f} / {ov['p95']:.2f} / {ov['p99']:.2f}"
|
||||
)
|
||||
by_type = stats.get("by_token_type", {})
|
||||
if by_type:
|
||||
print(f"\n by token type:")
|
||||
print(f" {'type':<12} {'count':>8} {'mean_loss':>10} {'ppl':>8}")
|
||||
print(f" {'-' * 12} {'-' * 8} {'-' * 10} {'-' * 8}")
|
||||
for ttype, s in by_type.items():
|
||||
print(
|
||||
f" {ttype:<12} {s['num_tokens']:>8,} "
|
||||
f"{s['mean_loss']:>10.4f} {s['ppl']:>8.2f}"
|
||||
)
|
||||
|
||||
|
||||
def main(
|
||||
param_path: str,
|
||||
input_path: str,
|
||||
output_dir: str,
|
||||
text_key: str,
|
||||
batch_size: int,
|
||||
max_length: int,
|
||||
token_level: bool,
|
||||
max_samples: Optional[int],
|
||||
device: str = "cuda",
|
||||
dtype: str = "bfloat16",
|
||||
):
|
||||
print(f"Loading model from {param_path} ...")
|
||||
model = AutoModel.from_pretrained(param_path)
|
||||
tokenizer = AutoTokenizer.from_pretrained(param_path)
|
||||
torch_dtype = getattr(torch, dtype)
|
||||
model.to(device=device, dtype=torch_dtype)
|
||||
model.eval()
|
||||
|
||||
input_files = _collect_input_files(input_path)
|
||||
if not input_files:
|
||||
print(f"No input files found at {input_path}")
|
||||
return
|
||||
|
||||
print(f"Found {len(input_files)} file(s) to evaluate")
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
all_stats = {}
|
||||
for filepath in input_files:
|
||||
label = os.path.splitext(os.path.basename(filepath))[0]
|
||||
items = _load_items(filepath)
|
||||
if not items:
|
||||
print(f" [{label}] empty, skipping")
|
||||
continue
|
||||
|
||||
token_output = (
|
||||
os.path.join(output_dir, f"{label}_tokens.jsonl") if token_level else None
|
||||
)
|
||||
|
||||
stats = process_file(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
items=items,
|
||||
text_key=text_key,
|
||||
batch_size=batch_size,
|
||||
max_length=max_length,
|
||||
token_level=token_level,
|
||||
max_samples=max_samples,
|
||||
output_file=token_output,
|
||||
label=label,
|
||||
device=device,
|
||||
)
|
||||
all_stats[label] = stats
|
||||
print_stats(label, stats)
|
||||
|
||||
if token_output:
|
||||
print(f" token-level output: {token_output}")
|
||||
|
||||
summary_path = os.path.join(output_dir, "summary.json")
|
||||
with open(summary_path, "w", encoding="utf-8") as f:
|
||||
json.dump(all_stats, f, ensure_ascii=False, indent=2)
|
||||
print(f"\nSummary saved to {summary_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Perplexity and token-level loss evaluation on JSONL/JSON data."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--param_path", type=str, required=True, help="Path to the model directory."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--input_path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to input file, glob pattern, or directory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Directory for output files (summary.json + per-file token JSONL).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text_key",
|
||||
type=str,
|
||||
default="text",
|
||||
help="Key for the text field in the input data.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size", type=int, default=4, help="Batch size for evaluation."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_length",
|
||||
type=int,
|
||||
default=2048,
|
||||
help="Maximum sequence length (tokens). Longer sequences are truncated.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--token_level",
|
||||
action="store_true",
|
||||
help="Store per-token log_probs and token type analysis. "
|
||||
"Default: off (only aggregate stats).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_samples",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Maximum number of samples per file (random subsample). Default: all.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--device",
|
||||
type=str,
|
||||
default="cuda" if torch.cuda.is_available() else "cpu",
|
||||
help="Device for model inference.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dtype",
|
||||
type=str,
|
||||
default="bfloat16" if torch.cuda.is_available() else "float32",
|
||||
help="Torch dtype for model weights.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
with torch.inference_mode():
|
||||
main(
|
||||
param_path=args.param_path,
|
||||
input_path=args.input_path,
|
||||
output_dir=args.output_dir,
|
||||
text_key=args.text_key,
|
||||
batch_size=args.batch_size,
|
||||
max_length=args.max_length,
|
||||
token_level=args.token_level,
|
||||
max_samples=args.max_samples,
|
||||
device=args.device,
|
||||
dtype=args.dtype,
|
||||
)
|
||||
@@ -0,0 +1,153 @@
|
||||
"""ROUGE evaluation (manual implementation, no external deps).
|
||||
|
||||
Computes ROUGE-1, ROUGE-2, ROUGE-L precision, recall, and F1.
|
||||
|
||||
Usage::
|
||||
|
||||
# Batch evaluation from JSONL (each line: {"reference": ..., "candidate": ...})
|
||||
python scripts/eval/evaluate_rouge.py --data_path preds.jsonl --output results.json
|
||||
|
||||
# As a library
|
||||
from scripts.eval.evaluate_rouge import compute_rouge
|
||||
scores = compute_rouge("the cat sat on the mat", "the cat sat")
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from collections import Counter
|
||||
from typing import Dict, List, Tuple
|
||||
|
||||
|
||||
def _tokenize(text: str) -> List[str]:
|
||||
return text.split()
|
||||
|
||||
|
||||
def _ngrams(tokens: List[str], n: int) -> Counter:
|
||||
return Counter(zip(*[tokens[i:] for i in range(n)]))
|
||||
|
||||
|
||||
def _lcs(x: List[str], y: List[str]) -> int:
|
||||
m, n = len(x), len(y)
|
||||
dp = [[0] * (n + 1) for _ in range(m + 1)]
|
||||
for i in range(1, m + 1):
|
||||
xi = x[i - 1]
|
||||
dpi = dp[i]
|
||||
dpi_1 = dp[i - 1]
|
||||
for j in range(1, n + 1):
|
||||
if xi == y[j - 1]:
|
||||
dpi[j] = dpi_1[j - 1] + 1
|
||||
else:
|
||||
dpi[j] = dpi_1[j] if dpi_1[j] > dpi[j - 1] else dpi[j - 1]
|
||||
return dp[m][n]
|
||||
|
||||
|
||||
def _f1(precision: float, recall: float) -> float:
|
||||
if precision + recall == 0:
|
||||
return 0.0
|
||||
return 2 * precision * recall / (precision + recall)
|
||||
|
||||
|
||||
def _rouge_n(ref_tokens: List[str], cand_tokens: List[str], n: int) -> Dict[str, float]:
|
||||
ref_ngrams = _ngrams(ref_tokens, n)
|
||||
cand_ngrams = _ngrams(cand_tokens, n)
|
||||
|
||||
overlap = sum((cand_ngrams & ref_ngrams).values())
|
||||
cand_total = sum(cand_ngrams.values())
|
||||
ref_total = sum(ref_ngrams.values())
|
||||
|
||||
precision = overlap / cand_total if cand_total > 0 else 0.0
|
||||
recall = overlap / ref_total if ref_total > 0 else 0.0
|
||||
f1 = _f1(precision, recall)
|
||||
|
||||
return {"precision": precision, "recall": recall, "f1": f1}
|
||||
|
||||
|
||||
def _rouge_l(ref_tokens: List[str], cand_tokens: List[str]) -> Dict[str, float]:
|
||||
lcs_len = _lcs(ref_tokens, cand_tokens)
|
||||
ref_len = len(ref_tokens)
|
||||
cand_len = len(cand_tokens)
|
||||
|
||||
recall = lcs_len / ref_len if ref_len > 0 else 0.0
|
||||
precision = lcs_len / cand_len if cand_len > 0 else 0.0
|
||||
f1 = _f1(precision, recall)
|
||||
|
||||
return {"precision": precision, "recall": recall, "f1": f1}
|
||||
|
||||
|
||||
def compute_rouge(
|
||||
reference: str, candidate: str, n: int = 2
|
||||
) -> Dict[str, Dict[str, float]]:
|
||||
"""Compute ROUGE-N (1..n) and ROUGE-L scores.
|
||||
|
||||
Returns::
|
||||
|
||||
{
|
||||
"rouge-1": {"precision": ..., "recall": ..., "f1": ...},
|
||||
"rouge-2": {"precision": ..., "recall": ..., "f1": ...},
|
||||
"rouge-l": {"precision": ..., "recall": ..., "f1": ...},
|
||||
}
|
||||
"""
|
||||
ref_tokens = _tokenize(reference)
|
||||
cand_tokens = _tokenize(candidate)
|
||||
|
||||
results = {}
|
||||
for i in range(1, n + 1):
|
||||
results[f"rouge-{i}"] = _rouge_n(ref_tokens, cand_tokens, i)
|
||||
results["rouge-l"] = _rouge_l(ref_tokens, cand_tokens)
|
||||
return results
|
||||
|
||||
|
||||
def evaluate_file(data_path: str) -> Dict:
|
||||
with open(data_path, "r", encoding="utf-8") as f:
|
||||
pairs = [json.loads(line) for line in f if line.strip()]
|
||||
|
||||
agg = {
|
||||
k: {"precision": 0.0, "recall": 0.0, "f1": 0.0}
|
||||
for k in ("rouge-1", "rouge-2", "rouge-l")
|
||||
}
|
||||
per_item = []
|
||||
|
||||
for item in pairs:
|
||||
ref = item["reference"]
|
||||
cand = item["candidate"]
|
||||
scores = compute_rouge(ref, cand)
|
||||
per_item.append({**item, "scores": scores})
|
||||
for k, v in scores.items():
|
||||
agg[k]["precision"] += v["precision"]
|
||||
agg[k]["recall"] += v["recall"]
|
||||
agg[k]["f1"] += v["f1"]
|
||||
|
||||
n = len(pairs)
|
||||
for k in agg:
|
||||
agg[k] = {m: v / n for m, v in agg[k].items()}
|
||||
|
||||
return {"num_samples": n, "aggregate": agg, "per_item": per_item}
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="ROUGE evaluation")
|
||||
parser.add_argument(
|
||||
"--data_path", required=True, help="JSONL with reference/candidate per line"
|
||||
)
|
||||
parser.add_argument("--output", type=str, default=None, help="Output JSON path")
|
||||
args = parser.parse_args()
|
||||
|
||||
results = evaluate_file(args.data_path)
|
||||
agg = results["aggregate"]
|
||||
|
||||
print(f"Samples: {results['num_samples']}")
|
||||
print()
|
||||
for metric in ("rouge-1", "rouge-2", "rouge-l"):
|
||||
s = agg[metric]
|
||||
print(
|
||||
f" {metric:8s} P={s['precision']:.4f} R={s['recall']:.4f} F1={s['f1']:.4f}"
|
||||
)
|
||||
|
||||
if args.output:
|
||||
with open(args.output, "w", encoding="utf-8") as f:
|
||||
json.dump(results, f, indent=2, ensure_ascii=False)
|
||||
print(f"\nSaved to {args.output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
+129
-79
@@ -1,12 +1,13 @@
|
||||
"""Benchmark AutoRegressiveLM with KVCache"""
|
||||
|
||||
import argparse
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.config import AutoRegressiveLMConfig
|
||||
from astrai.inference import KVCache
|
||||
from astrai.inference import ContiguousCache, PageCache
|
||||
from astrai.model.transformer import AutoRegressiveLM
|
||||
|
||||
|
||||
@@ -24,41 +25,14 @@ class GenerationBenchmark:
|
||||
config: AutoRegressiveLMConfig,
|
||||
device: str = "cuda",
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
page_size: int = 128,
|
||||
cache_type: str = "contiguous",
|
||||
):
|
||||
self.config = config
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
self.cache_type = cache_type
|
||||
self.model = AutoRegressiveLM(config).to(device=device, dtype=dtype)
|
||||
self.model.eval()
|
||||
head_dim = config.dim // config.n_heads
|
||||
n_pages = (config.max_len * 4 + page_size - 1) // page_size
|
||||
self._page_cache = KVCache(
|
||||
config.n_layers,
|
||||
n_pages,
|
||||
page_size,
|
||||
config.n_kv_heads,
|
||||
head_dim,
|
||||
device,
|
||||
dtype,
|
||||
)
|
||||
|
||||
def _prepare_inputs(self, batch_size: int, prompt_length: int, total_length: int):
|
||||
prompt_ids = torch.randint(
|
||||
low=0,
|
||||
high=self.config.vocab_size,
|
||||
size=(batch_size, prompt_length),
|
||||
device=self.device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
gen_ids = torch.randint(
|
||||
low=0,
|
||||
high=self.config.vocab_size,
|
||||
size=(batch_size, total_length - prompt_length),
|
||||
device=self.device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
return prompt_ids, gen_ids
|
||||
|
||||
@torch.inference_mode()
|
||||
def run_prefill_benchmark(
|
||||
@@ -68,8 +42,12 @@ class GenerationBenchmark:
|
||||
num_trials: int = 10,
|
||||
) -> BenchmarkResult:
|
||||
for _ in range(3):
|
||||
prompt_ids, _ = self._prepare_inputs(
|
||||
batch_size, prompt_length, prompt_length
|
||||
prompt_ids = torch.randint(
|
||||
0,
|
||||
self.config.vocab_size,
|
||||
(batch_size, prompt_length),
|
||||
device=self.device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
_ = self.model(prompt_ids)
|
||||
torch.cuda.synchronize()
|
||||
@@ -78,12 +56,15 @@ class GenerationBenchmark:
|
||||
total_tokens = batch_size * prompt_length * num_trials
|
||||
|
||||
for trial in range(num_trials):
|
||||
prompt_ids, _ = self._prepare_inputs(
|
||||
batch_size, prompt_length, prompt_length
|
||||
prompt_ids = torch.randint(
|
||||
0,
|
||||
self.config.vocab_size,
|
||||
(batch_size, prompt_length),
|
||||
device=self.device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
start = torch.cuda.Event(enable_timing=True)
|
||||
end = torch.cuda.Event(enable_timing=True)
|
||||
|
||||
start.record()
|
||||
_ = self.model(prompt_ids)
|
||||
end.record()
|
||||
@@ -107,6 +88,7 @@ class GenerationBenchmark:
|
||||
"prompt_length": prompt_length,
|
||||
"dtype": str(self.dtype),
|
||||
"device": self.device,
|
||||
"cache": "none",
|
||||
},
|
||||
)
|
||||
|
||||
@@ -120,29 +102,56 @@ class GenerationBenchmark:
|
||||
) -> BenchmarkResult:
|
||||
total_time = 0.0
|
||||
total_tokens = batch_size * gen_length * num_trials
|
||||
page_size = self._page_cache.page_size
|
||||
|
||||
for trial in range(num_trials):
|
||||
prompt_ids, gen_ids = self._prepare_inputs(
|
||||
batch_size,
|
||||
prompt_length,
|
||||
prompt_length + gen_length,
|
||||
)
|
||||
|
||||
n_pages = (prompt_length + gen_length + page_size - 1) // page_size
|
||||
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(
|
||||
[pages[i * n_pages : (i + 1) * n_pages] for i in range(batch_size)],
|
||||
dtype=torch.long,
|
||||
prompt_ids = torch.randint(
|
||||
0,
|
||||
self.config.vocab_size,
|
||||
(batch_size, prompt_length),
|
||||
device=self.device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
gen_ids = torch.randint(
|
||||
0,
|
||||
self.config.vocab_size,
|
||||
(batch_size, gen_length),
|
||||
device=self.device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
|
||||
cv = self._page_cache.bind(page_table, total_len=prompt_length)
|
||||
head_dim = self.config.dim // self.config.n_heads
|
||||
max_seq = prompt_length + gen_length
|
||||
|
||||
if self.cache_type == "contiguous":
|
||||
cache = ContiguousCache(
|
||||
self.config.n_layers,
|
||||
batch_size,
|
||||
max_seq,
|
||||
self.config.n_kv_heads,
|
||||
head_dim,
|
||||
self.device,
|
||||
self.dtype,
|
||||
)
|
||||
else:
|
||||
page_size = 128
|
||||
n_pages = (max_seq + page_size - 1) // page_size * batch_size
|
||||
cache = PageCache(
|
||||
self.config.n_layers,
|
||||
n_pages,
|
||||
page_size,
|
||||
self.config.n_kv_heads,
|
||||
head_dim,
|
||||
self.device,
|
||||
self.dtype,
|
||||
)
|
||||
|
||||
task_ids = [f"b{i}" for i in range(batch_size)]
|
||||
for tid in task_ids:
|
||||
cache.task_alloc(tid, [0] * max_seq)
|
||||
for p in range(max_seq):
|
||||
cache.task_extend(tid, p)
|
||||
|
||||
cv = cache.bind_tasks(task_ids, prompt_length, self.device)
|
||||
_ = self.model(
|
||||
prompt_ids,
|
||||
paged_cache=cv,
|
||||
@@ -152,37 +161,35 @@ class GenerationBenchmark:
|
||||
.unsqueeze(0)
|
||||
.expand(batch_size, -1),
|
||||
)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
|
||||
start = torch.cuda.Event(enable_timing=True)
|
||||
end = torch.cuda.Event(enable_timing=True)
|
||||
|
||||
start.record()
|
||||
current_pos = prompt_length
|
||||
|
||||
for i in range(gen_length):
|
||||
input_token = gen_ids[:, i : i + 1]
|
||||
cv = self._page_cache.bind(page_table, total_len=current_pos + 1)
|
||||
pos = prompt_length + i
|
||||
cv = cache.bind_tasks(task_ids, pos + 1, self.device)
|
||||
_ = self.model(
|
||||
input_token,
|
||||
gen_ids[:, i : i + 1],
|
||||
paged_cache=cv,
|
||||
position_ids=torch.full(
|
||||
(batch_size, 1),
|
||||
current_pos,
|
||||
pos,
|
||||
dtype=torch.long,
|
||||
device=self.device,
|
||||
),
|
||||
)
|
||||
current_pos += 1
|
||||
|
||||
end.record()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
for tid in task_ids:
|
||||
cache.task_free(tid)
|
||||
|
||||
trial_time = start.elapsed_time(end) / 1000
|
||||
total_time += trial_time
|
||||
|
||||
for idx in pages:
|
||||
self._page_cache._pool.free(idx)
|
||||
|
||||
print(
|
||||
f" Trial {trial + 1}/{num_trials}: {gen_length} tokens in {trial_time:.3f}s "
|
||||
f"({gen_length / trial_time:.1f} tok/s)"
|
||||
@@ -199,6 +206,7 @@ class GenerationBenchmark:
|
||||
"gen_length": gen_length,
|
||||
"dtype": str(self.dtype),
|
||||
"device": self.device,
|
||||
"cache": self.cache_type,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -216,6 +224,42 @@ def print_benchmark_result(result: BenchmarkResult):
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="AutoRegressiveLM benchmark")
|
||||
parser.add_argument(
|
||||
"--device", type=str, default="cuda", help="Device (default: cuda)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dtype",
|
||||
type=str,
|
||||
default="bfloat16",
|
||||
choices=["bfloat16", "float16", "float32"],
|
||||
help="Dtype",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cache",
|
||||
type=str,
|
||||
default="contiguous",
|
||||
choices=["contiguous", "paged"],
|
||||
help="KV cache type",
|
||||
)
|
||||
parser.add_argument("--batch_size", type=int, default=4, help="Batch size")
|
||||
parser.add_argument("--prompt_length", type=int, default=512, help="Prompt length")
|
||||
parser.add_argument("--gen_length", type=int, default=128, help="Generation length")
|
||||
parser.add_argument("--num_trials", type=int, default=5, help="Number of trials")
|
||||
parser.add_argument(
|
||||
"--prefill_only", action="store_true", help="Run prefill benchmark only"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--decode_only", action="store_true", help="Run decoding benchmark only"
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
dtype_map = {
|
||||
"bfloat16": torch.bfloat16,
|
||||
"float16": torch.float16,
|
||||
"float32": torch.float32,
|
||||
}
|
||||
|
||||
config = AutoRegressiveLMConfig(
|
||||
vocab_size=10000,
|
||||
dim=1536,
|
||||
@@ -227,23 +271,29 @@ if __name__ == "__main__":
|
||||
norm_eps=1e-5,
|
||||
)
|
||||
|
||||
benchmark = GenerationBenchmark(config)
|
||||
benchmark = GenerationBenchmark(
|
||||
config, device=args.device, dtype=dtype_map[args.dtype], cache_type=args.cache
|
||||
)
|
||||
|
||||
print("=" * 80)
|
||||
print("Running AutoRegressiveLM Generation Benchmark (KVCache)")
|
||||
print(
|
||||
f"Running AutoRegressiveLM Benchmark (device={args.device}, dtype={args.dtype})"
|
||||
)
|
||||
print("=" * 80)
|
||||
|
||||
prefill_result = benchmark.run_prefill_benchmark(
|
||||
batch_size=4,
|
||||
prompt_length=512,
|
||||
num_trials=5,
|
||||
)
|
||||
print_benchmark_result(prefill_result)
|
||||
if not args.decode_only:
|
||||
prefill_result = benchmark.run_prefill_benchmark(
|
||||
batch_size=args.batch_size,
|
||||
prompt_length=args.prompt_length,
|
||||
num_trials=args.num_trials,
|
||||
)
|
||||
print_benchmark_result(prefill_result)
|
||||
|
||||
gen_result = benchmark.run_decoding_benchmark(
|
||||
batch_size=4,
|
||||
prompt_length=512,
|
||||
gen_length=128,
|
||||
num_trials=5,
|
||||
)
|
||||
print_benchmark_result(gen_result)
|
||||
if not args.prefill_only:
|
||||
gen_result = benchmark.run_decoding_benchmark(
|
||||
batch_size=args.batch_size,
|
||||
prompt_length=args.prompt_length,
|
||||
gen_length=args.gen_length,
|
||||
num_trials=args.num_trials,
|
||||
)
|
||||
print_benchmark_result(gen_result)
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user