24 Commits
Author SHA1 Message Date
ViperEkura 523eacf5fe release: v1.3.4
- refactor: 分页 KV cache(PagedCache+CacheView)替换固定 slot,删除 PrefixCache
- refactor: 推理引擎控制逻辑重写,修复连续批处理核心缺陷、线程安全问题
- refactor: KV 缓存槽位下沉到注意力层,移除 _remap_kv / _writeback_kv
- refactor: 统一采样路径为 SamplingPipeline batch tensor,删除 apply_sampling_strategies
- refactor: 设计模式优化 inference 模块导入结构(cache/sampling 独立)
- feat: 推理引擎前缀缓存(KV cache 复用)
- feat: OpenAI 兼容 chat completion API(流式+非流式+usage)
- feat: Anthropic 兼容 /v1/messages API,移除旧版 /generate 端点
- feat: GRPO CLI 接入 + on-policy,OpenAI API top_k 参数化
- feat: Checkpoint 支持 extra 通用扩展数据
- feat: Docker Compose 一键部署(GPU/CPU 双模式)
- feat: GRPO 训练参数补充,批处理训练参数表
- fix: 调度器延迟优化 — 移除 5ms 睡眠,修复 refill 任务丢失
- fix: CLI 参数缺失/重复、device_ids 越界、generate 参数名不一致
- fix: 长对话截断方向错误,保留最新 token 而非最早
- fix: remove_task 未释放 KV cache slot 导致第二轮对话死锁
- fix: KV cache 槽位索引错位、版本校验缺失、注意力掩码
- fix: scheduler 越界 bug,SchedulerCallback 回调阶段修正
- perf: _Result 改用 Condition.wait_for 消除非流式 CPU 空转
- perf: decode 每步张量预分配;input_ids 改用一次构建代替逐元素赋值
- refactor: 移除 device_ids 参数,统一 CUDA_VISIBLE_DEVICES
- docs: 更新文档以匹配分页 KV cache 等代码重构
- docs: 修正多处文档错误、补充训练参数说明
2026-05-10 15:59:18 +08:00
ViperEkura cffedaad5e perf: 消除非流式推理 CPU 空转并减少 decode GPU 张量冗余分配
- engine.py: _Result 改用 threading.Condition.wait_for 替代
  Event busy-wait,非流式模式线程被内核挂起而非 1760 万次空转
- scheduler.py: _execute_decode 将 temperature/top_k/top_p 张量
  移至循环外预先分配,避免每步重复 torch.tensor();input_ids
  改用 torch.empty 避免不必要的 zero 初始化(两处均为完全覆盖)
- _execute_prefill: input_ids 同改为 torch.empty
2026-05-10 15:32:11 +08:00
ViperEkura 3583c46b66 feat: 推理引擎前缀缓存(KV cache 复用)
- cache.py: 新增模块级 page_hash() 多项式滚动哈希函数;PagedCache 新增
  record_page/lookup_prefix/inc_ref,free() 自动清理哈希映射
- scheduler.py: Task 新增 _prefix_cached_tokens;_refill_active_batch 先查
  缓存命中页(inc_ref)再分配剩余页;合并 _execute_prefill 为单一方法,
  按 (prompt_len, start_pos) 分组批量执行全量/部分 prefill;
  _record_page_hashes 注册完整页哈希;修复 device/dtype 默认值从硬编码
  改为 None(自动检测模型设备)
- test: mock model 补充 dtype/device 适配自动检测
2026-05-09 23:53:57 +08:00
ViperEkura ca4e6b907c feat: Checkpoint 支持 extra 通用扩展数据,用户通过函数自定义保存/恢复优化器等状态
- serialization.py: Checkpoint 新增 extra: dict 字段,
  save() 写入 extra.pt,load() 自动恢复
- train_callback.py: CheckpointCallback 新增 save_extra_fn
  参数,用户传入 (context) -> dict 决定保存哪些额外状态
- train_context.py: TrainContextBuilder 新增 load_extra_fn
  参数,用户传入 (extra, context) 从 checkpoint 恢复状态
2026-05-09 15:50:38 +08:00
ViperEkura db99d8b254 fix: 修复文档多处不准确 + inference scheduler 越界 bug + SchedulerCallback 回调阶段修正
文档 (6 个文件):
- design.md: 15+ 处修正 — persistent_key_values→paged_cache,
  MLA 字段重写, Server/ParallelSetup 不存在类移除,
  关系箭头方向修复, SchedulerCallback 阶段修正等
- dataflow.md: 重写数据流图和描述, 修复训练回调顺序、
  数据键名、MLA 归属、MetricTracker 等错误
- introduction.md: 层数 32→24, MLP 图双 Linear 修正,
  默认值/响应字段/health 端点修复
- params.md: 补充 grpo 及 4 个 GRPO 参数
- README.md / README-zh-CN.md: generate.py 补全必需参数,
  删除重复注释, HuggingFace 声明修正

代码 (2 个文件):
- scheduler.py: n_pages 池加 page_size 余量防止越界;
  decode 前预分配页
- train_callback.py: SchedulerCallback 从 on_step_end 改
  回 on_batch_end (按 batch 步进学习率)
2026-05-09 15:40:17 +08:00
ViperEkura b98c9cefdc refactor: 移除 device_ids 参数设计,统一通过 CUDA_VISIBLE_DEVICES 控制 GPU 分配;更新 README 训练示例
- setup.py: 移除 device_ids 参数,setup_parallel 直接用 rank 作为设备索引
- train_config.py: 移除 device_ids 字段
- trainer.py: 不再传递 device_ids
- train.py: ddp_wrap 用 get_rank() 直接取值
- README.md, README-zh-CN.md: 训练示例改为多行命令风格,去掉参数表格
2026-05-09 14:55:43 +08:00
ViperEkura 283bcaf2ff fix: 修复 CLI 参数缺失/重复、device_ids 越界、generate 参数名不一致、scheduler 时序、非流式截断等 bug
- train.py: 补上 --batch_size、--grpo_clip_eps,删除 3 处重复 --group_size
- generate.py: --model_dir 改为 --param_path 对齐 README
- automodel.py: from_pretrained 新增 strict 参数(默认 True)
- parallel/setup.py: 修复 device_ids 索引越界
- train_callback.py: scheduler.step() 移至 on_step_end
- test_train_strategy.py: 测试中补 optimizer.step()
- engine.py: 非流式改为循环等待所有任务完成,补 remove_task 清理
- scheduler.py: Task 添加 _pages_freed 标志,杜绝双重释放
- trainer.py: accumulation_steps=0 时 clamp 为 1
- tokenizer.py: save_pretrained 添加 _tokenizer is None 检查
- benchmark.py: 修复 ModelConfig 过时 import 路径
- inference/__init__.py: 修复 stale docstring
2026-05-09 14:36:42 +08:00
ViperEkura bc7c82977e feat: GRPO CLI 接入 + on-policy,OpenAI API top_k 参数化,补充训练参数表
- train.py 新增 --train_type=grpo 及参数 (--grpo_clip_eps, --grpo_kl_coef, --group_size, --grpo_sync_interval, --start_epoch)
- GRPOStrategy 统一 on-policy 模式,ratio = exp(logπ_θ - logπ_ref),PPO 裁剪目标,sync_interval 自动同步 ref_model
- ChatCompletionRequest 新增 top_k 参数,不再硬编码
- 补充 README 完整训练参数表(含此前缺失的 max_grad_norm / adamw / window_size / stride 等)
2026-05-09 12:22:33 +08:00
ViperEkura 34a511e36e feat: 新增 Docker Compose 一键部署,支持 GPU/CPU 双模式 2026-05-09 11:57:46 +08:00
ViperEkura d73f52a2f8 feat: 新增 Anthropic 兼容 /v1/messages API,移除旧版 /generate 端点
- 新增 /v1/messages 端点,兼容 Anthropic Messages API 格式
- 支持流式 SSE(message_start → content_block_delta → message_stop)
- 支持 system 顶层提示词与 stop_sequences 停止序列
- 新增 AnthropicMessage / MessagesRequest Pydantic 模型
- 移除旧版 /generate 端点及相关测试用例
- 更新 README.md / README-zh-CN.md / introduction.md 文档
2026-05-09 11:47:22 +08:00
ViperEkura 9d96b0431d docs: 更新文档以匹配分页 KV cache 等代码重构 2026-05-08 22:41:13 +08:00
ViperEkura f81e2b4a73 feat: OpenAI 兼容的 chat completion API(流式+非流式+usage) 2026-05-08 21:54:55 +08:00
ViperEkura 4e324d8f26 fix: benchmark 改用 PagedCache 替代已删除的 persistent_key_values 2026-05-08 21:26:55 +08:00
ViperEkura 6ed0506491 fix: 减少调度器延迟 — 移除解码路径 5ms 睡眠,修复 refill 任务丢失 bug 2026-05-08 21:13:52 +08:00
ViperEkura 30cc2d67a4 refactor: 分页 KV cache 替换固定 slot,删除 PrefixCache 及相关死代码
- 用 PagedCache + CacheView 替换固定 slot 式 KV cache,attention 层只通过 page_table 间接索引
- 删除 PrefixCache(radix tree)及 scheduler 中所有 prefix cache 命中/插入/释放逻辑
- 删除无用函数:pin、version、free_count、_mark_seq_mask 及 seq_mask 分配
- 修复 write 在多页 prefill 时 offset 为负导致 chunk 计算错误
- _make_page_table_tensor 改用 list 拼接一次 tensor,去掉逐元素赋值
- 清理 model 接口参数:kv_cache, slot_indices → paged_cache(CacheView)
- 精简 docstring 为单行,删除冗余 section 注释和旧代码
- 修复 test_scheduler_concurrency.py 缺少 import pytest
2026-05-08 20:44:05 +08:00
ViperEkura 7ddebf2cd9 refactor: 统一采样路径为 Strategy + batch tensor,删除 apply_sampling_strategies
- TemperatureStrategy / TopKStrategy / TopPStrategy 支持 Union[float, Tensor]
- SamplingPipeline.sample() 一条调用完成 apply + softmax + multinomial
- 新增 sample() 独立函数作为 scheduler 入口
- scheduler decode 改为 batch tensor 参数传递,支持任意 batch size
- 删除 apply_sampling_strategies(被 sample() 取代)
2026-05-08 19:07:14 +08:00
ViperEkura 78dc2bd41c docs: 修正文档错误并补充训练参数说明
- README: 补充训练参数速查表,完善训练命令示例
- design.md: 同步 inference 类图(SlotAllocator、GenerationParams、采样策略等
  新增类),修正参数名和类型错误,统一泛型符号
- params.md: 修正默认值(batch_size=1、num_workers=4),移除不存在参数
  (grpo_*、model_type、resume_dir),补充完整示例
- dataflow.md: _RadixNode 命名修正
2026-05-08 18:07:57 +08:00
ViperEkura 44d7a4e959 refactor: 设计模式优化 inference 模块导入结构
- 新建 cache.py:SlotAllocator 对象池 + PrefixCacheManager

- 新建 sampling.py:Temperature/TopK/TopP 可组合策略

- TaskStatus 改用 Enum,GenerationParams 值对象模式

- _STOP 移至 cache.py,解除 engine→scheduler 轻量耦合

- 更新测试导入路径,ruff 格式检查通过
2026-05-08 16:57:57 +08:00
ViperEkura c4401512f2 fix: 修复长对话截断方向错误,保留最新 token 而非最早
- add_task 中 prompt 超长时改为保留末尾 token(prompt_ids[-max_prompt_len:])
  而非开头 token,确保多轮对话时模型能看到最近的提问上下文
2026-05-08 15:52:48 +08:00
ViperEkura a6f5ff3b37 fix: 修复 remove_task 未释放 KV cache slot 导致第二轮对话死锁
- remove_task() 现在释放 KV cache slot 和 prefix cache 引用
- _refill_active_batch 中 alloc 失败时将剩余 task 推回 waiting_queue
- 主循环增加 try/except 异常兜底,发送 _STOP 给所有 task
- 重构:server.py 全局变量改为 ServerState 类;automodel.py
  使用 Registry 替代裸 dict;合并 TrainContextBuilder 的 with_*
  方法到 build()
2026-05-08 14:53:04 +08:00
ViperEkura ffff05b2c6 refactor: 替换魔法字符串为_STOP sentinel,修复generator清理逻辑 2026-05-06 20:37:16 +08:00
ViperEkura b89f8436ea refactor: 将KV缓存槽位映射下沉到模型注意力层,移除_remap_kv和_writeback_kv 2026-05-06 20:01:22 +08:00
ViperEkura 123f25e339 fix: 修复KV缓存槽位索引错位、版本校验缺失与注意力掩码问题,合并预填充方法 2026-05-06 19:51:14 +08:00
ViperEkura 520de3ebe8 refactor: 重构推理引擎控制逻辑,修复连续批处理核心缺陷
- 修复 decode 阶段新任务覆盖已有任务的严重缺陷
- 修复线程安全问题(热路径无锁竞争)
- 修复前缀缓存引用计数管理不当导致缓存被驱逐
- 修复 pad_id 缺失导致全量 prefill 崩溃
- 修复 RoPE 位置错乱(不同位置任务共用 start_pos)
- 新增 slot 版本追踪实现前缀缓存零拷贝复用
- 新增异步流式生成接口避免阻塞事件循环
- 添加完整英文文档字符串
2026-05-06 16:04:06 +08:00
34 changed files with 2147 additions and 1650 deletions
+1
View File
@@ -15,6 +15,7 @@
!/.gitattributes !/.gitattributes
!/.dockerignore !/.dockerignore
!/Dockerfile !/Dockerfile
!/docker-compose.yml
!/assets/** !/assets/**
!/CONTRIBUTING.md !/CONTRIBUTING.md
!/LICENSE !/LICENSE
+47 -14
View File
@@ -27,9 +27,6 @@
## 📖 Table of Contents ## 📖 Table of Contents
<details open>
<summary><b>English</b></summary>
- [Features](#features) - [Features](#features)
- [Quick Start](#quick-start) - [Quick Start](#quick-start)
- [Documentation](#documentation) - [Documentation](#documentation)
@@ -37,8 +34,6 @@
- [Community](#community) - [Community](#community)
- [License](#license) - [License](#license)
</details>
--- ---
<a id="english"></a> <a id="english"></a>
@@ -51,7 +46,8 @@
- 💡 **Easy to Use**: Simple API with comprehensive examples and demos. - 💡 **Easy to Use**: Simple API with comprehensive examples and demos.
- 📦 **Lightweight**: Minimal dependencies, easy to deploy. - 📦 **Lightweight**: Minimal dependencies, easy to deploy.
- 🔬 **ResearchFriendly**: Modular design, easy to experiment with new ideas. - 🔬 **ResearchFriendly**: Modular design, easy to experiment with new ideas.
- 🤗 **HuggingFace Integration**: Compatible with HuggingFace models and datasets. - 🤗 **HuggingFace-Style API**: AutoModel/AutoTokenizer APIs inspired by HuggingFace for easy model and tokenizer loading.
- 🔌 **Dual API Compatibility**: Supports both OpenAI and Anthropic chat completion APIs out of the box.
### Quick Start ### Quick Start
@@ -72,16 +68,26 @@ pip install -e ".[dev]"
#### Train a Model #### Train a Model
```bash ```bash
python scripts/tools/train.py \ CUDA_VISIBLE_DEVICES=0,1,2,3 python scripts/tools/train.py \
--train_type=seq \ --train_type seq \
--data_root_path=/path/to/dataset \ --data_root_path /path/to/dataset \
--param_path=/path/to/param_path --param_path /path/to/model \
--batch_size 4 \
--accumulation_steps 8 \
--max_lr 3e-4 \
--warmup_steps 1000 \
--n_epoch 1
``` ```
Full reference at [Parameter Guide](assets/docs/params.md).
#### Generate Text #### Generate Text
```bash ```bash
python scripts/tools/generate.py --param_path=/path/to/param_path python scripts/tools/generate.py \
--param_path /path/to/model \
--input_json_file /path/to/input.json \
--output_json_file /path/to/output.json
``` ```
#### Docker #### Docker
@@ -104,13 +110,19 @@ docker run --gpus all -p 8000:8000 astrai:latest \
# Run with volume mount for data # Run with volume mount for data
docker run --gpus all -v /path/to/data:/data -it astrai:latest docker run --gpus all -v /path/to/data:/data -it astrai:latest
# Docker Compose (GPU, default)
docker compose up -d
# Docker Compose (CPU only)
docker compose --profile cpu up -d
``` ```
> **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`. > **Note**: `--gpus all` is required for CUDA support. Without it, `torch.cuda.is_available()` will return `False`.
#### Start HTTP Server #### Start HTTP Server
Start the inference server with OpenAI-compatible HTTP API: Start the inference server with OpenAI and Anthropic-compatible HTTP API:
```bash ```bash
python -m scripts.tools.server --port 8000 --device cuda python -m scripts.tools.server --port 8000 --device cuda
@@ -119,7 +131,7 @@ python -m scripts.tools.server --port 8000 --device cuda
Make requests: Make requests:
```bash ```bash
# Chat API (OpenAI compatible) # OpenAI-compatible
curl -X POST http://localhost:8000/v1/chat/completions \ curl -X POST http://localhost:8000/v1/chat/completions \
-H "Content-Type: application/json" \ -H "Content-Type: application/json" \
-d '{ -d '{
@@ -127,7 +139,7 @@ curl -X POST http://localhost:8000/v1/chat/completions \
"max_tokens": 512 "max_tokens": 512
}' }'
# Streaming response # OpenAI-compatible streaming
curl -X POST http://localhost:8000/v1/chat/completions \ curl -X POST http://localhost:8000/v1/chat/completions \
-H "Content-Type: application/json" \ -H "Content-Type: application/json" \
-d '{ -d '{
@@ -136,6 +148,27 @@ curl -X POST http://localhost:8000/v1/chat/completions \
"max_tokens": 500 "max_tokens": 500
}' }'
# Anthropic-compatible
curl -X POST http://localhost:8000/v1/messages \
-H "Content-Type: application/json" \
-d '{
"model": "astrai",
"system": "You are a helpful assistant.",
"messages": [{"role": "user", "content": "Hello"}],
"max_tokens": 512
}'
# Anthropic-compatible streaming with stop sequences
curl -X POST http://localhost:8000/v1/messages \
-H "Content-Type: application/json" \
-d '{
"model": "astrai",
"messages": [{"role": "user", "content": "Write a story"}],
"max_tokens": 500,
"stream": true,
"stop_sequences": ["The end"]
}'
# Health check # Health check
curl http://localhost:8000/health curl http://localhost:8000/health
``` ```
+47 -9
View File
@@ -52,7 +52,8 @@
- 💡 **易用**: 简洁的 API 与丰富的示例、演示。 - 💡 **易用**: 简洁的 API 与丰富的示例、演示。
- 📦 **轻量**: 依赖少,部署简单。 - 📦 **轻量**: 依赖少,部署简单。
- 🔬 **研究友好**: 模块化设计,便于实验新想法。 - 🔬 **研究友好**: 模块化设计,便于实验新想法。
- 🤗 **HuggingFace 集成**: 兼容 HuggingFace 模型与数据集 - 🤗 **HuggingFace 风格 API**: HuggingFace 的 AutoModel/AutoTokenizer 接口,方便加载模型和分词器
- 🔌 **双 API 兼容**: 同时支持 OpenAI 和 Anthropic 聊天补全 API,开箱即用。
### 快速开始 ### 快速开始
@@ -73,16 +74,26 @@ pip install -e ".[dev]"
#### 训练模型 #### 训练模型
```bash ```bash
python scripts/tools/train.py \ CUDA_VISIBLE_DEVICES=0,1,2,3 python scripts/tools/train.py \
--train_type=seq \ --train_type seq \
--data_root_path=/path/to/dataset \ --data_root_path /path/to/dataset \
--param_path=/path/to/param_path --param_path /path/to/model \
--batch_size 4 \
--accumulation_steps 8 \
--max_lr 3e-4 \
--warmup_steps 1000 \
--n_epoch 1
``` ```
完整参数列表见[参数说明](./params.md)。
#### 文本生成 #### 文本生成
```bash ```bash
python scripts/tools/generate.py --param_path=/path/to/param_path python scripts/tools/generate.py \
--param_path /path/to/model \
--input_json_file /path/to/input.json \
--output_json_file /path/to/output.json
``` ```
#### Docker #### Docker
@@ -105,13 +116,19 @@ docker run --gpus all -p 8000:8000 astrai:latest \
# 挂载数据卷 # 挂载数据卷
docker run --gpus all -v /path/to/data:/data -it astrai:latest docker run --gpus all -v /path/to/data:/data -it astrai:latest
# Docker ComposeGPU,默认)
docker compose up -d
# Docker Compose(仅 CPU
docker compose --profile cpu up -d
``` ```
> **注意**: 必须使用 `--gpus all` 才能启用 CUDA 支持,否则 `torch.cuda.is_available()` 将返回 `False`。 > **注意**: 必须使用 `--gpus all` 才能启用 CUDA 支持,否则 `torch.cuda.is_available()` 将返回 `False`。
#### 启动 HTTP 服务 #### 启动 HTTP 服务
启动推理服务器,支持 OpenAI 兼容的 HTTP API 启动推理服务器,支持 OpenAI 和 Anthropic 兼容的 HTTP API
```bash ```bash
python -m scripts.tools.server --port 8000 --device cuda python -m scripts.tools.server --port 8000 --device cuda
@@ -120,7 +137,7 @@ python -m scripts.tools.server --port 8000 --device cuda
发起请求: 发起请求:
```bash ```bash
# Chat APIOpenAI 兼容 # OpenAI 兼容
curl -X POST http://localhost:8000/v1/chat/completions \ curl -X POST http://localhost:8000/v1/chat/completions \
-H "Content-Type: application/json" \ -H "Content-Type: application/json" \
-d '{ -d '{
@@ -128,7 +145,7 @@ curl -X POST http://localhost:8000/v1/chat/completions \
"max_tokens": 512 "max_tokens": 512
}' }'
# 流式响应 # OpenAI 兼容流式
curl -X POST http://localhost:8000/v1/chat/completions \ curl -X POST http://localhost:8000/v1/chat/completions \
-H "Content-Type: application/json" \ -H "Content-Type: application/json" \
-d '{ -d '{
@@ -137,6 +154,27 @@ curl -X POST http://localhost:8000/v1/chat/completions \
"max_tokens": 500 "max_tokens": 500
}' }'
# Anthropic 兼容
curl -X POST http://localhost:8000/v1/messages \
-H "Content-Type: application/json" \
-d '{
"model": "astrai",
"system": "你是一个乐于助人的助手。",
"messages": [{"role": "user", "content": "你好"}],
"max_tokens": 512
}'
# Anthropic 兼容流式并设置停止序列
curl -X POST http://localhost:8000/v1/messages \
-H "Content-Type: application/json" \
-d '{
"model": "astrai",
"messages": [{"role": "user", "content": "写个故事"}],
"max_tokens": 500,
"stream": true,
"stop_sequences": ["结束"]
}'
# 健康检查 # 健康检查
curl http://localhost:8000/health curl http://localhost:8000/health
``` ```
+156 -188
View File
@@ -7,14 +7,12 @@ This document describes the data flow of the AstrAI project (a training and infe
AstrAI adopts a modular design with the following main components: AstrAI adopts a modular design with the following main components:
- **Dataset Module** (`astrai/dataset/`): Dataset, sampler, serialization tools - **Dataset Module** (`astrai/dataset/`): Dataset, sampler, serialization tools
- **Model Module** (`astrai/model/`): AutoModel, Transformer model and its submodules - **Model Module** (`astrai/model/`): AutoModel, Transformer model and its submodules
- **Training Module** (`astrai/trainer/`): Trainer, training context, strategies, schedulers - **Training Module** (`astrai/trainer/`): Trainer, training context, strategies, schedulers, callbacks, metric utilities
- **Inference Module** (`astrai/inference/`): Inference engine with continuous batching, streaming generation - **Inference Module** (`astrai/inference/`): Inference engine with continuous batching, streaming generation
- **Config Module** (`astrai/config/`): Model, training, scheduler, and other configurations - **Config Module** (`astrai/config/`): ModelConfig, TrainConfig
- **Factory Module** (`astrai/factory/`): Registry, BaseFactory for component registration - **Factory Module** (`astrai/factory/`): Registry, BaseFactory for component registration
- **Parallel Module** (`astrai/parallel/`): Distributed training support - **Parallel Module** (`astrai/parallel/`): Distributed training support
- **Serialization Module** (`astrai/serialization/`): HDF5 data loading, checkpoint management - **Serialization** (`astrai/serialization.py`): HDF5 data loading, checkpoint management
The data flow can generally be divided into two main lines: **Training Data Flow** and **Inference Data Flow**.
## Data Flow Diagram ## Data Flow Diagram
@@ -23,38 +21,36 @@ flowchart LR
subgraph A[Data Preparation] subgraph A[Data Preparation]
direction TB direction TB
A1[Raw Text] --> A2[AutoTokenizer] A1[Raw Text] --> A2[AutoTokenizer]
A2 --> A3[Serialize to .h5 files] A2 --> A3[Tokenized .h5 files]
A3 --> A4[BaseDataset] A3 --> A4[BaseDataset]
A4 --> A5[ResumableDistributedSampler] A4 --> A5[ResumableDistributedSampler]
A5 --> A6[PyTorch DataLoader] A5 --> A6[DataLoader]
end end
subgraph B[Training] subgraph B[Training]
direction TB direction TB
B1[Batch Data] --> B2[TrainContextBuilder] B1[DataLoader] --> B2[BaseStrategy]
B2 --> B3[TrainContext] B2 --> B3[Transformer Forward]
B3 --> B4[BaseStrategy] B3 --> B4[Loss + Backward]
B4 --> B5[Transformer] B4 --> B5[Gradient Accumulation]
B5 --> B6[Compute Loss] B5 -->|every accum_steps| B6[Optimizer Step]
B6 --> B7[Backward] B6 --> B7[LR Scheduler]
B7 --> B8[Optimizer] B7 -->|next batch| B2
B8 --> B9[LRScheduler] B6 --> B8[CheckpointCallback]
B9 --> B10[CheckpointCallback]
end end
subgraph C[Inference] subgraph C[Inference]
direction TB direction TB
C1[Checkpoint] --> C2[AutoModel] C1[Checkpoint] --> C2[AutoModel]
C2 --> C3[Transformer + Tokenizer] C1 --> C3[AutoTokenizer]
C3 --> C4[GenerationRequest + apply_chat_template] C2 --> C4[InferenceEngine]
C4 --> C5[InferenceEngine] C3 --> C4
C5 --> C6[InferenceScheduler] C4 --> C5[InferenceScheduler]
C6 --> C7[apply_sampling_strategies] C5 --> C6[Transformer Forward]
C7 --> C8[Transformer Forward] C6 --> C7[sample]
C8 --> C9[KV Cache + Prefix Cache] C7 --> C8{End?}
C9 --> C10{End Condition?} C8 -->|No| C6
C10 -->|No| C8 C8 -->|Yes| C9[Generated Text]
C10 -->|Yes| C11[Output Text]
end end
A --> B A --> B
@@ -63,207 +59,179 @@ flowchart LR
## Detailed Module Descriptions ## Detailed Module Descriptions
### 1. Dataset Module ### 1. Serialization (`astrai/serialization.py`)
#### 1.1 Serialization (`serialization.py`) - **`save_h5`**: Saves tensors by groups as HDF5 files (`.h5`), each key maps to a list of tensors
- **`save_h5`**: Saves multiple tensors by groups as HDF5 files (`.h5`), each key corresponds to a list of tensors - **`load_h5`**: Loads `.h5` files, returns `Dict[str, List[Tensor]]`, supports shared memory
- **`load_h5`**: Loads `.h5` files, returns `Dict[str, List[Tensor]]`, supports shared memory (`share_memory=True`) - **`Checkpoint`**: Encapsulates model state dict + epoch + iteration; uses safetensors
- **`Checkpoint` class**: Encapsulates model state dict, training epoch, iteration count; supports safetensors format for saving and loading
#### 1.2 Dataset (`dataset.py`) ### 2. Dataset Module
- **`BaseDataset`**: Abstract base class, defines common logic for window sampling, stride, etc.
- **`BaseSegmentFetcher`** and **`MultiSegmentFetcher`**: Efficiently fetch data from specified index ranges in multiple segments
- **`DatasetFactory`**: Factory pattern, supports dynamic registration of dataset types (`seq`, `sft`, `dpo`, `grpo`)
- After dataset loading, multiple data keys (such as `"sequence"`, `"mask"`) are managed through `MultiSegmentFetcher`
#### 1.3 Sampler (`sampler.py`) #### 2.1 Dataset (`dataset.py`)
- **`ResumableDistributedSampler`**: Resumable sampler supporting distributed training - **`BaseDataset`**: Abstract base class for windowed sequence sampling
- Records current epoch and iteration position, enabling training resume from breakpoints - **`BaseSegmentFetcher` / `MultiSegmentFetcher`**: Fetch tensor segments by index range
- Supports shuffle and drop_last options - **`DatasetFactory`**: Creates dataset instances by `train_type` (`seq`, `sft`, `dpo`, `grpo`)
- Data keys: `"sequence"` (SEQ), `"loss_mask"` (SFT), `"chosen_mask"/"rejected_mask"` (DPO), `"masks"` (GRPO)
### 2. Model Module #### 2.2 Sampler (`sampler.py`)
- **`ResumableDistributedSampler`**: Tracks `epoch` and `iter` for breakpoint resume; supports shuffle and drop_last
#### 2.1 Transformer / AutoModel (`transformer.py`, `automodel.py`) ### 3. Model Module
- **`AutoModel`**: Base class for autoregressive language models with `from_pretrained()` and `save_pretrained()` methods
- **`Transformer`**: Core autoregressive decoder architecture (registered via `@AutoModel.register('transformer')`)
- Contains embedding layer, multi-layer `DecoderBlock`, RMSNorm, and linear output head
- Supports weight tying (`tie_weight=True`) to reduce parameter count
- Uses Rotary Position Embedding (RoPE) to inject position information
- Supports loading from safetensors format with automatic model type detection from `config.json`
#### 2.2 Submodules (`module.py`) #### 3.1 Transformer / AutoModel
- **`RotaryEmbedding`**: Generates RoPE cos/sin cache - **`AutoModel`**: Base class with `from_pretrained()` / `save_pretrained()`
- **`DecoderBlock`**: Contains multi-head attention (supports GQA and MLA), feedforward network (FFN), residual connections - **`Transformer`**: Decoder-only architecture, registered via `@AutoModel.register('transformer')`
- **`GQA`**: Grouped Query Attention implementation - Embedding → N×DecoderBlock → RMSNorm → Linear lm_head
- **`MLA`**: Multi-Latent Attention implementation (like Qwen2-VL) - RoPE position encoding, optional weight tying
- **`MLP`**: Feed-forward network with SiLU activation and gated mechanism
- **`RMSNorm`**: Layer normalization variant
- **`Linear`**, **`Embedding`**: Custom linear layer and embedding layer, supporting parallelism wrappers
### 3. Training Module #### 3.2 Submodules (`module.py`)
- **`DecoderBlock`**: GQA attention + residual + MLP + RMSNorm
- **`GQA`**: Grouped Query Attention (also `MLA` for multi-latent attention)
- **`MLP`**: `SiLU(gate(x)) * up(x)` → down projection
- **`RotaryEmbedding`**: RoPE cos/sin cache
- **`RMSNorm`**: Layer normalization
#### 3.1 Training Context (`train_context.py`) ### 4. Training Module
- **`TrainContext`**: Data class encapsulating all components needed for training (model, optimizer, data loader, strategy, etc.)
- **`TrainContextBuilder`**: Builder pattern, progressively assembles training context, supports resume from checkpoint
#### 3.2 Trainer (`trainer.py`) #### 4.1 Training Context (`train_context.py`)
- **`Trainer`**: Main training loop, manages callbacks (progress bar, checkpoint, metric logging, gradient clipping, scheduler) - **`TrainContext`**: Dataclass holding model, optimizer, dataloader, strategy, scheduler, checkpoint state
- Supports distributed training (launches multi-process via `spawn_parallel_fn`) - **`TrainContextBuilder`**: Builder pattern — takes checkpoint for resume, builds all components
- Training steps include:
1. `on_train_begin` → 2. `on_epoch_begin` → 3. `on_batch_begin` → 4. Forward/loss calculation → 5. `on_batch_end` → 6. Gradient accumulation → 7. `on_step_begin` → 8. Optimizer update → 9. `on_step_end` → 10. `on_epoch_end`
#### 3.3 Strategy (`strategy.py`) #### 4.2 Trainer (`trainer.py`)
- **`BaseStrategy`**: Defines training strategy interface
- **`SEQStrategy`**: Standard next-token prediction training
- **`SFTStrategy`**: Supervised Fine-tuning with loss masking
- **`DPOStrategy`**: Direct Preference Optimization
- **`GRPOStrategy`**: Group Relative Policy Optimization
- Strategy receives batch data, executes model forward pass, loss calculation, returns loss tensor
- Created dynamically by `StrategyFactory` according to configuration
#### 3.4 Scheduler (`schedule.py`) The training loop is nested: **epoch****batch** (with step phase interspersed):
- **`BaseScheduler`**: Abstract base class defining learning rate scheduling interface
- **`CosineScheduler`**: Cosine decay scheduler with warmup
- **`SGDRScheduler`**: Stochastic Gradient Descent with Warm Restarts
- **`SchedulerFactory`**: Factory pattern, supports registration of various schedulers
- Scheduler is automatically created according to configuration and bound to optimizer
#### 3.5 Callbacks (`train_callback.py`) ```
- **`TrainCallback`**: Protocol interface for trainer callbacks on_train_begin
- **`CheckpointCallback`**: Saves model checkpoints at configurable intervals on_epoch_begin
- **`ProgressBarCallback`**: Displays training progress for each batch:
- **`MetricLoggerCallback`**: Logs training metrics to JSON files if iteration % accumulation_steps == 0: ← step phase
- **`GradientClippingCallback`**: Clips gradient norms on_step_begin → optimizer.step() → zero_grad → on_step_end
- **`SchedulerCallback`**: Steps learning rate scheduler ← batch phase
on_batch_begin → strategy(batch) → loss → backward → on_batch_end
iteration += 1
### 4. Factory Module on_epoch_end
on_train_end
```
#### 4.1 Registry and BaseFactory (`factory.py`) Key points:
- **`Registry`**: Flexible registry for component classes with category and priority support - `on_step_*` wraps optimizer step (fires every `accumulation_steps` batches)
- **`BaseFactory`**: Generic factory class for component registration and creation - `on_batch_*` wraps loss computation (fires every batch)
- Supports decorator-based registration pattern for extensible components - `SchedulerCallback` fires on `on_batch_end` — LR scheduler steps every batch
- Provides methods for registration, retrieval, and listing with filtering - `GradientClippingCallback` fires on `on_step_begin`
### 5. Parallel Module #### 4.3 Strategy (`strategy.py`)
- **`SEQStrategy`**: Next-token prediction, cross-entropy with label smoothing
- **`SFTStrategy`**: Supervised fine-tuning with loss masking
- **`DPOStrategy`**: Direct Preference Optimization with reference model
- **`GRPOStrategy`**: Group Relative Policy Optimization with clipped ratio
#### 5.1 Setup (`setup.py`) #### 4.4 Scheduler (`schedule.py`)
- **`spawn_parallel_fn`**: Spawns multiple processes for distributed training using PyTorch multiprocessing - **`CosineScheduler`**: Cosine decay + linear warmup
- **`setup_parallel`**: Context manager for initializing distributed process group (NCCL/CCL backend) - **`SGDRScheduler`**: Cosine annealing with warm restarts
- **`only_on_rank`**: Decorator to execute functions only on specific ranks - Created by `SchedulerFactory` and bound to optimizer
- **`get_rank`**: Returns current process rank in distributed group
- **`get_world_size`**: Returns total number of processes in distributed group
- **`get_current_device`**: Returns current device from environment
#### 5.2 Parallel Layers (`module.py`) #### 4.5 Callbacks
- **`ParallelModel`**: Base class for parallel models with process group - **`CheckpointCallback`**: Saves safetensors at `ckpt_interval` iterations
- **`ColumnParallelLinear`**: Column-parallel linear layer with input splitting and output gathering - **`ProgressBarCallback`**: tqdm progress display
- **`RowParallelLinear`**: Row-parallel linear layer with output reduction - **`MetricLoggerCallback`**: Writes JSONL metrics to `{ckpt_dir}/logs/`
- **`GradientClippingCallback`**: `clip_grad_norm_` on `on_step_begin`
- **`SchedulerCallback`**: `scheduler.step()` on `on_batch_end`
### 6. Inference Module ### 5. Inference Module
#### 6.1 Inference Engine (`engine.py`) #### 5.1 Inference Engine (`engine.py`)
- **`InferenceEngine`**: Unified inference interface, supports streaming and non-streaming generation - **`InferenceEngine`**: Facade over scheduler; provides `generate()`, `generate_with_request()`, `generate_async()`
- **`InferenceScheduler`**: Continuous batching scheduler with dynamic batch composition - Accepts `prompt: str | List[str]`, returns generator (stream) or string (non-stream)
- **`GenerationRequest`**: Encapsulates generation parameters (top_k, top_p, temperature, max_len, messages, etc.)
- **`messages` format**: List of message dictionaries with `role` (system/user/assistant) and `content`
- **`apply_chat_template`** (from `tokenizer.py`): Converts messages into prompt string using ChatML format
- Provides streaming (`stream=True`) and non-streaming (`stream=False`) generation interfaces
- Supports continuous batching with `max_batch_size` and `max_seq_len` parameters
- Uses separate model and tokenizer initialization for flexibility
#### 6.2 Scheduler (`scheduler.py`) #### 5.2 Scheduler 4-Phase Loop (`scheduler.py`)
- **`Task`**: Individual generation task with state management (PENDING, RUNNING, FINISHED, ABORTED)
- **`TaskStatus`**: Task state enumeration
- **`apply_sampling_strategies`**: Applies temperature, top-k, top-p sampling to logits
- **`PrefixCacheManager`**: Radix tree-based prefix cache with LRU eviction for efficient KV cache reuse
- **`RadixNode`**: Tree node structure for prefix caching
- Continuous batching: new requests can join at any time, completed requests are released immediately
#### 6.3 Server (`server.py`) Background thread runs continuously:
- FastAPI-based HTTP inference server
- OpenAI-compatible `/v1/chat/completions` endpoint
- Health check and statistics endpoints
- Supports both streaming and non-streaming responses
### 7. Tokenizer Module ```
1. Cleanup → Remove finished tasks, free KV cache pages
2. Refill → Pop from waiting_queue, alloc pages, add to active
3. Prefill → Group active tasks by prompt_len, run full forward pass
4. Decode → Pick largest same-position group, run single-token forward
```
#### 7.1 Tokenizer (`tokenizer.py`) - **`Task`**: Tracks prompt_ids, output_ids, page_table, status (PENDING/RUNNING/FINISHED/ABORTED)
- Implemented based on HuggingFace tokenizers library (Byte-Level BPE) - **`PagedCache`**: Bitmask-based page allocator with page-table-indirected read/write
- **`AutoTokenizer`**: Auto-loading tokenizer class - **`CacheView`**: Batch view bundling cache + page table for attention layers
- Supports special tokens: `<begin▁of▁sentence>`, `<end▁of▁sentence>`, `<|▁pad▁|>`, `<im▁start>`, `<im▁end>` - **`sample()`**: Temperature → top-k → top-p → multinomial
- Provides `encode`/`decode` methods for mutual conversion between text and token IDs
- Uses `AutoTokenizer` for loading pre-trained tokenizers
#### 7.2 Chat Template (`chat_template.py`) #### 5.3 Server (`server.py`)
- **`ChatTemplate`**: Jinja2-based chat template with rendering support - FastAPI with OpenAI `/v1/chat/completions` and Anthropic `/v1/messages` endpoints
- Handles multi-role message formatting (system, user, assistant) - Streaming via SSE, health check at `/health`, stats at `/stats`
- Supports dynamic prompts and generation prompts
## Training Data Flow - Detailed Steps ### 6. Tokenizer Module
- **`AutoTokenizer`**: Wraps HuggingFace tokenizers (BBPE); `encode`/`decode`/`apply_chat_template`
- **`ChatTemplate`**: Jinja2-based template rendering for multi-turn chat
### 7. Factory & Parallel
- **`Registry` / `BaseFactory`**: Decorator-based component registration
- **`spawn_parallel_fn`**: Multi-process DDP launcher with NCCL backend
- **`ParallelModel` / `ColumnParallelLinear` / `RowParallelLinear`**: Tensor model parallelism
## Training Data Flow — Detailed Steps
1. **Data Preparation** 1. **Data Preparation**
- Raw text is converted to token ID sequences through AutoTokenizer - Raw text → token IDs via `AutoTokenizer.encode()`
- Token ID sequences (possibly with masks, labels, etc.) are saved by groups as `.h5` files - Save as `.h5` files (groups of tensor lists per data key)
- Files can contain multiple segments, each segment corresponds to a tensor
2. **Dataset Loading** 2. **Dataset Loading**
- `BaseDataset`'s `load` method calls `load_h5`, obtaining `segments` dictionary - `BaseDataset.load()` calls `load_h5()`, builds `MultiSegmentFetcher`
- Create `MultiSegmentFetcher` to manage data for multiple keys - Sliding window of `window_size` with `stride` determines sample boundaries
- Calculate total sample count, and determine start/end indices for each sample based on window size and stride
3. **Sampling and Batch Loading** 3. **Sampling & Batching**
- `ResumableDistributedSampler` generates index sequence based on current epoch and iteration position - `ResumableDistributedSampler` produces shuffled index sequences
- PyTorch `DataLoader` uses sampler to get indices, calls dataset's `__getitem__` to get actual data - `DataLoader` fetches `[batch_size, window_size]` tensors via `__getitem__`
- Batch data shape is `[batch_size, window_size]` (or varies according to specific dataset type)
4. **Strategy Forward and Loss Calculation** 4. **Strategy Forward**
- Batch data is passed to strategy (such as `SEQStrategy`) - Strategy receives batch, calls `Transformer.forward()` for logits
- Strategy internally calls `Transformer` model, obtaining logits - Computes task-specific loss (cross-entropy, DPO, GRPO)
- Calculate cross-entropy loss (or DPO loss, etc.) according to task type
- Return loss tensor
5. **Backpropagation and Optimization** 5. **Backward & Accumulation**
- Loss is normalized by dividing by accumulation steps, then `loss.backward()` is executed - `loss = raw_loss / accumulation_steps`
- After accumulating `accumulation_steps` batches, optimizer `step()` and `zero_grad()` are executed - `loss.backward()` accumulates gradients
- Learning rate scheduler updates learning rate after each step - Every `accumulation_steps` batches: `optimizer.step()``zero_grad()`
- Every batch: `scheduler.step()` updates learning rate
6. **Checkpoint Saving** 6. **Checkpoint**
- `CheckpointCallback` saves checkpoints at set intervals - `CheckpointCallback` saves `model.state_dict()` + metadata to safetensors at `ckpt_interval` iterations
- Checkpoints contain model state dict, current epoch, iteration, and other metadata - Does NOT save optimizer/scheduler state (resume resets those)
- Saved in safetensors format, ensuring safety and efficiency
## Inference Data Flow - Detailed Steps ## Inference Data Flow Detailed Steps
1. **Model Loading** 1. **Model Loading**
- Load `Transformer` model from checkpoint via `AutoModel.from_pretrained()` - `AutoModel.from_pretrained(path)` loads weights from safetensors
- Set model to evaluation mode (`model.eval()`), enable inference mode (`torch.inference_mode`) - `torch.inference_mode()` wraps generation
2. **Prompt Construction and Encoding** 2. **Prompt Construction**
- User messages (list of dict with role and content) are converted to ChatML format string through `apply_chat_template` method in tokenizer - Messages `apply_chat_template(messages, tokenize=False)` → prompt string
- Tokenizer encodes prompt string to token ID sequence `input_ids` - `tokenizer.encode(prompt)` → token IDs (truncated to `max_prompt_len`)
- For batch generation, use `pad_sequence` for padding
3. **Autoregressive Generation Loop** 3. **Continuous Batching Loop**
- Initialize KV cache (optional) and prefix cache - **Cleanup**: Finished tasks → `stream_callback(STOP)`, free KV pages
- Loop until generating `max_len` tokens or encountering stop token: - **Refill**: Pop from waiting queue, `PagedCache.alloc_n()` for prompt pages
- Input current `input_ids` (or cached new token) to model, obtain `logits` - **Prefill**: Group by prompt length, run full forward with `start_pos=0`
- Apply `apply_sampling_strategies` (temperature, top-k, top-p) to `logits` - **Decode**: Pick position group with most tasks, single-token forward:
- Sample next token ID from the processed distribution - Model forward → `logits``sample()` → next token ID
- Append new token to `input_ids`, while updating KV cache - Append to `output_ids`, update `output_tokens`
- For streaming generation, yield each token to caller immediately - `_maybe_alloc_page()` grows page table as needed
- `stream_callback(token)` for streaming clients
4. **Decoding and Output** 4. **Output**
- Decode generated token ID sequence to text through tokenizer - `tokenizer.decode(output_ids)` → text
- Remove special tokens, return plain text response - Return to caller (streaming: token-by-token; non-streaming: complete string)
## Checkpoint and Serialization ## Checkpoint & Serialization
- **Training Checkpoint**: Saves model parameters, optimizer state, scheduler state, current epoch and iteration - **Training Checkpoint**: safetensors weights + epoch/iteration metadata. Optimizer/scheduler state is NOT persisted.
- **Model Parameters**: Supports safetensors format, automatically handles special logic like weight tying during loading - **Inference Loading**: `AutoModel.from_pretrained()` loads from the same safetensors format.
- **Dataset Serialization**: HDF5 format supports efficient random access and shared memory, suitable for large-scale pre-training data - **Dataset Serialization**: HDF5 with shared memory support for large-scale pre-training data.
## Summary > Document Update Time: 2026-05-09
The data flow design of AstrAI reflects the characteristics of modularity, extensibility, and resumability. The training data flow supports large-scale distributed training through chunk loading, resumable sampling, gradient accumulation, and other mechanisms; the inference data flow achieves efficient text generation using KV cache, prefix caching, and sampling strategies. Clear interfaces between modules facilitate customization and extension.
> Document Update Time: 2026-04-09
+125 -100
View File
@@ -50,7 +50,6 @@ classDiagram
+str master_port +str master_port
+Callable parallel_wrapper +Callable parallel_wrapper
+Callable state_dict_fn +Callable state_dict_fn
+List[int] device_ids
+str device_type +str device_type
+dict extra_kwargs +dict extra_kwargs
+validate() +validate()
@@ -85,8 +84,8 @@ classDiagram
} }
class BaseSegmentFetcher { class BaseSegmentFetcher {
+List~Tensor~ segments +List[Tensor] segments
+List~int~ cum_lengths +List[int] cum_lengths
+int total_length +int total_length
+fetch_data(begin_idx, end_idx) Tensor +fetch_data(begin_idx, end_idx) Tensor
} }
@@ -99,8 +98,8 @@ classDiagram
} }
class ResumableDistributedSampler { class ResumableDistributedSampler {
+int start_epoch +int epoch
+int start_iter +int iter
} }
class DatasetFactory { class DatasetFactory {
@@ -109,7 +108,9 @@ classDiagram
+create(train_type, window_size, stride) BaseDataset +create(train_type, window_size, stride) BaseDataset
+load(train_type, load_path, window_size, stride) BaseDataset +load(train_type, load_path, window_size, stride) BaseDataset
} }
}
namespace serialization {
class Checkpoint { class Checkpoint {
+dict state_dict +dict state_dict
+int epoch +int epoch
@@ -122,7 +123,7 @@ classDiagram
namespace model { namespace model {
class AutoModel { class AutoModel {
+ModelConfig config +ModelConfig config
+Dict _registry +Registry _registry
+register(model_type) decorator +register(model_type) decorator
+get_model_class(model_type) Type +get_model_class(model_type) Type
+from_pretrained(path, disable_random_init) nn.Module +from_pretrained(path, disable_random_init) nn.Module
@@ -137,7 +138,7 @@ classDiagram
+ModuleList layers +ModuleList layers
+RMSNorm norm +RMSNorm norm
+Linear lm_head +Linear lm_head
+forward(input_ids, input_mask, persistent_key_values, start_pos) Dict +forward(input_ids, input_mask, paged_cache, start_pos) Dict
+load_state_dict(state_dict) +load_state_dict(state_dict)
+state_dict() +state_dict()
} }
@@ -147,7 +148,7 @@ classDiagram
+RMSNorm input_norm +RMSNorm input_norm
+MLP mlp +MLP mlp
+RMSNorm post_attention_norm +RMSNorm post_attention_norm
+forward(x, rotary_emb, attention_mask, kv_cache, start_pos) Tensor +forward(x, rotary_emb, attention_mask, paged_cache, start_pos) Tensor
} }
class GQA { class GQA {
@@ -156,18 +157,20 @@ classDiagram
+int head_dim +int head_dim
+Linear q_proj, k_proj, v_proj, o_proj +Linear q_proj, k_proj, v_proj, o_proj
+RMSNorm q_norm, k_norm +RMSNorm q_norm, k_norm
+forward(x, rotary_emb, mask, kv_cache, start_pos) Tensor +forward(x, rotary_emb, mask, paged_cache, start_pos) Tensor
} }
class MLA { class MLA {
+int n_heads +int n_heads
+int n_kv_heads +int n_kv_heads
+int head_dim +int head_dim
+Linear q_a_proj, q_b_proj, q_c_proj +int kv_lora_rank
+Linear kv_a_proj, kv_b_proj, kv_c_proj +int qk_nope_head_dim
+int qk_rope_head_dim
+Linear q_proj, kv_a_proj, kv_b_proj
+Linear o_proj +Linear o_proj
+RMSNorm q_norm, k_norm +RMSNorm kv_norm
+forward(x, rotary_emb, mask, kv_cache, start_pos) Tensor +forward(x, rotary_emb, mask, paged_cache, start_pos) Tensor
} }
class MLP { class MLP {
@@ -191,7 +194,7 @@ classDiagram
+int dim +int dim
+int max_len +int max_len
+float base +float base
+forward(x, start_pos) Tuple~Tensor, Tensor~ +forward(x, start_pos) Tuple[Tensor, Tensor]
} }
class Embedding { class Embedding {
@@ -202,14 +205,14 @@ classDiagram
namespace tokenize { namespace tokenize {
class AutoTokenizer { class AutoTokenizer {
+List~str~ stop_ids +List[int] stop_ids
+int bos_id +int bos_id
+int eos_id +int eos_id
+int pad_id +int pad_id
+vocab_size int +vocab_size int
+encode(tokens, out_ids, add_special_tokens) List~int~ +encode(tokens, out_ids, add_special_tokens) List[int]
+decode(tokens, skip_special_tokens) str +decode(tokens, skip_special_tokens) str
+apply_chat_template(messages, tokenize) Union~str, List[int]~ +apply_chat_template(messages, tokenize) Union[str, List[int]]
+set_chat_template(template) +set_chat_template(template)
+load(path) +load(path)
+from_pretrained(path) AutoTokenizer +from_pretrained(path) AutoTokenizer
@@ -218,7 +221,7 @@ classDiagram
class ChatTemplate { class ChatTemplate {
+String template_str +String template_str
+render(messages, add_generation_prompt) str +render(messages, system_prompt, **extra_variables) str
+from_string(template) ChatTemplate +from_string(template) ChatTemplate
} }
} }
@@ -228,7 +231,7 @@ classDiagram
+Dict _entries +Dict _entries
+register(name, component_cls, category, priority) +register(name, component_cls, category, priority)
+get(name) Type +get(name) Type
+list_names() List~str~ +list_names() List[str]
} }
class BaseFactory { class BaseFactory {
@@ -242,10 +245,10 @@ classDiagram
namespace trainer { namespace trainer {
class Trainer { class Trainer {
+TrainConfig train_config +TrainConfig train_config
+List~TrainCallback~ callbacks +List[TrainCallback] callbacks
+train(checkpoint) +train(checkpoint)
+_build_context(checkpoint) TrainContext +_build_context(checkpoint) TrainContext
+_get_default_callbacks() List~TrainCallback~ +_get_default_callbacks() List[TrainCallback]
} }
class TrainContext { class TrainContext {
@@ -265,8 +268,6 @@ classDiagram
class TrainContextBuilder { class TrainContextBuilder {
+TrainConfig config +TrainConfig config
+with_checkpoint(checkpoint) TrainContextBuilder +with_checkpoint(checkpoint) TrainContextBuilder
+with_dataloader() TrainContextBuilder
+with_strategy() TrainContextBuilder
+build() TrainContext +build() TrainContext
} }
@@ -308,7 +309,7 @@ classDiagram
} }
class BaseScheduler { class BaseScheduler {
+get_lr() List~float~ +get_lr() List[float]
+step() +step()
} }
@@ -390,12 +391,9 @@ classDiagram
+InferenceScheduler scheduler +InferenceScheduler scheduler
+int max_batch_size +int max_batch_size
+Optional int max_seq_len +Optional int max_seq_len
+int max_prefix_len
+int cache_capacity
+Tensor kv_cache
+Tensor seq_mask
+generate(prompt, stream, max_tokens, temperature, top_p, top_k) Union[Generator, str, List[str]] +generate(prompt, stream, max_tokens, temperature, top_p, top_k) Union[Generator, str, List[str]]
+generate_with_request(request) Union[Generator, str, List[str]] +generate_with_request(request) Union[Generator, str, List[str]]
+generate_async(prompt, max_tokens, temperature, top_p, top_k) AsyncGenerator
+get_stats() Dict +get_stats() Dict
+shutdown() +shutdown()
} }
@@ -403,10 +401,11 @@ classDiagram
class InferenceScheduler { class InferenceScheduler {
+nn.Module model +nn.Module model
+AutoTokenizer tokenizer +AutoTokenizer tokenizer
+ModelConfig config +PagedCache page_cache
+Tuple kv_cache +int max_batch_size
+Tensor seq_mask +int max_seq_len
+PrefixCacheManager prefix_cache +int max_prompt_len
+int page_size
+List waiting_queue +List waiting_queue
+List active_tasks +List active_tasks
+add_task(prompt, max_tokens, temperature, top_p, top_k, stream_callback) str +add_task(prompt, max_tokens, temperature, top_p, top_k, stream_callback) str
@@ -416,22 +415,26 @@ classDiagram
+get_stats() Dict +get_stats() Dict
} }
class PrefixCacheManager { class PagedCache {
+RadixNode root +int page_size
+int max_capacity +int _free_mask
+List lru +List[int] _refs
+insert(token_ids, slot) +Tensor k_cache
+find_longest_prefix(token_ids) Tuple[int, int] +Tensor v_cache
+release(token_ids) +alloc() int
+alloc_n(n) List[int]
+free(idx)
+bind(page_table, total_len) CacheView
+write(layer_id, page_table, start_pos, k, v)
+gather(layer_id, page_table) Tuple[Tensor, Tensor]
} }
class RadixNode { class CacheView {
+Dict children +PagedCache _cache
+int hash +Tensor _page_table
+int slot +int _total_len
+int ref_count +write(layer_id, start_pos, k, v)
+float last_access +gather(layer_id) Tuple[Tensor, Tensor]
+List token_sequence
} }
class Task { class Task {
@@ -445,38 +448,71 @@ classDiagram
+List output_ids +List output_ids
+int input_tokens +int input_tokens
+int output_tokens +int output_tokens
+int slot +List[int] page_table
+int n_pages
+float arrival_time
+float finish_time
+Callable stream_callback +Callable stream_callback
+int next_pos
+is_finished(stop_ids) bool +is_finished(stop_ids) bool
} }
class TaskStatus { class TaskStatus {
+str PENDING <<enumeration>>
+str RUNNING PENDING
+str FINISHED RUNNING
+str ABORTED FINISHED
} ABORTED
class Server {
+start()
+predict(request)
} }
class GenerationRequest { class GenerationRequest {
+List[Dict] messages
+GenerationParams params
+bool stream
}
class GenerationParams {
<<value object>>
+int top_k +int top_k
+float top_p +float top_p
+float temperature +float temperature
+int max_len +int max_tokens
+List~Dict~ messages }
+stream bool
class BaseSamplingStrategy {
<<abstract>>
+apply(logits, filter_value) Tensor
}
class TemperatureStrategy {
+float temperature
+apply(logits, filter_value) Tensor
}
class TopKStrategy {
+int top_k
+apply(logits, filter_value) Tensor
}
class TopPStrategy {
+float top_p
+apply(logits, filter_value) Tensor
}
class SamplingPipeline {
+List strategies
+apply(logits, filter_value) Tensor
+sample(logits, filter_value) Tensor
} }
class _Result { class _Result {
+List~str~ tokens +List[str] tokens
+List~str~ results +List[str] results
+List~bool~ done_flags +List[bool] _done
+append(token, idx) +append(token, idx)
+get_results() List~str~ +get_results() List[str]
+pop_all() List[str]
+wait(timeout) bool
} }
class ChatMessage { class ChatMessage {
@@ -485,28 +521,21 @@ classDiagram
} }
class ChatCompletionRequest { class ChatCompletionRequest {
+List~ChatMessage~ messages +List[ChatMessage] messages
+float temperature +float temperature
+float top_p +float top_p
+int top_k +int top_k
+int max_tokens +int max_tokens
+bool stream +bool stream
+Optional~str~ system_prompt +Optional[str] stop
} +Optional[int] n
class CompletionResponse {
+str id
+str object
+int created
+str model
+List~Dict~ choices
} }
} }
namespace parallel { namespace parallel {
class ParallelSetup { class ParallelFunctions {
+spawn_parallel_fn(fn, nprocs) +spawn_parallel_fn(fn, nprocs)
+setup_parallel(rank, world_size, backend, master_addr, master_port, device_type, device_ids) +setup_parallel(rank, world_size, backend, master_addr, master_port, device_type)
} }
class ParallelModel { class ParallelModel {
@@ -539,10 +568,10 @@ classDiagram
Trainer --> TrainContextBuilder : builds Trainer --> TrainContextBuilder : builds
Trainer --> TrainCallback : manages Trainer --> TrainCallback : manages
TrainContextBuilder --> TrainContext : creates TrainContextBuilder --> TrainContext : creates
Checkpoint ..> Checkpoint : saves/loads
TrainContext --> Checkpoint : manages TrainContext --> Checkpoint : manages
TrainContext --> BaseStrategy : uses TrainContext --> BaseStrategy : uses
TrainContext --> BaseScheduler : uses TrainContext --> BaseScheduler : uses
AutoModel --> ModelConfig : contains
SchedulerFactory ..> BaseScheduler : creates SchedulerFactory ..> BaseScheduler : creates
BaseScheduler <|-- CosineScheduler BaseScheduler <|-- CosineScheduler
BaseScheduler <|-- SGDRScheduler BaseScheduler <|-- SGDRScheduler
@@ -553,30 +582,32 @@ classDiagram
TrainCallback <|-- ProgressBarCallback TrainCallback <|-- ProgressBarCallback
TrainCallback <|-- MetricLoggerCallback TrainCallback <|-- MetricLoggerCallback
InferenceEngine --> InferenceScheduler : uses InferenceEngine --> InferenceScheduler : uses
InferenceEngine --> GenerationRequest : uses
GenerationRequest --> GenerationParams : contains
InferenceScheduler --> Task : manages InferenceScheduler --> Task : manages
Task --> TaskStatus : uses
InferenceScheduler --> TaskStatus : uses InferenceScheduler --> TaskStatus : uses
InferenceScheduler --> PagedCache : uses
InferenceScheduler --> Transformer : uses InferenceScheduler --> Transformer : uses
InferenceEngine --> Transformer : uses InferenceEngine --> Transformer : uses
InferenceEngine --> GenerationRequest : uses InferenceEngine --> _Result : uses
Server --> InferenceEngine : uses BaseSamplingStrategy <|-- TemperatureStrategy
Server --> ChatMessage : uses BaseSamplingStrategy <|-- TopKStrategy
Server --> ChatCompletionRequest : uses BaseSamplingStrategy <|-- TopPStrategy
Server --> CompletionResponse : uses SamplingPipeline --> BaseSamplingStrategy : composes
ParallelSetup --> Trainer : enables
BaseDataset <|-- SEQDataset BaseDataset <|-- SEQDataset
BaseDataset <|-- SFTDataset BaseDataset <|-- SFTDataset
BaseDataset <|-- DPODataset BaseDataset <|-- DPODataset
BaseDataset <|-- GRPODataset BaseDataset <|-- GRPODataset
DatasetFactory ..> BaseDataset : creates DatasetFactory ..> BaseDataset : creates
BaseSegmentFetcher --> MultiSegmentFetcher : used by MultiSegmentFetcher --> BaseSegmentFetcher : uses
MultiSegmentFetcher --> BaseDataset : used by BaseDataset --> MultiSegmentFetcher : uses
AutoModel <|-- Transformer AutoModel <|-- Transformer
AutoModel --> ModelConfig : contains AutoModel --> ModelConfig : contains
Transformer --> DecoderBlock : uses Transformer --> DecoderBlock : uses
Transformer --> RotaryEmbedding : uses Transformer --> RotaryEmbedding : uses
Transformer --> Embedding : uses Transformer --> Embedding : uses
DecoderBlock --> GQA : uses DecoderBlock --> GQA : uses
DecoderBlock --> MLA : uses
DecoderBlock --> MLP : uses DecoderBlock --> MLP : uses
DecoderBlock --> RMSNorm : uses DecoderBlock --> RMSNorm : uses
TrainContextBuilder --> ResumableDistributedSampler : creates TrainContextBuilder --> ResumableDistributedSampler : creates
@@ -584,9 +615,6 @@ classDiagram
ParallelModel <|-- RowParallelLinear ParallelModel <|-- RowParallelLinear
ParallelModel <|-- ColumnParallelLinear ParallelModel <|-- ColumnParallelLinear
AutoTokenizer --> ChatTemplate : uses AutoTokenizer --> ChatTemplate : uses
InferenceScheduler --> PrefixCacheManager : uses
InferenceScheduler --> RadixNode : uses
Checkpoint ..> Checkpoint : saves/loads
TrainConfig --> DatasetFactory : selects TrainConfig --> DatasetFactory : selects
TrainConfig --> SchedulerFactory : selects TrainConfig --> SchedulerFactory : selects
TrainConfig --> CallbackFactory : selects TrainConfig --> CallbackFactory : selects
@@ -602,12 +630,13 @@ classDiagram
| Module | Components | Description | | Module | Components | Description |
|--------|------------|-------------| |--------|------------|-------------|
| **astrai.config** | ModelConfig, TrainConfig | Configuration management | | **astrai.config** | ModelConfig, TrainConfig | Configuration management |
| **astrai.dataset** | BaseDataset, SEQDataset, SFTDataset, DPODataset, GRPODataset, BaseSegmentFetcher, MultiSegmentFetcher, ResumableDistributedSampler, DatasetFactory, Checkpoint | Dataset loading and management | | **astrai.dataset** | BaseDataset, SEQDataset, SFTDataset, DPODataset, GRPODataset, BaseSegmentFetcher, MultiSegmentFetcher, ResumableDistributedSampler, DatasetFactory | Dataset loading and management |
| **astrai.serialization** | Checkpoint, save_h5, load_h5 | Model serialization and checkpoint management |
| **astrai.model** | AutoModel, Transformer, DecoderBlock, GQA, MLA, MLP, RMSNorm, Linear, RotaryEmbedding, Embedding | Neural network model | | **astrai.model** | AutoModel, Transformer, DecoderBlock, GQA, MLA, MLP, RMSNorm, Linear, RotaryEmbedding, Embedding | Neural network model |
| **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template | | **astrai.tokenize** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
| **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy, StrategyFactory, BaseScheduler, SchedulerFactory, TrainCallback, CallbackFactory | Training workflow management | | **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy, StrategyFactory, BaseScheduler, SchedulerFactory, TrainCallback, CallbackFactory | Training workflow management |
| **astrai.inference** | InferenceEngine, InferenceScheduler, Task, TaskStatus, Server, GenerationRequest, PrefixCacheManager, ChatMessage, ChatCompletionRequest, CompletionResponse | Inference service with continuous batching | | **astrai.inference** | InferenceEngine, InferenceScheduler, PagedCache, CacheView, Task, TaskStatus, GenerationParams, GenerationRequest, BaseSamplingStrategy, TemperatureStrategy, TopKStrategy, TopPStrategy, SamplingPipeline, ChatMessage, ChatCompletionRequest | Inference service with continuous batching and paged KV cache |
| **astrai.parallel** | ParallelSetup, ColumnParallelLinear, RowParallelLinear | Distributed parallel | | **astrai.parallel** | ParallelFunctions, ParallelModel, ColumnParallelLinear, RowParallelLinear | Distributed parallel |
| **astrai.factory** | Registry, BaseFactory | Generic component registration | | **astrai.factory** | Registry, BaseFactory | Generic component registration |
### Design Patterns ### Design Patterns
@@ -618,8 +647,10 @@ classDiagram
| **Builder** | `TrainContextBuilder` | Chain-building training context, step-by-step initialization of components | | **Builder** | `TrainContextBuilder` | Chain-building training context, step-by-step initialization of components |
| **Factory** | `StrategyFactory`, `SchedulerFactory`, `DatasetFactory`, `CallbackFactory`, `BaseFactory` | Decorator registration mechanism, dynamically create training strategies, schedulers, datasets, and callbacks | | **Factory** | `StrategyFactory`, `SchedulerFactory`, `DatasetFactory`, `CallbackFactory`, `BaseFactory` | Decorator registration mechanism, dynamically create training strategies, schedulers, datasets, and callbacks |
| **Observer** | `TrainCallback`, `CallbackFactory` | Callback mechanism for training process monitoring (checkpoint, early stopping, metrics) | | **Observer** | `TrainCallback`, `CallbackFactory` | Callback mechanism for training process monitoring (checkpoint, early stopping, metrics) |
| **Singleton** | `TrainContext` | Training process global state management | | **Context** | `TrainContext` | Training process state container with model, optimizer, scheduler and checkpoint |
| **Registry** | `BaseFactory`, `Registry` | Generic component registration with category and priority support | | **Registry** | `BaseFactory`, `Registry` | Generic component registration with category and priority support |
| **Object Pool** | `PagedCache` | Page-based KV cache with O(1) alloc/free via bitmask |
| **Strategy (Sampling)** | `BaseSamplingStrategy`, `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations with temperature, top-k, top-p |
| **Producer-Consumer** | `InferenceScheduler`, `Task`, `waiting_queue`, `active_tasks` | Continuous batching with dynamic task queue management | | **Producer-Consumer** | `InferenceScheduler`, `Task`, `waiting_queue`, `active_tasks` | Continuous batching with dynamic task queue management |
| **Event-Driven** | `threading.Event`, `_task_event` | Non-blocking wait mechanism for task scheduling using Python's `threading` module | | **Event-Driven** | `threading.Event`, `_task_event` | Non-blocking wait mechanism for task scheduling using Python's `threading` module |
| **AutoModel Registry** | `AutoModel`, `Transformer` | Model type registration and dynamic loading via decorator pattern | | **AutoModel Registry** | `AutoModel`, `Transformer` | Model type registration and dynamic loading via decorator pattern |
@@ -630,8 +661,8 @@ classDiagram
1. **Configuration → Training**: `TrainConfig` contains `ModelConfig`, holds model, dataset, optimizer and other references 1. **Configuration → Training**: `TrainConfig` contains `ModelConfig`, holds model, dataset, optimizer and other references
2. **Training Flow**: `Trainer``TrainContextBuilder``TrainContext`, uses `BaseStrategy` to compute loss 2. **Training Flow**: `Trainer``TrainContextBuilder``TrainContext`, uses `BaseStrategy` to compute loss
3. **Strategy Selection**: `StrategyFactory` creates corresponding strategy instance based on `train_type` 3. **Strategy Selection**: `StrategyFactory` creates corresponding strategy instance based on `train_type`
4. **Inference Flow**: `Server``InferenceEngine``InferenceScheduler``Transformer`, supports continuous batching with streaming/non-streaming 4. **Inference Flow**: `InferenceEngine``InferenceScheduler``Transformer`, uses `PagedCache` for paged KV cache management and `SamplingPipeline` for efficient continuous batching with streaming/non-streaming
5. **Distributed Support**: `ParallelSetup` provides multi-process training capability for `Trainer` 5. **Distributed Support**: `spawn_parallel_fn` and `setup_parallel` provide multi-process training capability for `Trainer`
6. **Dataset Loading**: `DatasetFactory` creates datasets (SEQDataset, SFTDataset, DPODataset, GRPODataset), supports HDF5 loading via `BaseSegmentFetcher` and `MultiSegmentFetcher` 6. **Dataset Loading**: `DatasetFactory` creates datasets (SEQDataset, SFTDataset, DPODataset, GRPODataset), supports HDF5 loading via `BaseSegmentFetcher` and `MultiSegmentFetcher`
7. **Checkpoint Management**: `Checkpoint` handles model state serialization/deserialization with safetensors 7. **Checkpoint Management**: `Checkpoint` handles model state serialization/deserialization with safetensors
8. **Scheduler Support**: `SchedulerFactory` creates learning rate schedulers (CosineScheduler, SGDRScheduler) 8. **Scheduler Support**: `SchedulerFactory` creates learning rate schedulers (CosineScheduler, SGDRScheduler)
@@ -675,12 +706,6 @@ $$
L_{\text{GRPO}} = -\mathbb{E} \left[ \min\left( \frac{\pi_\theta(a|s)}{\pi_{\text{ref}}(a|s)} \cdot A, \text{clip}\left(\frac{\pi_\theta(a|s)}{\pi_{\text{ref}}(a|s)}, 1-\epsilon, 1+\epsilon\right) \cdot A \right) \right] + \lambda \cdot D_{KL} L_{\text{GRPO}} = -\mathbb{E} \left[ \min\left( \frac{\pi_\theta(a|s)}{\pi_{\text{ref}}(a|s)} \cdot A, \text{clip}\left(\frac{\pi_\theta(a|s)}{\pi_{\text{ref}}(a|s)}, 1-\epsilon, 1+\epsilon\right) \cdot A \right) \right] + \lambda \cdot D_{KL}
$$ $$
In this implementation, an off-policy approach is used ($\pi_\theta = \pi_{\text{ref}}$), and the policy loss simplifies to:
$$
L_{\text{policy}} = -\mathbb{E}[A]
$$
The KL divergence term uses mean squared error approximation: The KL divergence term uses mean squared error approximation:
$$ $$
+67 -32
View File
@@ -2,7 +2,7 @@
### 1. Model Architecture ### 1. Model Architecture
This model uses the Transformer architecture with GQA mechanism (q_head=24, kv_head=4), which saves KV cache memory compared to traditional MHA. The model is built by stacking 32 layers of Transformer blocks, with 1.0 billion parameters. Transformer is an autoregressive model that calculates the relationship between all previous tokens to obtain the probability distribution of the next token. This model uses the Transformer architecture with GQA mechanism (q_head=24, kv_head=4), which saves KV cache memory compared to traditional MHA. The model is built by stacking 24 layers of Transformer blocks, with 1.0 billion parameters. Transformer is an autoregressive model that calculates the relationship between all previous tokens to obtain the probability distribution of the next token.
The model now uses the **AutoModel** base class for flexible loading and saving: The model now uses the **AutoModel** base class for flexible loading and saving:
@@ -48,14 +48,15 @@ flowchart TB
S --> T[+] S --> T[+]
H --> T H --> T
T --> U[RMSNorm] T --> U[RMSNorm]
U --> V[Linear] U --> V["Linear (gate)"]
V --> W[SiLU] U --> W["Linear (up)"]
V --> X[×] V --> X[SiLU]
W --> X X --> Y[×]
X --> Y[Linear] W --> Y
Y --> Z[+] Y --> Z["Linear (down)"]
T --> Z Z --> AA[+]
Z --> AA[x'] T --> AA
AA --> BB[x']
end end
classDef main fill:#e6f3ff,stroke:#0066cc; classDef main fill:#e6f3ff,stroke:#0066cc;
@@ -168,8 +169,6 @@ from astrai.inference import InferenceEngine, GenerationRequest
engine = InferenceEngine( engine = InferenceEngine(
model=model, model=model,
tokenizer=tokenizer, tokenizer=tokenizer,
max_batch_size=8,
max_seq_len=4096,
) )
# Use GenerationRequest with messages format # Use GenerationRequest with messages format
@@ -222,12 +221,11 @@ curl -X POST http://localhost:8000/v1/chat/completions \
| Parameter | Type | Default | Description | | Parameter | Type | Default | Description |
|-----------|------|---------|-------------| |-----------|------|---------|-------------|
| `messages` | List[dict] | Required | Chat messages with role and content | | `messages` | List[dict] | Required | Chat messages with role and content |
| `temperature` | float | 0.8 | Sampling temperature (0.0-2.0) | | `temperature` | float | 1.0 | Sampling temperature (0.0-2.0) |
| `top_p` | float | 0.95 | Nucleus sampling threshold | | `top_p` | float | 1.0 | Nucleus sampling threshold |
| `top_k` | int | 50 | Top-k sampling parameter | | `top_k` | int | 50 | Top-k sampling parameter |
| `max_tokens` | int | 2048 | Maximum tokens to generate | | `max_tokens` | int | 1024 | Maximum tokens to generate |
| `stream` | bool | false | Enable streaming response | | `stream` | bool | false | Enable streaming response |
| `system_prompt` | str | None | System prompt override |
**Response (non-streaming):** **Response (non-streaming):**
```json ```json
@@ -242,7 +240,12 @@ curl -X POST http://localhost:8000/v1/chat/completions \
"message": {"role": "assistant", "content": "Hello! I'm doing well..."}, "message": {"role": "assistant", "content": "Hello! I'm doing well..."},
"finish_reason": "stop" "finish_reason": "stop"
} }
] ],
"usage": {
"prompt_tokens": 20,
"completion_tokens": 15,
"total_tokens": 35
}
} }
``` ```
@@ -262,25 +265,57 @@ curl -X POST http://localhost:8000/v1/chat/completions \
The server uses Server-Sent Events (SSE) with content type `text/event-stream`. The server uses Server-Sent Events (SSE) with content type `text/event-stream`.
### Simple Generation Endpoint ### Anthropic-Compatible Endpoint
For basic text generation without chat format: The server also provides an Anthropic-compatible endpoint at `/v1/messages`:
```bash ```bash
curl -X POST "http://localhost:8000/generate?query=Hello&max_len=1000" \ curl -X POST http://localhost:8000/v1/messages \
-H "Content-Type: application/json"
```
Or with conversation history:
```bash
curl -X POST "http://localhost:8000/generate" \
-H "Content-Type: application/json" \ -H "Content-Type: application/json" \
-d '{ -d '{
"query": "What is AI?", "model": "astrai",
"history": [["Hello", "Hi there!"], ["How are you?", "I'm doing well"]], "system": "You are a helpful assistant.",
"temperature": 0.8, "messages": [{"role": "user", "content": "Hello, how are you?"}],
"max_len": 2048 "max_tokens": 2048
}'
```
Response:
```json
{
"id": "msg_abc123...",
"type": "message",
"role": "assistant",
"model": "astrai",
"content": [{"type": "text", "text": "Hello! I am doing well..."}],
"stop_reason": "end_turn",
"stop_sequence": null,
"usage": {"input_tokens": 20, "output_tokens": 15}
}
```
Streaming:
```bash
curl -X POST http://localhost:8000/v1/messages \
-H "Content-Type: application/json" \
-d '{
"model": "astrai",
"system": "You are a helpful assistant.",
"messages": [{"role": "user", "content": "Write a short poem"}],
"max_tokens": 500,
"stream": true
}'
```
Supports `stop_sequences` for early termination:
```bash
curl -X POST http://localhost:8000/v1/messages \
-H "Content-Type: application/json" \
-d '{
"model": "astrai",
"messages": [{"role": "user", "content": "Write a story"}],
"max_tokens": 500,
"stop_sequences": ["The end", "THE END"]
}' }'
``` ```
@@ -290,10 +325,10 @@ Monitor server and model status:
```bash ```bash
curl http://localhost:8000/health curl http://localhost:8000/health
# {"status": "ok", "model_loaded": true, "engine_ready": true} # {"status": "ok", "model_loaded": true}
curl http://localhost:8000/stats curl http://localhost:8000/stats
# {"requests_total": 10, "tokens_generated": 5000, ...} # {"total_tasks": 10, "total_tokens": 5000, "active_tasks": 1, "waiting_queue": 0}
``` ```
> Document Update Time: 2026-04-09 > Document Update Time: 2026-04-09
+63 -46
View File
@@ -4,70 +4,87 @@
### Basic Parameters ### Basic Parameters
| Parameter | Description | Default Value | | Parameter | Description | Default |
|-----------|-------------|---------------| |-----------|-------------|---------|
| `--train_type` | Training type (seq, sft, dpo, grpo) | required | | `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`) | required |
| `--model_type` | Model type for AutoModel loading (e.g., transformer) | transformer |
| `--data_root_path` | Dataset root directory | required | | `--data_root_path` | Dataset root directory | required |
| `--param_path` | Model parameters or checkpoint path | required | | `--param_path` | Model parameters or checkpoint path | required |
| `--n_epoch` | Total training epochs | 1 | | `--n_epoch` | Total training epochs | 1 |
| `--batch_size` | Batch size | 4 | | `--batch_size` | Batch size | 1 |
| `--accumulation_steps` | Gradient accumulation steps | 1 | | `--accumulation_steps` | Gradient accumulation steps between optimizer steps | 1 |
### Learning Rate Scheduling ### Learning Rate Scheduling
| Parameter | Description | Default Value | | Parameter | Description | Default |
|-----------|-------------|---------------| |-----------|-------------|---------|
| `--warmup_steps` | Warmup steps | 1000 | | `--warmup_steps` | Warmup steps | 1000 |
| `--max_lr` | Maximum learning rate (warmup + cosine decay) | 3e-4 | | `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
| `--max_grad_norm` | Maximum gradient norm | 1.0 | | `--max_grad_norm` | Maximum gradient norm for clipping | 1.0 |
### Checkpoint ### Optimizer (AdamW)
| Parameter | Description | Default Value | | Parameter | Description | Default |
|-----------|-------------|---------------| |-----------|-------------|---------|
| `--ckpt_interval` | Checkpoint save interval (iterations) | 5000 |
| `--ckpt_dir` | Checkpoint save directory | checkpoint |
| `--resume_dir` | Resume training from specified path | - |
### Optimizer Parameters
| Parameter | Description | Default Value |
|-----------|-------------|---------------|
| `--adamw_beta1` | AdamW beta1 | 0.9 | | `--adamw_beta1` | AdamW beta1 | 0.9 |
| `--adamw_beta2` | AdamW beta2 | 0.95 | | `--adamw_beta2` | AdamW beta2 | 0.95 |
| `--adamw_weight_decay` | AdamW weight decay | 0.01 | | `--adamw_weight_decay` | AdamW weight decay | 0.01 |
### Data Loading ### Data Loading
| Parameter | Description | Default Value | | Parameter | Description | Default |
|-----------|-------------|---------------| |-----------|-------------|---------|
| `--random_seed` | Random seed | 3407 | | `--window_size` | Max input sequence length | model config `max_len` |
| `--num_workers` | DataLoader workers | 0 | | `--stride` | Stride for sliding window over sequences | None |
| `--prefetch_factor` | Prefetch factor for dataloader | None | | `--random_seed` | Random seed for reproducibility | 3407 |
| `--pin_memory` | Enable pin_memory | False | | `--num_workers` | DataLoader worker processes | 4 |
| `--no_pin_memory` | Disable pin_memory | - | | `--no_pin_memory` | Disable pin_memory (enabled by default) | (flag) |
### Checkpoint & Resume
| Parameter | Description | Default |
|-----------|-------------|---------|
| `--ckpt_interval` | Iterations between checkpoints | 5000 |
| `--ckpt_dir` | Checkpoint save directory | checkpoint |
| `--start_epoch` | Resume from epoch (0 = from scratch) | 0 |
| `--start_batch` | Resume from batch iteration | 0 |
### Distributed Training ### Distributed Training
| Parameter | Description | Default Value | | Parameter | Description | Default |
|-----------|-------------|---------------| |-----------|-------------|---------|
| `--nprocs` | Number of GPUs | 1 | | `--nprocs` | Number of GPUs / processes | 1 |
| `--device_type` | Device type (cuda/cpu) | cuda | | `--device_type` | Device type | cuda |
### Other Parameters ### Strategy-specific
| Parameter | Description | Default Value | | Parameter | Description | Default | Used by |
|-----------|-------------|---------------| |-----------|-------------|---------|---------|
| `--window_size` | Maximum input sequence length | model config max_len | | `--dpo_beta` | DPO beta value | 0.1 | `dpo` |
| `--stride` | Input sequence stride | - | | `--label_smoothing` | Label smoothing for cross-entropy loss | 0.1 | `seq`, `sft` |
| `--dpo_beta` | DPO beta value | 0.1 | | `--group_size` | GRPO group size | 4 | `grpo` |
| `--grpo_clip_eps` | GRPO clip epsilon | 0.2 | | `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo` |
| `--grpo_kl_coef` | GRPO KL coefficient | 0.01 | | `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo` |
| `--grpo_group_size` | GRPO group size | 4 | | `--grpo_sync_interval` | GRPO ref_model sync interval (steps) | 200 | `grpo` |
| `--label_smoothing` | Label smoothing parameter | 0.1 |
| `--start_epoch` | Starting epoch | 0 | ### Usage Example
| `--start_batch` | Starting batch | 0 |
```bash
python scripts/tools/train.py \
--train_type seq \
--data_root_path /path/to/dataset \
--param_path /path/to/model \
--n_epoch 3 \
--batch_size 4 \
--accumulation_steps 8 \
--max_lr 3e-4 \
--warmup_steps 2000 \
--max_grad_norm 1.0 \
--ckpt_interval 5000 \
--ckpt_dir ./checkpoints \
--num_workers 4 \
--nprocs 1 \
--device_type cuda
```
--- ---
@@ -89,14 +106,14 @@
```python ```python
import torch import torch
from astrai.model import AutoModel from astrai.model import AutoModel
from astrai.tokenize import Tokenizer from astrai.tokenize import AutoTokenizer
from astrai.inference import InferenceEngine, GenerationRequest from astrai.inference import InferenceEngine, GenerationRequest
# Load model using AutoModel # Load model using AutoModel
model = AutoModel.from_pretrained("your_model_dir") model = AutoModel.from_pretrained("your_model_dir")
# Load tokenizer # Load tokenizer
tokenizer = Tokenizer("your_model_dir") tokenizer = AutoTokenizer.from_pretrained("your_model_dir")
# Create engine with separate model and tokenizer # Create engine with separate model and tokenizer
engine = InferenceEngine( engine = InferenceEngine(
+1 -1
View File
@@ -1,4 +1,4 @@
__version__ = "1.3.3" __version__ = "1.3.4"
__author__ = "ViperEkura" __author__ = "ViperEkura"
from astrai.config import ( from astrai.config import (
+1 -4
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Callable, List, Optional from typing import Callable, Optional
import torch.nn as nn import torch.nn as nn
from torch.optim import Optimizer from torch.optim import Optimizer
@@ -74,9 +74,6 @@ class TrainConfig:
) )
# others # others
device_ids: Optional[List[int]] = field(
default=None, metadata={"help": "Device ids for distributed training."}
)
device_type: str = field( device_type: str = field(
default="cuda", metadata={"help": "Device type for distributed training."} default="cuda", metadata={"help": "Device type for distributed training."}
) )
+28 -7
View File
@@ -1,25 +1,46 @@
"""Inference module for continuous batching.""" """Inference module for continuous batching.
Layers:
- engine.py: Facade (InferenceEngine), Value Object (GenerationParams, GenerationRequest)
- scheduler.py: Continuous-batching loop, Task state machine, TaskStatus enum
- cache.py: PagedCache (page-table-indirected KV cache with alloc/free)
- sampling.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy)
- server.py: FastAPI HTTP server (OpenAI-compatible endpoints)
"""
from astrai.inference.engine import ( from astrai.inference.engine import (
GenerationParams,
GenerationRequest, GenerationRequest,
InferenceEngine, InferenceEngine,
) )
from astrai.inference.sampling import (
BaseSamplingStrategy,
SamplingPipeline,
TemperatureStrategy,
TopKStrategy,
TopPStrategy,
sample,
)
from astrai.inference.scheduler import ( from astrai.inference.scheduler import (
InferenceScheduler, InferenceScheduler,
Task, Task,
TaskStatus, TaskStatus,
apply_sampling_strategies,
) )
__all__ = [ __all__ = [
# Engine # Engine / Requests
"InferenceEngine", "InferenceEngine",
"GenerationRequest",
"GenerationParams",
# Scheduler # Scheduler
"InferenceScheduler", "InferenceScheduler",
"Task", "Task",
"TaskStatus", "TaskStatus",
# Request # Sampling (Strategy pattern)
"GenerationRequest", "sample",
# Sampling "BaseSamplingStrategy",
"apply_sampling_strategies", "TemperatureStrategy",
"TopKStrategy",
"TopPStrategy",
"SamplingPipeline",
] ]
+174
View File
@@ -0,0 +1,174 @@
"""Page-based KV cache with page-table-indirected read/write.
Provides:
- PagedCache: paged KV cache combining page pool and tensor storage.
"""
from typing import Dict, List, Tuple
import torch
from torch import Tensor
STOP = object()
def page_hash(token_ids: List[int], page_idx: int, page_size: int) -> int:
start = page_idx * page_size
end = min(start + page_size, len(token_ids))
h = 0
for i in range(start, end):
h = (h * 31 + token_ids[i]) & 0xFFFFFFFFFFFFFFFF
return h
class PagedCache:
"""Paged KV cache with page-table-indirected read/write.
Combines:
- Page pool (ref-counted alloc/free via bitmask)
- KV tensor storage (k_cache, v_cache)
- Prefix-cache hash lookup (page_content_hash -> physical_page_idx)
Call :meth:`bind` to obtain a batch view for the attention layers.
"""
def __init__(
self,
n_layers: int,
n_pages: int,
page_size: int,
n_kv_heads: int,
head_dim: int,
device: torch.device,
dtype: torch.dtype,
):
self.page_size = page_size
self._free_mask = (1 << n_pages) - 1
self._refs: List[int] = [0] * n_pages
self.k_cache = torch.empty(
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
device=device,
dtype=dtype,
)
self.v_cache = torch.empty(
(n_layers, n_pages, page_size, n_kv_heads, head_dim),
device=device,
dtype=dtype,
)
self._page_to_hash: Dict[int, int] = {}
self._hash_to_page: Dict[int, int] = {}
def record_page(
self, page_idx: int, token_ids: List[int], logical_page_idx: int
) -> None:
h = page_hash(token_ids, logical_page_idx, self.page_size)
old_h = self._page_to_hash.pop(page_idx, None)
if old_h is not None:
self._hash_to_page.pop(old_h, None)
self._page_to_hash[page_idx] = h
self._hash_to_page[h] = page_idx
def lookup_prefix(self, token_ids: List[int]) -> List[int]:
full_pages = len(token_ids) // self.page_size
hits: List[int] = []
for i in range(full_pages):
h = page_hash(token_ids, i, self.page_size)
p = self._hash_to_page.get(h)
if p is None:
break
hits.append(p)
return hits
def inc_ref(self, idx: int) -> None:
self._refs[idx] += 1
def alloc(self) -> int:
lsb = self._free_mask & -self._free_mask
if lsb == 0:
return -1
idx = lsb.bit_length() - 1
self._free_mask ^= lsb
self._refs[idx] = 1
return idx
def alloc_n(self, n: int) -> List[int]:
pages = [self.alloc() for _ in range(n)]
if any(p < 0 for p in pages):
for p in pages:
if p >= 0:
self.free(p)
return []
return pages
def free(self, idx: int) -> None:
self._refs[idx] -= 1
if self._refs[idx] == 0:
self._free_mask |= 1 << idx
h = self._page_to_hash.pop(idx, None)
if h is not None:
self._hash_to_page.pop(h, None)
def bind(self, page_table: Tensor, total_len: int = 0) -> "CacheView":
return CacheView(self, page_table, total_len)
def write(
self, layer_id: int, page_table: Tensor, start_pos: int, k: Tensor, v: Tensor
) -> None:
seq_len = k.size(1)
if seq_len == 0:
return
page_size = self.page_size
written = 0
first_page = start_pos // page_size
last_page = (start_pos + seq_len - 1) // page_size
for pi in range(first_page, last_page + 1):
phys_pages = page_table[:, pi]
page_start = pi * page_size
write_start = max(page_start, start_pos)
write_end = min(page_start + page_size, start_pos + seq_len)
offset = write_start - page_start
chunk = write_end - write_start
self.k_cache[layer_id, phys_pages, offset : offset + chunk] = k[
:, written : written + chunk
]
self.v_cache[layer_id, phys_pages, offset : offset + chunk] = v[
:, written : written + chunk
]
written += chunk
def gather(self, layer_id: int, page_table: Tensor) -> Tuple[Tensor, Tensor]:
k_parts, v_parts = [], []
for pi in range(page_table.size(1)):
phys_pages = page_table[:, pi]
if not (phys_pages >= 0).any():
break
k_parts.append(self.k_cache[layer_id, phys_pages])
v_parts.append(self.v_cache[layer_id, phys_pages])
k = torch.cat(k_parts, dim=1)
v = torch.cat(v_parts, dim=1)
return k, v
class CacheView:
"""Per-batch view that bundles PagedCache + page_table + total_len.
Attention layers receive this as ``paged_cache`` and only see
``write()`` / ``gather()``, never raw page tables or length params.
"""
__slots__ = ("_cache", "_page_table", "_total_len")
def __init__(self, cache: PagedCache, page_table: Tensor, total_len: int = 0):
self._cache = cache
self._page_table = page_table
self._total_len = total_len
def write(self, layer_id: int, start_pos: int, k: Tensor, v: Tensor) -> None:
self._cache.write(layer_id, self._page_table, start_pos, k, v)
def gather(self, layer_id: int) -> Tuple[Tensor, Tensor]:
k, v = self._cache.gather(layer_id, self._page_table)
if self._total_len:
k = k[:, : self._total_len]
v = v[:, : self._total_len]
return k, v
+282 -114
View File
@@ -1,21 +1,42 @@
"""Unified inference engine.""" """Unified inference engine for continuous batching.
Layers:
- GenerationParams: Immutable value object for sampling parameters.
- GenerationRequest: User-facing request DTO with validation.
- _Result: Thread-safe token accumulator (Observer pattern).
- InferenceEngine: Facade over InferenceScheduler + async wrapper.
"""
import asyncio
import gc import gc
import logging
import threading import threading
from typing import Any, Dict, Generator, List, Optional, Union from dataclasses import dataclass
from typing import Any, AsyncGenerator, Dict, Generator, List, Optional, Union
import torch import torch
import torch.nn as nn import torch.nn as nn
from astrai.inference.cache import STOP
from astrai.inference.scheduler import InferenceScheduler from astrai.inference.scheduler import InferenceScheduler
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
logger = logging.getLogger(__name__)
@dataclass(frozen=True)
class GenerationParams:
"""Immutable value object for sampling hyperparameters."""
top_k: int = 50
top_p: float = 1.0
temperature: float = 1.0
max_tokens: int = 1024
class GenerationRequest: class GenerationRequest:
"""Request parameters for text generation.""" """Request parameters for text generation.
Encapsulates messages, sampling parameters (via GenerationParams),
and streaming preference for a single generation request.
"""
def __init__( def __init__(
self, self,
@@ -26,17 +47,44 @@ class GenerationRequest:
max_len: int = 1024, max_len: int = 1024,
stream: bool = False, stream: bool = False,
): ):
self.messages = messages """Initializes a generation request.
self.top_k = top_k
self.top_p = top_p
self.temperature = temperature
self.max_len = max_len
self.stream = stream
Args:
messages: Conversation history as list of {"role": ..., "content": ...}.
top_k: Top-k sampling count (0 disables).
top_p: Nucleus sampling probability threshold.
temperature: Sampling temperature.
max_len: Maximum tokens to generate.
stream: Whether to return output as a token stream.
"""
self.messages = messages
self.params = GenerationParams(
top_k=top_k,
top_p=top_p,
temperature=temperature,
max_tokens=max_len,
)
self.stream = stream
self._validate() self._validate()
@property
def top_k(self) -> int:
return self.params.top_k
@property
def top_p(self) -> float:
return self.params.top_p
@property
def temperature(self) -> float:
return self.params.temperature
@property
def max_len(self) -> int:
return self.params.max_tokens
def _validate(self): def _validate(self):
"""Validate request parameters.""" """Validates sampling parameter ranges."""
if not (isinstance(self.top_k, int) and self.top_k >= 0): if not (isinstance(self.top_k, int) and self.top_k >= 0):
raise ValueError("top_k must be a non-negative integer") raise ValueError("top_k must be a non-negative integer")
if not (0.0 <= self.top_p <= 1.0): if not (0.0 <= self.top_p <= 1.0):
@@ -46,50 +94,101 @@ class GenerationRequest:
class _Result: class _Result:
"""Unified result holder for streaming/non-streaming modes.""" """Thread-safe token accumulator for streaming and non-streaming modes.
def __init__(self, count: int = 1, stream: bool = False): Supports multiple concurrent generation tasks with per-index result tracking.
self._stream = stream Uses a threading.Condition for efficient completion notification
self._lock = threading.Lock() and a threading.Event for streaming wakeup.
"""
def __init__(self, count: int = 1):
"""Initializes the accumulator.
Args:
count: Number of concurrent generation tasks to track.
"""
self._cond = threading.Condition()
self._event = threading.Event() self._event = threading.Event()
self.tokens: List[str] = [] self.tokens: List[str] = []
self.results: List[str] = [""] * count if count > 1 else [""] self.results: List[str] = [""] * count
self.done_flags: List[bool] = [False] * count self._done: List[bool] = [False] * count
self._completed_count = 0 self._completed = 0
self._total = count
def append(self, token: str, idx: int = 0): def append(self, token: str, idx: int = 0):
with self._lock: """Appends a token to the result buffer.
if self._stream:
In non-streaming mode, tokens are concatenated into results[idx].
The sentinel STOP marks a task as complete.
Args:
token: The decoded token string, or STOP sentinel.
idx: Index of the generation task this token belongs to.
"""
with self._cond:
self.tokens.append(token) self.tokens.append(token)
else: if token is not STOP:
if token == "[DONE]":
if not self.done_flags[idx]:
self.done_flags[idx] = True
self._completed_count += 1
if self._completed_count == len(self.results):
self._event.set()
else:
self.results[idx] += token self.results[idx] += token
else:
if not self._done[idx]:
self._done[idx] = True
self._completed += 1
self._cond.notify_all()
self._event.set() self._event.set()
def pop_all(self) -> List[str]: def pop_all(self) -> List[str]:
with self._lock: """Returns and clears all accumulated tokens.
tokens = self.tokens.copy()
self.tokens.clear()
if not tokens:
self._event.clear()
return tokens
def wait(self, timeout: float = None) -> bool: Returns:
List of token strings since the last call.
"""
with self._cond:
out = self.tokens.copy()
self.tokens.clear()
if not out:
self._event.clear()
return out
def wait(self, timeout: Optional[float] = None) -> bool:
"""Blocks until new tokens arrive or the timeout expires.
Args:
timeout: Maximum wait time in seconds (None = infinite).
Returns:
True if the event was set (new data available), False on timeout.
"""
return self._event.wait(timeout=timeout) return self._event.wait(timeout=timeout)
def wait_completion(self) -> None:
"""Blocks until all tasks complete (non-streaming).
Uses a Condition to sleep efficiently instead of busy-waiting.
The calling thread is parked until a STOP signal arrives.
"""
with self._cond:
self._cond.wait_for(lambda: self._completed >= self._total)
def get_results(self) -> List[str]: def get_results(self) -> List[str]:
with self._lock: """Returns all accumulated results for non-streaming mode.
Returns:
List of complete generated strings, one per task index.
"""
with self._cond:
return self.results.copy() return self.results.copy()
class InferenceEngine: class InferenceEngine:
"""Unified inference engine for continuous batching.""" """Unified inference engine backed by continuous-batching scheduler.
Usage:
with InferenceEngine(model, tokenizer) as engine:
for token in engine.generate("hello", stream=True):
print(token, end="")
text = engine.generate("hello")
"""
def __init__( def __init__(
self, self,
@@ -97,55 +196,37 @@ class InferenceEngine:
tokenizer: AutoTokenizer, tokenizer: AutoTokenizer,
max_batch_size: int = 1, max_batch_size: int = 1,
max_seq_len: Optional[int] = None, max_seq_len: Optional[int] = None,
max_prefix_len: int = 512, max_prompt_len: int = 2048,
cache_capacity: int = 1000, page_size: int = 128,
): ):
""" """Initializes the inference engine.
Initialize inference engine with separate model and tokenizer.
Args: Args:
model: The language model for inference (nn.Module, e.g., Transformer) model: The model instance.
tokenizer: The tokenizer for encoding/decoding text tokenizer: The tokenizer instance.
config: Model configuration max_batch_size: Maximum number of concurrent tasks.
max_batch_size: Maximum batch size for continuous batching max_seq_len: Maximum sequence length.
max_seq_len: Maximum sequence length (defaults to config.max_len) max_prompt_len: Maximum prompt tokens.
max_prefix_len: Maximum prefix length for cache (default: 512) compile: Whether to compile the model with torch.compile.
cache_capacity: Maximum number of cached prefixes (default: 1000) page_size: Number of tokens per KV cache page.
""" """
self.model = model self.model = model
self.tokenizer = tokenizer self.tokenizer = tokenizer
# Get device and dtype from model parameters
try:
first_param = next(model.parameters())
device = first_param.device
dtype = first_param.dtype
except StopIteration:
# Model has no parameters, use default device/dtype
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
dtype = torch.float32
self.scheduler = InferenceScheduler( self.scheduler = InferenceScheduler(
model=self.model, model=self.model,
tokenizer=self.tokenizer, tokenizer=self.tokenizer,
max_batch_size=max_batch_size, max_batch_size=max_batch_size,
max_seq_len=max_seq_len, max_seq_len=max_seq_len,
max_prefix_len=max_prefix_len, max_prompt_len=max_prompt_len,
cache_capacity=cache_capacity, page_size=page_size,
device=device,
dtype=dtype,
) )
self.kv_cache = self.scheduler.kv_cache
self.seq_mask = self.scheduler.seq_mask
self.scheduler.start() self.scheduler.start()
def __enter__(self): def __enter__(self):
return self return self
def __exit__(self, exc_type, exc_val, exc_tb): def __exit__(self, exc_type, exc_val, exc_tb):
"""Handle exceptions on exit."""
self.shutdown() self.shutdown()
return False return False
@@ -157,46 +238,106 @@ class InferenceEngine:
temperature: float = 1.0, temperature: float = 1.0,
top_p: float = 1.0, top_p: float = 1.0,
top_k: int = 50, top_k: int = 50,
abort_on_exception: bool = True,
) -> Union[Generator[str, None, None], str, List[str]]: ) -> Union[Generator[str, None, None], str, List[str]]:
"""Unified generation interface. """Generates text from a prompt.
Args: Args:
abort_on_exception: If True, abort the generation when consumer prompt: Single string or list of strings for batch generation.
stops iterating (GeneratorExit/StopIteration). Default: True. stream: If True, returns a generator yielding tokens one by one.
max_tokens: Maximum number of tokens to generate.
temperature: Sampling temperature.
top_p: Nucleus sampling probability threshold.
top_k: Top-k sampling count (0 disables).
Returns:
Generator (stream=True), single string (non-stream, single prompt),
or list of strings (non-stream, batch prompts).
""" """
is_batch = isinstance(prompt, list) is_batch = isinstance(prompt, list)
prompts = prompt if is_batch else [prompt] prompts = prompt if is_batch else [prompt]
if stream: if stream:
return self._generate_streaming( return self._generate_streaming(
prompts, prompts, is_batch, max_tokens, temperature, top_p, top_k
is_batch,
max_tokens,
temperature,
top_p,
top_k,
abort_on_exception,
) )
else: else:
return self._generate_non_streaming( return self._generate_non_streaming(
prompts, is_batch, max_tokens, temperature, top_p, top_k prompts, is_batch, max_tokens, temperature, top_p, top_k
) )
def generate_async(
self,
prompt: str,
max_tokens: int = 1024,
temperature: float = 1.0,
top_p: float = 1.0,
top_k: int = 50,
) -> AsyncGenerator[str, None]:
"""Async streaming generator that does not block the event loop.
Runs the synchronous generator in a background thread pool executor,
yielding tokens to the async consumer as they arrive.
Args:
prompt: Input text to generate from.
max_tokens: Maximum tokens to generate.
temperature: Sampling temperature.
top_p: Nucleus sampling threshold.
top_k: Top-k sampling count.
Yields:
Decoded token strings as they are generated.
"""
sync_gen = self._generate_streaming(
[prompt], False, max_tokens, temperature, top_p, top_k
)
async def _agen():
loop = asyncio.get_event_loop()
while True:
token = await loop.run_in_executor(None, self._next_token, sync_gen)
if token is None:
break
yield token
return _agen()
@staticmethod
def _next_token(gen: Generator) -> Optional[str]:
"""Retrieves the next token from a synchronous generator.
Args:
gen: A synchronous generator yielding token strings.
Returns:
The next token, or None if the generator is exhausted.
"""
try:
return next(gen)
except StopIteration:
return None
def generate_with_request( def generate_with_request(
self, request: GenerationRequest self, request: GenerationRequest
) -> Union[Generator[str, None, None], str, List[str]]: ) -> Union[Generator[str, None, None], str, List[str]]:
"""Generate with GenerationRequest object.""" """Generates text from a structured GenerationRequest.
# Use tokenizer's chat template with messages
prompt = self.tokenizer.apply_chat_template(request.messages, tokenize=False)
Applies the chat template to the request's messages before generation.
Args:
request: A GenerationRequest with messages and parameters.
Returns:
Generator, string, or list of strings (see generate()).
"""
prompt = self.tokenizer.apply_chat_template(request.messages, tokenize=False)
return self.generate( return self.generate(
prompt=prompt, prompt=prompt,
stream=request.stream, stream=request.stream,
max_tokens=request.max_len, max_tokens=request.params.max_tokens,
temperature=request.temperature, temperature=request.params.temperature,
top_p=request.top_p, top_p=request.params.top_p,
top_k=request.top_k, top_k=request.params.top_k,
) )
def _generate_streaming( def _generate_streaming(
@@ -207,18 +348,27 @@ class InferenceEngine:
temperature: float, temperature: float,
top_p: float, top_p: float,
top_k: int, top_k: int,
abort_on_exception: bool = True, ) -> Generator[str, None, None]:
) -> Union[Generator[str, None, None], List[Generator[str, None, None]]]: """Internal streaming generator.
"""Generate with streaming output.
Polls the _Result accumulator in a loop, yielding tokens as they arrive.
Cleans up the scheduler task on GeneratorExit.
Args: Args:
abort_on_exception: If True, abort the task when generator is prompts: List of prompts (only first is used; batch not yet supported).
stopped early by consumer (GeneratorExit/StopIteration). is_batch: If True, raises NotImplementedError.
max_tokens: Maximum tokens to generate.
temperature: Sampling temperature.
top_p: Nucleus sampling threshold.
top_k: Top-k sampling count.
Yields:
Decoded token strings.
""" """
if is_batch: if is_batch:
raise NotImplementedError("Batch streaming is not implemented yet") raise NotImplementedError("Batch streaming not yet supported")
result = _Result(stream=True) result = _Result()
task_id = self.scheduler.add_task( task_id = self.scheduler.add_task(
prompt=prompts[0], prompt=prompts[0],
@@ -226,7 +376,7 @@ class InferenceEngine:
temperature=temperature, temperature=temperature,
top_p=top_p, top_p=top_p,
top_k=top_k, top_k=top_k,
stream_callback=result.append, stream_callback=lambda tok: result.append(tok, 0),
) )
def gen(): def gen():
@@ -234,17 +384,14 @@ class InferenceEngine:
while True: while True:
tokens = result.pop_all() tokens = result.pop_all()
for token in tokens: for token in tokens:
if token == "[DONE]": if token is STOP:
return return
yield token yield token
result.wait(timeout=0.05) if not result.wait(timeout=0.05):
except Exception: pass
# Consumer stopped iterating - abort the task finally:
if abort_on_exception:
self.scheduler.remove_task(task_id) self.scheduler.remove_task(task_id)
raise
gen.task_id = task_id
return gen() return gen()
def _generate_non_streaming( def _generate_non_streaming(
@@ -256,36 +403,57 @@ class InferenceEngine:
top_p: float, top_p: float,
top_k: int, top_k: int,
) -> Union[str, List[str]]: ) -> Union[str, List[str]]:
"""Generate without streaming.""" """Internal non-streaming generator.
Submits all prompts to the scheduler and waits for all to complete.
Args:
prompts: List of prompt strings.
is_batch: Whether multiple prompts were provided.
max_tokens: Maximum tokens to generate.
temperature: Sampling temperature.
top_p: Nucleus sampling threshold.
top_k: Top-k sampling count.
Returns:
Single string for one prompt, list of strings for batch.
"""
result = _Result(count=len(prompts)) result = _Result(count=len(prompts))
task_ids = []
for i, p in enumerate(prompts): for i, p in enumerate(prompts):
# Create closure to capture current index value using factory function
def make_callback(idx):
def callback(token):
result.append(idx, token)
return callback def make_cb(idx):
return lambda tok: result.append(tok, idx)
self.scheduler.add_task( task_id = self.scheduler.add_task(
prompt=p, prompt=p,
max_tokens=max_tokens, max_tokens=max_tokens,
temperature=temperature, temperature=temperature,
top_p=top_p, top_p=top_p,
top_k=top_k, top_k=top_k,
stream_callback=make_callback(i), stream_callback=make_cb(i),
) )
task_ids.append(task_id)
result.wait() result.wait_completion()
results = result.get_results()
return results if is_batch else results[0] for task_id in task_ids:
self.scheduler.remove_task(task_id)
res = result.get_results()
return res if is_batch else res[0]
def get_stats(self) -> Dict[str, Any]: def get_stats(self) -> Dict[str, Any]:
"""Get engine statistics.""" """Returns current engine statistics.
Returns:
Dict with total_tasks, total_tokens, active_tasks, waiting_queue.
"""
return self.scheduler.get_stats() return self.scheduler.get_stats()
def shutdown(self) -> None: def shutdown(self) -> None:
"""Shutdown the engine and release all resources.""" """Shuts down the engine, stops the scheduler, and frees GPU memory."""
self.scheduler.stop() self.scheduler.stop()
if torch.cuda.is_available(): if torch.cuda.is_available():
torch.cuda.empty_cache() torch.cuda.empty_cache()
+178
View File
@@ -0,0 +1,178 @@
"""Composable sampling strategies for logit transformation.
Implements the Strategy pattern: each sampling technique
(temperature, top-k, top-p) is a pluggable strategy that
can be composed into a pipeline.
All strategies accept both scalar and per-sample tensor
parameters, so a single pipeline works for any batch size.
"""
from abc import ABC, abstractmethod
from typing import List, Union
import torch
from torch import Tensor
class BaseSamplingStrategy(ABC):
"""Abstract base for a logit transformation strategy."""
@abstractmethod
def apply(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
"""Applies the strategy to logits.
Args:
logits: Raw logits tensor (batch, vocab_size).
filter_value: Value assigned to filtered-out positions.
Returns:
Transformed logits tensor.
"""
class TemperatureStrategy(BaseSamplingStrategy):
"""Divides logits by temperature to control randomness.
Args:
temperature: Scalar or ``[batch]`` tensor.
"""
def __init__(self, temperature: Union[float, Tensor] = 1.0):
self.temperature = temperature
def apply(self, logits, filter_value=-float("inf")):
t = self.temperature
if isinstance(t, Tensor):
if (t != 1.0).any():
logits = logits / t.to(logits.device, non_blocking=True).view(-1, 1)
elif t != 1.0:
logits = logits / t
return logits
class TopKStrategy(BaseSamplingStrategy):
"""Keeps only the top-k logits, setting the rest to filter_value.
Args:
top_k: Scalar or ``[batch]`` tensor (0 disables).
"""
def __init__(self, top_k: Union[int, Tensor] = 0):
self.top_k = top_k
def apply(self, logits, filter_value=-float("inf")):
tk = self.top_k
if isinstance(tk, Tensor):
max_k = int(tk.max().item())
if max_k <= 0:
return logits
k = min(max_k, logits.size(-1))
elif tk > 0:
k = min(tk, logits.size(-1))
else:
return logits
thresholds = torch.topk(logits, k, dim=-1)[0][..., -1:]
logits[logits < thresholds] = filter_value
return logits
class TopPStrategy(BaseSamplingStrategy):
"""Nucleus (top-p) filtering: keeps the smallest set of tokens whose
cumulative probability exceeds top_p.
Args:
top_p: Scalar or ``[batch]`` tensor (1.0 disables).
"""
def __init__(self, top_p: Union[float, Tensor] = 1.0):
self.top_p = top_p
def _apply(self, logits, top_p, filter_value):
sorted_logits, sorted_indices = torch.sort(logits, descending=True, dim=-1)
cum_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
remove = cum_probs > top_p
remove[..., 1:] = remove[..., :-1].clone()
remove[..., 0] = False
mask = torch.zeros_like(logits, dtype=torch.bool)
mask.scatter_(1, sorted_indices, remove)
logits[mask] = filter_value
return logits
def apply(self, logits, filter_value=-float("inf")):
tp = self.top_p
if isinstance(tp, Tensor):
tp = tp.to(logits.device, non_blocking=True)
if (tp < 1.0).any():
logits = self._apply(logits, tp.view(-1, 1), filter_value)
elif tp < 1.0:
logits = self._apply(logits, tp, filter_value)
return logits
class SamplingPipeline(BaseSamplingStrategy):
"""Composes multiple sampling strategies into a single transformation.
Strategies are applied sequentially in the order they are provided,
matching the original temperature -> top-k -> top-p ordering.
Usage::
pipeline = SamplingPipeline([
TemperatureStrategy(0.8),
TopKStrategy(50),
TopPStrategy(0.95),
])
logits = pipeline.apply(logits)
token = pipeline.sample(logits) # softmax + multinomial
"""
def __init__(self, strategies: List[BaseSamplingStrategy]):
self.strategies = strategies
def apply(self, logits, filter_value=-float("inf")):
for strategy in self.strategies:
logits = strategy.apply(logits, filter_value)
return logits
@torch.no_grad()
def sample(self, logits: Tensor, filter_value: float = -float("inf")) -> Tensor:
"""Apply strategies then sample (softmax + multinomial).
Args:
logits: Raw logits ``[batch, vocab_size]``.
Returns:
Sampled token IDs ``[batch]``.
"""
return torch.multinomial(
torch.softmax(self.apply(logits, filter_value), dim=-1),
num_samples=1,
).squeeze(-1)
@torch.inference_mode()
def sample(
logits: Tensor,
temperature: Union[float, Tensor] = 1.0,
top_k: Union[int, Tensor] = 0,
top_p: Union[float, Tensor] = 1.0,
filter_value: float = -float("inf"),
) -> Tensor:
"""Apply sampling strategies then sample (softmax + multinomial).
Shortcut for ``SamplingPipeline(...).sample(logits)``.
Args:
logits: Raw logits ``[batch, vocab_size]``.
Returns:
Sampled token IDs ``[batch]``.
"""
return SamplingPipeline(
[
TemperatureStrategy(temperature),
TopKStrategy(top_k),
TopPStrategy(top_p),
]
).sample(logits, filter_value)
+208 -434
View File
@@ -1,148 +1,25 @@
"""Inference scheduler for continuous batching.""" """Inference scheduler for single-GPU continuous batching with paged KV cache."""
import logging
import threading import threading
import time import time
import uuid import uuid
from enum import Enum
from typing import Any, Callable, Dict, List, Optional, Tuple from typing import Any, Callable, Dict, List, Optional, Tuple
import torch import torch
from torch import Tensor from torch import Tensor
from astrai.inference.cache import STOP, PagedCache
from astrai.inference.sampling import sample
from astrai.model.automodel import AutoModel from astrai.model.automodel import AutoModel
from astrai.tokenize import AutoTokenizer from astrai.tokenize.tokenizer import AutoTokenizer
logger = logging.getLogger(__name__)
class RadixNode: class TaskStatus(Enum):
"""Radix tree node for prefix cache.""" """Task states in the continuous batching lifecycle."""
def __init__(self):
self.children: Dict[int, "RadixNode"] = {} # token_id -> child node
self.hash: Optional[int] = None # 64-bit hash of the prefix
self.slot: int = -1 # KV Cache slot, valid only for leaf nodes
self.ref_count: int = 0 # number of tasks referencing this prefix
self.last_access: float = 0.0 # timestamp for LRU
self.token_sequence: list = [] # full token sequence from root to this node
class PrefixCacheManager:
"""Prefix cache manager using Radix tree with LRU eviction."""
def __init__(self, max_capacity: int = 1000, base: int = 131, mod: int = 10**9 + 7):
self.root = RadixNode()
self.base = base
self.mod = mod
self.max_capacity = max_capacity
self.lru: List[Tuple[float, RadixNode]] = [] # (timestamp, node) for LRU
def insert(self, token_ids: Tuple[int, ...], slot: int) -> None:
"""Insert a prefix, increase ref_count if already exists, otherwise create new node."""
node = self.root
path = []
h = 0
for i, token_id in enumerate(token_ids):
if token_id not in node.children:
node.children[token_id] = RadixNode()
node = node.children[token_id]
h = (h * self.base + token_id) % self.mod
node.hash = h
path.append(token_id)
node.token_sequence = list(
path
) # store full sequence for exact verification
# Leaf node: set slot and increase ref_count
if node.slot == -1:
node.slot = slot
node.ref_count += 1
node.last_access = time.time()
self._update_lru(node)
self._evict_if_needed()
def find_longest_prefix(self, token_ids: List[int]) -> Optional[Tuple[int, int]]:
"""Find longest matching prefix, return (prefix_len, slot).
During traversal, compute hash per token and compare with node hash.
If hash matches, perform full token sequence verification to avoid
hash collision errors.
"""
node = self.root
best_len = 0
best_slot = -1
h = 0
for i, token_id in enumerate(token_ids):
if token_id not in node.children:
break
node = node.children[token_id]
h = (h * self.base + token_id) % self.mod
if node.hash == h: # hash matches
# Exact verification: compare full token sequence
if node.token_sequence == token_ids[: i + 1]:
best_len = i + 1
best_slot = node.slot
node.last_access = time.time()
self._update_lru(node)
if best_len > 0:
return (best_len, best_slot)
return None
def release(self, token_ids: Tuple[int, ...]) -> None:
"""Release reference to a prefix, decrease ref_count. If zero, mark as evictable."""
node = self.root
for token_id in token_ids:
if token_id not in node.children:
return
node = node.children[token_id]
if node.ref_count > 0:
node.ref_count -= 1
if node.ref_count == 0:
node.slot = -1 # slot can be reused
def _update_lru(self, node: RadixNode) -> None:
"""Update LRU list, move node to most recently used position."""
self.lru = [(ts, n) for (ts, n) in self.lru if n is not node]
self.lru.append((node.last_access, node))
def _evict_if_needed(self) -> None:
"""If cache entries exceed capacity, evict least recently used leaf nodes (ref_count must be 0)."""
if len(self.lru) <= self.max_capacity:
return
# Sort by timestamp
self.lru.sort(key=lambda x: x[0])
for ts, node in self.lru:
if node.ref_count == 0:
# Remove leaf node from tree (need to recursively delete empty branches)
self._remove_node(node)
self.lru.remove((ts, node))
if len(self.lru) <= self.max_capacity:
break
def _remove_node(
self,
node: RadixNode,
parent: Optional[RadixNode] = None,
child_key: Optional[int] = None,
) -> None:
"""Remove node from tree, including empty parent nodes."""
# First, recursively remove all children
for child_key, child_node in list(node.children.items()):
self._remove_node(child_node, node, child_key)
# Clear the node's leaf properties
node.slot = -1
node.hash = None
node.token_sequence = []
node.children.clear()
# If this node has no children and has a parent, remove the reference from parent
if parent is not None and child_key is not None and len(node.children) == 0:
if child_key in parent.children:
del parent.children[child_key]
class TaskStatus:
"""Task state for continuous batching."""
PENDING = "pending" PENDING = "pending"
RUNNING = "running" RUNNING = "running"
@@ -151,7 +28,7 @@ class TaskStatus:
class Task: class Task:
"""Individual task for continuous batching.""" """Represents a single generation request with paged KV cache tracking."""
def __init__( def __init__(
self, self,
@@ -174,60 +51,35 @@ class Task:
self.output_ids: List[int] = [] self.output_ids: List[int] = []
self.input_tokens: int = 0 self.input_tokens: int = 0
self.output_tokens: int = 0 self.output_tokens: int = 0
self.slot: int = -1 self.page_table: List[int] = []
self.prefix_len: int = 0 # prefix cache matched length self.n_pages: int = 0
self._prefix_cached_tokens: int = 0
self.arrival_time = time.time() self.arrival_time = time.time()
self.finish_time: Optional[float] = None self.finish_time: Optional[float] = None
self.stream_callback = stream_callback self.stream_callback = stream_callback
self._pages_freed: bool = False
@property
def next_pos(self) -> int:
return self.input_tokens + len(self.output_ids)
def is_finished(self, stop_ids: List[int]) -> bool: def is_finished(self, stop_ids: List[int]) -> bool:
"""Check if task is finished.""" if self.output_tokens >= self.max_tokens:
return ( return True
bool(self.output_ids and self.output_ids[-1] in stop_ids) if self.output_ids and self.output_ids[-1] in stop_ids:
or self.output_tokens >= self.max_tokens return True
) return False
def apply_sampling_strategies(
logits: Tensor,
temperature: float,
top_k: int,
top_p: float,
filter_value: float = -float("inf"),
) -> Tensor:
"""Apply sampling strategies to the logits tensor."""
# Clone logits to avoid inplace updates on inference tensor
logits = logits.clone()
if temperature != 1.0:
logits = logits / temperature
if top_k > 0:
top_k = min(top_k, logits.size(-1))
indices_to_remove = logits < torch.topk(logits, top_k, dim=-1)[0][..., -1, None]
logits[indices_to_remove] = filter_value
if top_p < 1.0:
sorted_logits, sorted_indices = torch.sort(logits, descending=True, dim=-1)
cumulative_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
sorted_indices_to_remove = cumulative_probs > top_p
sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
sorted_indices_to_remove[..., 0] = 0
indices_to_remove = torch.zeros_like(logits, dtype=torch.bool)
indices_to_remove.scatter_(
dim=1, index=sorted_indices, src=sorted_indices_to_remove
)
logits[indices_to_remove] = filter_value
return logits
class InferenceScheduler: class InferenceScheduler:
"""Inference scheduler with continuous batching support.""" """Continuous batching scheduler with paged KV cache.
Runs a background generation loop with four phases per iteration:
1. Cleanup finished tasks and release resources.
2. Refill active batch from the waiting queue.
3. Prefill newly activated tasks.
4. Decode the largest same-position group of active tasks.
"""
def __init__( def __init__(
self, self,
@@ -235,10 +87,10 @@ class InferenceScheduler:
tokenizer: AutoTokenizer, tokenizer: AutoTokenizer,
max_batch_size: int = 16, max_batch_size: int = 16,
max_seq_len: Optional[int] = None, max_seq_len: Optional[int] = None,
max_prefix_len: int = 512, max_prompt_len: int = 512,
cache_capacity: int = 1000, page_size: int = 64,
device: str = "cuda", device: Optional[str] = None,
dtype: torch.dtype = torch.bfloat16, dtype: Optional[torch.dtype] = None,
): ):
config = model.config config = model.config
@@ -246,42 +98,26 @@ class InferenceScheduler:
self.tokenizer = tokenizer self.tokenizer = tokenizer
self.max_batch_size = max_batch_size self.max_batch_size = max_batch_size
self.max_seq_len = max_seq_len or config.max_len self.max_seq_len = max_seq_len or config.max_len
self.max_prefix_len = max_prefix_len self.max_prompt_len = max_prompt_len
self.page_size = page_size
self.device = device or next(model.parameters()).device self.device = device or next(model.parameters()).device
self.dtype = dtype or next(model.parameters()).dtype self.dtype = dtype or next(model.parameters()).dtype
# Initialize prefix cache n_kv_heads = config.n_kv_heads
self.prefix_cache = PrefixCacheManager(max_capacity=cache_capacity)
num_kv_heads = config.n_kv_heads
head_dim = config.dim // config.n_heads head_dim = config.dim // config.n_heads
n_layers = config.n_layers n_layers = config.n_layers
n_pages = (
max_batch_size * (self.max_seq_len + page_size) + page_size - 1
) // page_size
k_cache = torch.empty( self.page_cache = PagedCache(
(
max_batch_size,
self.max_seq_len,
n_layers, n_layers,
num_kv_heads, n_pages,
page_size,
n_kv_heads,
head_dim, head_dim,
), self.device,
device=self.device, self.dtype,
dtype=self.dtype,
)
v_cache = torch.empty(
(
max_batch_size,
self.max_seq_len,
n_layers,
num_kv_heads,
head_dim,
),
device=self.device,
dtype=self.dtype,
)
self.kv_cache = (k_cache, v_cache)
self.seq_mask = torch.ones(
(max_batch_size, self.max_seq_len), device=self.device, dtype=torch.bool
) )
self.waiting_queue: List[Task] = [] self.waiting_queue: List[Task] = []
@@ -294,6 +130,9 @@ class InferenceScheduler:
self._total_tasks = 0 self._total_tasks = 0
self._total_tokens = 0 self._total_tokens = 0
def _n_pages_for(self, n_tokens: int) -> int:
return (n_tokens + self.page_size - 1) // self.page_size
def add_task( def add_task(
self, self,
prompt: str, prompt: str,
@@ -303,13 +142,10 @@ class InferenceScheduler:
top_k: int = 50, top_k: int = 50,
stream_callback: Optional[Callable[[str], None]] = None, stream_callback: Optional[Callable[[str], None]] = None,
) -> str: ) -> str:
"""Add a new task to the waiting queue."""
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}" task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
prompt_ids = self.tokenizer.encode(prompt) prompt_ids = self.tokenizer.encode(prompt)
if len(prompt_ids) > self.max_prompt_len:
# Truncate if exceeds max_prefix_len prompt_ids = prompt_ids[-self.max_prompt_len :]
if len(prompt_ids) > self.max_prefix_len:
prompt_ids = prompt_ids[: self.max_prefix_len]
task = Task( task = Task(
task_id=task_id, task_id=task_id,
@@ -321,16 +157,6 @@ class InferenceScheduler:
stream_callback=stream_callback, stream_callback=stream_callback,
) )
# Find longest matching prefix from cache
match = self.prefix_cache.find_longest_prefix(prompt_ids)
if match:
prefix_len, slot = match
task.prefix_len = prefix_len
task.slot = slot
else:
task.prefix_len = 0
task.slot = -1
with self._lock: with self._lock:
self.waiting_queue.append(task) self.waiting_queue.append(task)
self._total_tasks += 1 self._total_tasks += 1
@@ -339,13 +165,28 @@ class InferenceScheduler:
return task_id return task_id
def remove_task(self, task_id: str) -> None: def remove_task(self, task_id: str) -> None:
"""Remove a task from the scheduler."""
with self._lock: with self._lock:
removed_active = [t for t in self.active_tasks if t.task_id == task_id]
self.waiting_queue = [t for t in self.waiting_queue if t.task_id != task_id] self.waiting_queue = [t for t in self.waiting_queue if t.task_id != task_id]
self.active_tasks = [t for t in self.active_tasks if t.task_id != task_id] self.active_tasks = [t for t in self.active_tasks if t.task_id != task_id]
for task in removed_active:
if not task._pages_freed:
self._free_pages(task.page_table)
task.page_table.clear()
task.n_pages = 0
task._pages_freed = True
def _free_pages(self, indices: List[int]) -> None:
for idx in indices:
self.page_cache.free(idx)
def _record_page_hashes(self, task: Task, start_logical_page: int = 0) -> None:
full_pages = len(task.prompt_ids) // self.page_size
for i in range(start_logical_page, full_pages):
self.page_cache.record_page(task.page_table[i], task.prompt_ids, i)
def _remove_finished_tasks(self) -> None: def _remove_finished_tasks(self) -> None:
"""Remove finished tasks from active batch."""
finished = [] finished = []
for task in self.active_tasks: for task in self.active_tasks:
if task.is_finished(self.tokenizer.stop_ids): if task.is_finished(self.tokenizer.stop_ids):
@@ -355,280 +196,213 @@ class InferenceScheduler:
self._total_tokens += task.output_tokens self._total_tokens += task.output_tokens
for task in finished: for task in finished:
slot = task.slot if not task._pages_freed:
if slot >= 0 and slot < len(self.active_tasks): self._free_pages(task.page_table)
self.seq_mask[slot, :] = False task.page_table.clear()
task.n_pages = 0
# Release prefix cache reference task._pages_freed = True
if task.prefix_len > 0:
self.prefix_cache.release(tuple(task.prompt_ids[: task.prefix_len]))
task.slot = -1
self.active_tasks = [ self.active_tasks = [
t for t in self.active_tasks if t.status != TaskStatus.FINISHED t for t in self.active_tasks if t.status != TaskStatus.FINISHED
] ]
def _refill_active_batch(self) -> None: def _refill_active_batch(self) -> None:
"""Refill active batch with waiting tasks.""" available = self.max_batch_size - len(self.active_tasks)
available_slots = self.max_batch_size - len(self.active_tasks) if available <= 0:
if available_slots <= 0:
return return
to_add: List[Task] = []
with self._lock: with self._lock:
to_add = [ n = min(available, len(self.waiting_queue))
self.waiting_queue.pop(0) for _ in range(n):
for _ in range(min(available_slots, len(self.waiting_queue))) to_add.append(self.waiting_queue.pop(0))
]
failed: List[Task] = []
for task in to_add: for task in to_add:
task.slot = self._allocate_slot() prompt_len = len(task.prompt_ids)
hit_pages = self.page_cache.lookup_prefix(task.prompt_ids)
cached_tokens = len(hit_pages) * self.page_size
for p in hit_pages:
self.page_cache.inc_ref(p)
remaining = prompt_len - cached_tokens
n_new = self._n_pages_for(remaining) if remaining > 0 else 0
new_pages = self.page_cache.alloc_n(n_new) if n_new > 0 else []
if remaining > 0 and not new_pages:
for p in hit_pages:
self.page_cache.free(p)
failed.append(task)
continue
task.page_table = hit_pages + new_pages
task.n_pages = len(task.page_table)
task._prefix_cached_tokens = cached_tokens
task.status = TaskStatus.RUNNING task.status = TaskStatus.RUNNING
self.active_tasks.append(task) self.active_tasks.append(task)
def _allocate_slot(self) -> int: if failed:
"""Allocate an available slot for a task.""" with self._lock:
for i in range(self.max_batch_size): self.waiting_queue[:0] = failed
if not any(t.slot == i for t in self.active_tasks):
return i
return -1
def _execute_prefill(self, tasks: List[Task]) -> None: def _execute_prefill(
"""Execute Prefill phase with incremental prefill support.""" self, tasks: List[Task], prompt_len: int, start_pos: int = 0
if not tasks: ) -> None:
return tasks = sorted(tasks, key=lambda t: t.task_id)
batch_sz = len(tasks)
# Group tasks by prefix cache status seq_len = prompt_len - start_pos
fully_cached, partial, full = [], [], [] input_ids = torch.empty(batch_sz, seq_len, dtype=torch.long, device=self.device)
for task in tasks: input_mask = torch.ones(batch_sz, seq_len, dtype=torch.bool, device=self.device)
total_len, prefix_len = len(task.prompt_ids), task.prefix_len
if prefix_len == total_len:
fully_cached.append(task)
elif prefix_len > 0:
partial.append(task)
else:
full.append(task)
# Handle fully cached tasks for i, t in enumerate(tasks):
for t in fully_cached: input_ids[i] = torch.tensor(
t.input_tokens, t.output_tokens = len(t.prompt_ids), 0 t.prompt_ids[start_pos:prompt_len], device=self.device
if t.slot >= 0:
self.seq_mask[t.slot, : t.input_tokens] = True
if full:
self._execute_full_prefill(full)
if partial:
self._execute_partial_prefill(partial)
def _execute_full_prefill(self, tasks: List[Task]) -> None:
"""Execute full prefill for tasks without prefix cache."""
if not tasks:
return
tasks = sorted(tasks, key=lambda t: t.slot)
prompt_lens = [len(task.prompt_ids) for task in tasks]
max_len = max(prompt_lens)
input_ids = torch.zeros(
len(tasks), max_len, dtype=torch.long, device=self.device
)
for i, task in enumerate(tasks):
if len(task.prompt_ids) > 0:
input_ids[i, : len(task.prompt_ids)] = torch.tensor(
task.prompt_ids, device=self.device
) )
if self.tokenizer.pad_id is not None: page_tables = self._make_page_table_tensor(tasks)
input_mask = torch.ne(input_ids, self.tokenizer.pad_id)
else:
input_mask = torch.ones(
input_ids.shape, dtype=torch.bool, device=self.device
)
with torch.inference_mode(): with torch.inference_mode():
self.model( self.model(
input_ids, input_ids,
input_mask=input_mask, input_mask=input_mask,
start_pos=0, start_pos=start_pos,
persistent_key_values=self.kv_cache, paged_cache=self.page_cache.bind(page_tables, total_len=prompt_len),
) )
for i, task in enumerate(tasks): start_logical_page = start_pos // self.page_size
task.input_tokens = prompt_lens[i] for t in tasks:
task.output_tokens = 0 self._record_page_hashes(t, start_logical_page=start_logical_page)
# Insert new prefix into cache
self.prefix_cache.insert(tuple(task.prompt_ids), task.slot)
for task in tasks:
if task.slot >= 0:
self.seq_mask[task.slot, : task.input_tokens] = True
def _execute_partial_prefill(self, tasks: List[Task]) -> None:
"""Execute incremental prefill for tasks with partial prefix cache match."""
for task in tasks:
total_len = len(task.prompt_ids)
prefix_len = task.prefix_len
if prefix_len >= total_len:
task.input_tokens = total_len
task.output_tokens = 0
continue
# Get new tokens that need prefill
new_ids = task.prompt_ids[prefix_len:]
new_len = len(new_ids)
if new_len == 0:
task.input_tokens = total_len
task.output_tokens = 0
continue
# Build input for incremental prefill
input_ids = torch.tensor([new_ids], dtype=torch.long, device=self.device)
# Input mask should cover from position 0 to prefix_len + new_len
# The prefix part uses cached KV, new part needs computation
input_mask = torch.ones(
(1, prefix_len + new_len), dtype=torch.bool, device=self.device
)
with torch.inference_mode():
self.model(
input_ids,
input_mask=input_mask,
start_pos=prefix_len,
persistent_key_values=self.kv_cache,
)
task.input_tokens = total_len
task.output_tokens = 0
# Insert full prefix into cache (ref_count already increased in add_task)
self.prefix_cache.insert(tuple(task.prompt_ids), task.slot)
if task.slot >= 0:
self.seq_mask[task.slot, : task.input_tokens] = True
def _execute_decode(self, tasks: List[Task], start_pos: int) -> None: def _execute_decode(self, tasks: List[Task], start_pos: int) -> None:
"""Execute Decode phase."""
if not tasks: if not tasks:
return return
tasks = sorted(tasks, key=lambda t: t.slot) tasks = sorted(tasks, key=lambda t: t.task_id)
batch_sz = len(tasks)
input_ids = torch.zeros(len(tasks), dtype=torch.long, device=self.device) for t in tasks:
for i, task in enumerate(tasks): self._maybe_alloc_page(t, start_pos)
if task.output_ids:
input_ids[i] = task.output_ids[-1]
else:
input_ids[i] = task.prompt_ids[-1]
input_tensor = input_ids.unsqueeze(1) input_ids = torch.tensor(
active_mask = torch.ones((len(tasks), 1), dtype=torch.bool, device=self.device) [t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks],
dtype=torch.long,
device=self.device,
)
active_mask = torch.ones((batch_sz, 1), dtype=torch.bool, device=self.device)
page_tables = self._make_page_table_tensor(tasks)
total_len = start_pos + 1
temperatures = torch.tensor([t.temperature for t in tasks], device=self.device)
top_ks = torch.tensor([t.top_k for t in tasks], device=self.device)
top_ps = torch.tensor([t.top_p for t in tasks], device=self.device)
with torch.inference_mode(): with torch.inference_mode():
outputs = self.model( outputs = self.model(
input_tensor, input_ids.unsqueeze(1),
input_mask=active_mask, input_mask=active_mask,
persistent_key_values=self.kv_cache, paged_cache=self.page_cache.bind(page_tables, total_len=total_len),
start_pos=start_pos, start_pos=start_pos,
) )
logits = outputs["logits"][:, -1, :] logits = outputs["logits"][:, -1, :]
next_token_ids = [] next_tokens = sample(
for i, task in enumerate(tasks): logits,
logit = logits[i : i + 1] temperature=temperatures,
logit = apply_sampling_strategies( top_k=top_ks,
logit, top_p=top_ps,
task.temperature, ).tolist()
task.top_k,
task.top_p,
)
probs = torch.softmax(logit, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
next_token_ids.append(next_token.item())
for task, next_token in zip(tasks, next_token_ids): for t, ntok in zip(tasks, next_tokens):
task.output_ids.append(next_token) t.output_ids.append(ntok)
task.output_tokens += 1 t.output_tokens += 1
pos = t.input_tokens + t.output_tokens
self._maybe_alloc_page(t, pos)
if t.stream_callback:
t.stream_callback(self.tokenizer.decode([ntok]))
pos = task.input_tokens + task.output_tokens for t in tasks:
if task.slot >= 0 and pos < self.max_seq_len: if t.is_finished(self.tokenizer.stop_ids):
self.seq_mask[task.slot, pos] = True if t.stream_callback:
t.stream_callback(STOP)
if task.stream_callback: def _make_page_table_tensor(self, tasks: List[Task]) -> Tensor:
token_str = self.tokenizer.decode([next_token]) max_pages = max(t.n_pages for t in tasks)
task.stream_callback(token_str) rows = [t.page_table + [-1] * (max_pages - t.n_pages) for t in tasks]
return torch.tensor(rows, dtype=torch.long, device=self.device)
for task in tasks: def _maybe_alloc_page(self, task: Task, pos: int) -> None:
if task.output_tokens >= task.max_tokens or ( needed = self._n_pages_for(pos + 1)
task.output_ids and task.output_ids[-1] in self.tokenizer.stop_ids while task.n_pages < needed:
): p = self.page_cache.alloc()
if task.stream_callback: if p < 0:
task.stream_callback("[DONE]") break
task.page_table.append(p)
task.n_pages += 1
def _run_generation_loop(self) -> None: def _run_generation_loop(self) -> None:
"""Main generation loop.""" try:
while self._running: while self._running:
self._remove_finished_tasks() self._remove_finished_tasks()
self._refill_active_batch() self._refill_active_batch()
if not self.active_tasks: if not self.active_tasks and not self.waiting_queue:
self._task_event.wait(timeout=0.01)
self._task_event.clear() self._task_event.clear()
self._task_event.wait(timeout=1.0)
continue continue
new_tasks = [t for t in self.active_tasks if t.output_tokens == 0] to_prefill = [t for t in self.active_tasks if t.output_tokens == 0]
decode_tasks = [t for t in self.active_tasks if t.output_tokens > 0] if to_prefill:
for t in to_prefill:
t.input_tokens = len(t.prompt_ids)
if decode_tasks: groups: Dict[Tuple[int, int], List[Task]] = {}
start_pos = max(t.input_tokens + t.output_tokens for t in decode_tasks) for t in to_prefill:
else: key = (len(t.prompt_ids), t._prefix_cached_tokens)
start_pos = 0 groups.setdefault(key, []).append(t)
if new_tasks: for (prompt_len, start_pos), group in groups.items():
self._execute_prefill(new_tasks) if start_pos < prompt_len:
decode_tasks = new_tasks self._execute_prefill(group, prompt_len, start_pos)
start_pos = max(t.input_tokens for t in decode_tasks)
if decode_tasks: pos_groups: Dict[int, List[Task]] = {}
self._execute_decode(decode_tasks, start_pos) for t in self.active_tasks:
pos_groups.setdefault(t.next_pos, []).append(t)
if not self.active_tasks and not self.waiting_queue: if pos_groups:
self._task_event.wait(timeout=0.05) best_pos = max(pos_groups, key=lambda p: len(pos_groups[p]))
self._task_event.clear() self._execute_decode(pos_groups[best_pos], best_pos)
except Exception as e:
logger.error(f"Scheduler loop crashed: {e}", exc_info=True)
for task in self.active_tasks:
if task.stream_callback:
task.stream_callback(STOP)
for task in self.waiting_queue:
if task.stream_callback:
task.stream_callback(STOP)
raise
def start(self) -> None: def start(self) -> None:
"""Start the generation loop."""
if not self._running: if not self._running:
self._running = True self._running = True
self._loop_thread = threading.Thread(target=self._run_generation_loop) t = threading.Thread(target=self._run_generation_loop, daemon=True)
self._loop_thread.daemon = True t.start()
self._loop_thread.start() self._loop_thread = t
def stop(self) -> None: def stop(self) -> None:
"""Stop the generation loop."""
self._running = False self._running = False
self._task_event.set()
if hasattr(self, "_loop_thread"): if hasattr(self, "_loop_thread"):
self._loop_thread.join(timeout=1.0) self._loop_thread.join(timeout=2.0)
# Clear KV cache to free GPU memory
if self.kv_cache is not None:
k_cache, v_cache = self.kv_cache
if k_cache is not None:
k_cache.detach()
if v_cache is not None:
v_cache.detach()
# Clear seq mask
self.seq_mask.detach()
# Clear task lists
self.waiting_queue.clear() self.waiting_queue.clear()
self.active_tasks.clear() self.active_tasks.clear()
if torch.cuda.is_available():
torch.cuda.empty_cache()
def get_stats(self) -> Dict[str, Any]: def get_stats(self) -> Dict[str, Any]:
"""Get scheduler statistics."""
return { return {
"total_tasks": self._total_tasks, "total_tasks": self._total_tasks,
"total_tokens": self._total_tokens, "total_tokens": self._total_tokens,
+346 -181
View File
@@ -1,15 +1,14 @@
""" """
Inference Server with Continuous Batching Support OpenAI / Anthropic-compatible chat completion server backed by continuous-batching inference.
FastAPI server for inference with continuous batching.
Provides OpenAI-compatible chat completion endpoints.
""" """
import json import json
import logging import logging
import time
import uuid
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from pathlib import Path from pathlib import Path
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional, Union
import torch import torch
import uvicorn import uvicorn
@@ -23,13 +22,13 @@ from astrai.tokenize import AutoTokenizer
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# Global model parameter and engine (loaded once)
_engine: Optional[InferenceEngine] = None
_model_param: Optional[Any] = None
_project_root = Path(__file__).parent.parent.parent _project_root = Path(__file__).parent.parent.parent
# Server configuration (set before running server)
_server_config: Dict[str, Any] = { class ServerState:
def __init__(self):
self.engine: Optional[InferenceEngine] = None
self.config: Dict[str, Any] = {
"device": "cuda", "device": "cuda",
"dtype": torch.bfloat16, "dtype": torch.bfloat16,
"param_path": None, "param_path": None,
@@ -37,45 +36,80 @@ _server_config: Dict[str, Any] = {
} }
_state = ServerState()
class ChatMessage(BaseModel):
role: str
content: str
class ChatCompletionRequest(BaseModel):
"""OpenAI Chat Completion API request body."""
model: str = "astrai"
messages: List[ChatMessage]
temperature: Optional[float] = Field(default=1.0, ge=0.0, le=2.0)
top_p: Optional[float] = Field(default=1.0, ge=0.0, le=1.0)
top_k: Optional[int] = Field(default=50, ge=1)
stream: Optional[bool] = False
stop: Optional[Union[str, List[str]]] = None
max_tokens: Optional[int] = Field(default=2048, ge=1)
n: Optional[int] = Field(default=1, ge=1)
presence_penalty: Optional[float] = Field(default=0.0, ge=-2.0, le=2.0)
frequency_penalty: Optional[float] = Field(default=0.0, ge=-2.0, le=2.0)
logit_bias: Optional[Dict[int, float]] = None
user: Optional[str] = None
class AnthropicMessage(BaseModel):
role: str
content: Union[str, List[Dict[str, Any]]]
class MessagesRequest(BaseModel):
"""Anthropic Messages API request body."""
model: str = "astrai"
max_tokens: int = Field(default=1024, ge=1)
messages: List[AnthropicMessage]
system: Optional[str] = None
temperature: Optional[float] = Field(default=1.0, ge=0.0, le=2.0)
top_p: Optional[float] = Field(default=1.0, ge=0.0, le=1.0)
top_k: Optional[int] = Field(default=50, ge=1)
stream: Optional[bool] = False
stop_sequences: Optional[List[str]] = None
def configure_server( def configure_server(
device: str = "cuda", device: str = "cuda",
dtype: torch.dtype = torch.bfloat16, dtype: torch.dtype = torch.bfloat16,
param_path: Optional[Path] = None, param_path: Optional[Path] = None,
max_batch_size: int = 16, max_batch_size: int = 16,
): ):
"""Configure server settings before starting. _state.config.update(
device=device,
Args: dtype=dtype,
device: Device to load model on (e.g., "cuda", "cpu", "cuda:0") param_path=param_path,
dtype: Data type for model weights (e.g., torch.bfloat16, torch.float16) max_batch_size=max_batch_size,
param_path: Path to model parameters directory )
max_batch_size: Maximum batch size for continuous batching
"""
_server_config["device"] = device
_server_config["dtype"] = dtype
_server_config["param_path"] = param_path
_server_config["max_batch_size"] = max_batch_size
@asynccontextmanager @asynccontextmanager
async def lifespan(app: FastAPI): async def lifespan(app: FastAPI):
"""Lifespan context manager for startup and shutdown events."""
global _model_param, _engine
# Startup: Load model with configured settings
try: try:
load_model( load_model(
param_path=_server_config["param_path"], param_path=_state.config["param_path"],
device=_server_config["device"], device=_state.config["device"],
dtype=_server_config["dtype"], dtype=_state.config["dtype"],
max_batch_size=_server_config["max_batch_size"], max_batch_size=_state.config["max_batch_size"],
) )
except Exception as e: except Exception as e:
logger.error(f"Failed to load model: {e}") logger.error(f"Failed to load model: {e}")
raise raise
yield yield
# Shutdown: Cleanup engine if _state.engine:
if _engine: _state.engine.shutdown()
_engine.shutdown()
logger.info("Inference engine shutdown complete") logger.info("Inference engine shutdown complete")
@@ -88,203 +122,345 @@ def load_model(
dtype: torch.dtype = torch.bfloat16, dtype: torch.dtype = torch.bfloat16,
max_batch_size: int = 16, max_batch_size: int = 16,
): ):
"""Load model parameters and initialize inference engine."""
global _model_param, _engine
if param_path is None: if param_path is None:
param_path = _project_root / "params" param_path = _project_root / "params"
if not param_path.exists(): if not param_path.exists():
raise FileNotFoundError(f"Parameter directory not found: {param_path}") raise FileNotFoundError(f"Parameter directory not found: {param_path}")
# Load tokenizer separately
tokenizer = AutoTokenizer.from_pretrained(param_path) tokenizer = AutoTokenizer.from_pretrained(param_path)
_model_param = AutoModel.from_pretrained(param_path) model = AutoModel.from_pretrained(param_path)
_model_param.to(device=device, dtype=dtype) model.to(device=device, dtype=dtype)
logger.info(f"Model loaded on {device} with dtype {dtype}") logger.info(f"Model loaded on {device} with dtype {dtype}")
# Initialize inference engine with separate model and tokenizer _state.engine = InferenceEngine(
_engine = InferenceEngine( model=model,
model=_model_param,
tokenizer=tokenizer, tokenizer=tokenizer,
max_batch_size=max_batch_size, max_batch_size=max_batch_size,
) )
logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}") logger.info(f"Inference engine initialized with max_batch_size={max_batch_size}")
# Pydantic models for API request/response def _get_engine() -> InferenceEngine:
class ChatMessage(BaseModel): if _state.engine is None:
role: str # "user", "assistant", "system" raise HTTPException(status_code=503, detail="Engine not initialized")
content: str return _state.engine
class ChatCompletionRequest(BaseModel): def _make_chunk(
messages: List[ChatMessage] delta: Dict[str, str],
temperature: float = Field(0.8, ge=0.0, le=2.0) finish_reason: Optional[str] = None,
top_p: float = Field(0.95, ge=0.0, le=1.0) *,
top_k: int = Field(50, ge=0) resp_id: str,
max_tokens: int = Field(2048, ge=1) created: int,
stream: bool = False model: str,
system_prompt: Optional[str] = None index: int = 0,
) -> str:
"""Build a single SSE ``data:`` chunk matching OpenAI streaming format."""
class CompletionResponse(BaseModel): data = {
id: str = "chatcmpl-default" "id": resp_id,
object: str = "chat.completion" "object": "chat.completion.chunk",
created: int = 0 "created": created,
model: str = "astrai" "model": model,
choices: List[Dict[str, Any]] "choices": [
{
"index": index,
"delta": delta,
"finish_reason": finish_reason,
}
],
}
return f"data: {json.dumps(data, ensure_ascii=False)}\n\n"
@app.get("/health") @app.get("/health")
async def health(): async def health():
return { return {
"status": "ok", "status": "ok",
"model_loaded": _model_param is not None, "model_loaded": _state.engine is not None,
"engine_ready": _engine is not None,
} }
@app.get("/stats") @app.get("/stats")
async def get_stats(): async def get_stats():
"""Get inference engine statistics.""" return _get_engine().get_stats()
if _engine is None:
raise HTTPException(status_code=503, detail="Engine not initialized")
return _engine.get_stats()
@app.post("/v1/chat/completions", response_model=CompletionResponse) @app.post("/v1/chat/completions")
async def chat_completion(request: ChatCompletionRequest): async def chat_completion(request: ChatCompletionRequest):
"""OpenAI-compatible chat completion endpoint. """OpenAI-compatible chat completion endpoint (streaming + non-streaming)."""
engine = _get_engine()
resp_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
created = int(time.time())
model = request.model
Supports both streaming and non-streaming modes with continuous batching. prompt = engine.tokenizer.apply_chat_template(
"""
if _engine is None:
raise HTTPException(status_code=503, detail="Engine not initialized")
# Convert messages to prompt using engine's tokenizer
# Extract system prompt if present, then apply chat template
# Apply chat template directly with messages
prompt = _engine.tokenizer.apply_chat_template(
[{"role": m.role, "content": m.content} for m in request.messages], [{"role": m.role, "content": m.content} for m in request.messages],
tokenize=False, tokenize=False,
) )
prompt_tokens = len(engine.tokenizer.encode(prompt))
if request.stream: if request.stream:
# Streaming response (use synchronous generator) agen = engine.generate_async(
generator = _engine.generate(
prompt=prompt, prompt=prompt,
stream=True,
max_tokens=request.max_tokens, max_tokens=request.max_tokens,
temperature=request.temperature, temperature=request.temperature,
top_p=request.top_p, top_p=request.top_p,
top_k=request.top_k, top_k=request.top_k,
) )
def generate_stream(): async def event_stream():
for token in generator: yield _make_chunk(
if token == "[DONE]": {"role": "assistant"},
break finish_reason=None,
yield f"data: {json.dumps({'choices': [{'delta': {'content': token}}]})}\n\n" resp_id=resp_id,
created=created,
model=model,
)
completion_tokens = 0
async for token in agen:
yield _make_chunk(
{"content": token},
finish_reason=None,
resp_id=resp_id,
created=created,
model=model,
)
completion_tokens += 1
yield _make_chunk(
{},
finish_reason="stop",
resp_id=resp_id,
created=created,
model=model,
)
usage = {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
}
yield f"data: {json.dumps(usage, ensure_ascii=False)}\n\n"
yield "data: [DONE]\n\n" yield "data: [DONE]\n\n"
return StreamingResponse( return StreamingResponse(
generate_stream(), event_stream(),
media_type="text/event-stream", media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"}, headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
) )
else:
# Non-streaming response completion_tokens = 0
result = _engine.generate( chunks: List[str] = []
agen = engine.generate_async(
prompt=prompt,
max_tokens=request.max_tokens,
temperature=request.temperature,
top_p=request.top_p,
top_k=request.top_k,
)
async for token in agen:
chunks.append(token)
completion_tokens += 1
content = "".join(chunks)
return {
"id": resp_id,
"object": "chat.completion",
"created": created,
"model": model,
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": content},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
},
}
def _make_anthropic_sse(event: str, data: Dict[str, Any]) -> str:
return f"event: {event}\ndata: {json.dumps(data, ensure_ascii=False)}\n\n"
def _check_stop_sequence(text: str, stop_sequences: List[str]) -> Optional[str]:
for seq in stop_sequences:
if seq and seq in text:
return seq
return None
def _extract_text_content(content: Union[str, List[Dict[str, Any]]]) -> str:
if isinstance(content, str):
return content
if isinstance(content, list):
for block in content:
if isinstance(block, dict) and block.get("type") == "text":
return block.get("text", "")
return ""
def _build_anthropic_messages(
messages: List[AnthropicMessage], system: Optional[str]
) -> List[Dict[str, str]]:
result: List[Dict[str, str]] = []
if system:
result.append({"role": "system", "content": system})
for m in messages:
content = _extract_text_content(m.content)
if content:
result.append({"role": m.role, "content": content})
return result
@app.post("/v1/messages")
async def create_message(request: MessagesRequest):
"""Anthropic-compatible Messages API endpoint (streaming + non-streaming)."""
engine = _get_engine()
resp_id = f"msg_{uuid.uuid4().hex[:24]}"
model = request.model
chat_messages = _build_anthropic_messages(request.messages, request.system)
prompt = engine.tokenizer.apply_chat_template(chat_messages, tokenize=False)
prompt_tokens = len(engine.tokenizer.encode(prompt))
stop_sequences = request.stop_sequences or []
if request.stream:
agen = engine.generate_async(
prompt=prompt, prompt=prompt,
stream=False,
max_tokens=request.max_tokens, max_tokens=request.max_tokens,
temperature=request.temperature, temperature=request.temperature,
top_p=request.top_p, top_p=request.top_p,
top_k=request.top_k, top_k=request.top_k,
) )
# Build OpenAI-style response async def event_stream():
import time yield _make_anthropic_sse(
"message_start",
resp = CompletionResponse(
id=f"chatcmpl-{int(time.time())}",
created=int(time.time()),
choices=[
{ {
"type": "message_start",
"message": {
"id": resp_id,
"type": "message",
"role": "assistant",
"model": model,
"content": [],
"usage": {"input_tokens": prompt_tokens},
},
},
)
yield _make_anthropic_sse(
"content_block_start",
{
"type": "content_block_start",
"index": 0, "index": 0,
"message": {"role": "assistant", "content": result}, "content_block": {"type": "text", "text": ""},
"finish_reason": "stop", },
)
completion_tokens = 0
accumulated = ""
stopped_seq: Optional[str] = None
async for token in agen:
accumulated += token
completion_tokens += 1
matched = _check_stop_sequence(accumulated, stop_sequences)
if matched:
text = accumulated[: accumulated.rfind(matched)]
stopped_seq = matched
if text:
yield _make_anthropic_sse(
"content_block_delta",
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "text_delta", "text": text},
},
)
break
yield _make_anthropic_sse(
"content_block_delta",
{
"type": "content_block_delta",
"index": 0,
"delta": {"type": "text_delta", "text": token},
},
)
yield _make_anthropic_sse(
"content_block_stop",
{"type": "content_block_stop", "index": 0},
)
stop_reason = "stop_sequence" if stopped_seq else "end_turn"
yield _make_anthropic_sse(
"message_delta",
{
"type": "message_delta",
"delta": {"stop_reason": stop_reason, "stop_sequence": stopped_seq},
"usage": {"output_tokens": completion_tokens},
},
)
yield _make_anthropic_sse(
"message_stop",
{"type": "message_stop"},
)
return StreamingResponse(
event_stream(),
media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
)
completion_tokens = 0
chunks: List[str] = []
agen = engine.generate_async(
prompt=prompt,
max_tokens=request.max_tokens,
temperature=request.temperature,
top_p=request.top_p,
top_k=request.top_k,
)
stopped_seq: Optional[str] = None
accumulated = ""
async for token in agen:
chunks.append(token)
completion_tokens += 1
accumulated += token
matched = _check_stop_sequence(accumulated, stop_sequences)
if matched:
stopped_seq = matched
break
content = "".join(chunks)
if stopped_seq:
idx = content.rfind(stopped_seq)
if idx != -1:
content = content[:idx]
return {
"id": resp_id,
"type": "message",
"role": "assistant",
"model": model,
"content": [{"type": "text", "text": content}],
"stop_reason": "stop_sequence" if stopped_seq else "end_turn",
"stop_sequence": stopped_seq,
"usage": {
"input_tokens": prompt_tokens,
"output_tokens": completion_tokens,
},
} }
],
)
return resp
@app.post("/generate")
async def generate(
query: str,
history: Optional[List[List[str]]] = None,
temperature: float = 0.8,
top_p: float = 0.95,
top_k: int = 50,
max_len: int = 2048,
stream: bool = False,
):
"""Simple generation endpoint.
Args:
query: Input query string
history: Conversation history as list of [user, assistant] pairs
temperature: Sampling temperature
top_p: Top-p sampling parameter
top_k: Top-k sampling parameter
max_len: Maximum tokens to generate
stream: Enable streaming output
Returns:
dict: Generation result with response field
"""
if _engine is None:
raise HTTPException(status_code=503, detail="Engine not initialized")
# Build messages for chat template
messages = []
if history:
# Convert history format: List[List[str]] -> List[Dict]
for h in history:
if len(h) >= 2:
messages.append({"role": "user", "content": h[0]})
messages.append({"role": "assistant", "content": h[1]})
messages.append({"role": "user", "content": query})
# Use tokenizer's chat template
prompt = _engine.tokenizer.apply_chat_template(messages, tokenize=False)
if stream:
# Synchronous streaming
result = _engine.generate(
prompt=prompt,
stream=True,
max_tokens=max_len,
temperature=temperature,
top_p=top_p,
top_k=top_k,
)
def stream_generator():
for token in result:
yield token + "\n"
return StreamingResponse(stream_generator(), media_type="text/plain")
else:
result = _engine.generate(
prompt=prompt,
stream=False,
max_tokens=max_len,
temperature=temperature,
top_p=top_p,
top_k=top_k,
)
return {"response": result}
def run_server( def run_server(
@@ -296,17 +472,6 @@ def run_server(
param_path: Optional[Path] = None, param_path: Optional[Path] = None,
max_batch_size: int = 16, max_batch_size: int = 16,
): ):
"""Run the FastAPI server with uvicorn.
Args:
host: Server host address
port: Server port number
reload: Enable auto-reload for development
device: Device to load model on (e.g., "cuda", "cpu", "cuda:0")
dtype: Data type for model weights (e.g., torch.bfloat16, torch.float16)
param_path: Path to model parameters directory
max_batch_size: Maximum batch size for continuous batching
"""
configure_server( configure_server(
device=device, device=device,
dtype=dtype, dtype=dtype,
+9 -14
View File
@@ -4,12 +4,13 @@ AutoModel base class for model loading and saving.
from contextlib import contextmanager from contextlib import contextmanager
from pathlib import Path from pathlib import Path
from typing import Dict, Self, Type, Union from typing import Self, Type, Union
import safetensors.torch as st import safetensors.torch as st
import torch.nn as nn import torch.nn as nn
from astrai.config import ModelConfig from astrai.config import ModelConfig
from astrai.factory import Registry
@contextmanager @contextmanager
@@ -44,8 +45,7 @@ class AutoModel(nn.Module):
Provides model loading/saving and generation capabilities. Provides model loading/saving and generation capabilities.
""" """
# Model registry - stored as class attribute _registry = Registry()
_registry: Dict[str, Type["AutoModel"]] = {}
def __init__(self, config: ModelConfig): def __init__(self, config: ModelConfig):
super().__init__() super().__init__()
@@ -63,7 +63,7 @@ class AutoModel(nn.Module):
""" """
def decorator(sub_cls: Type["AutoModel"]) -> Type["AutoModel"]: def decorator(sub_cls: Type["AutoModel"]) -> Type["AutoModel"]:
cls._registry[model_type.lower()] = sub_cls cls._registry.register(model_type.lower(), sub_cls)
return sub_cls return sub_cls
return decorator return decorator
@@ -72,18 +72,19 @@ class AutoModel(nn.Module):
def get_model_class(cls, model_type: str) -> Type["AutoModel"]: def get_model_class(cls, model_type: str) -> Type["AutoModel"]:
"""Get model class by model_type string.""" """Get model class by model_type string."""
model_type = model_type.lower() model_type = model_type.lower()
if model_type not in cls._registry: if not cls._registry.contains(model_type):
available = list(cls._registry.keys()) available = cls._registry.list_names()
raise ValueError( raise ValueError(
f"Unknown model_type: {model_type}. Available: {available}" f"Unknown model_type: {model_type}. Available: {available}"
) )
return cls._registry[model_type] return cls._registry.get(model_type)
@classmethod @classmethod
def from_pretrained( def from_pretrained(
cls, cls,
path: Union[str, Path], path: Union[str, Path],
disable_random_init: bool = True, disable_random_init: bool = True,
strict: bool = True,
) -> nn.Module: ) -> nn.Module:
model_path = Path(path) model_path = Path(path)
@@ -96,14 +97,8 @@ class AutoModel(nn.Module):
else: else:
raise FileNotFoundError(f"Config file not found: {config_path}") raise FileNotFoundError(f"Config file not found: {config_path}")
# If called from base class, use model_type to determine actual model class
if cls is AutoModel:
model_type = config.model_type or "transformer" model_type = config.model_type or "transformer"
actual_cls = cls.get_model_class(model_type) actual_cls = cls.get_model_class(model_type)
else:
raise ValueError(
f"Cannot call from_pretrained() on subclass {cls.__name__}"
)
with _disable_random_init(enable=disable_random_init): with _disable_random_init(enable=disable_random_init):
model = actual_cls(config) model = actual_cls(config)
@@ -112,7 +107,7 @@ class AutoModel(nn.Module):
weights_path = model_path / "model.safetensors" weights_path = model_path / "model.safetensors"
if weights_path.exists(): if weights_path.exists():
state_dict = st.load_file(str(weights_path)) state_dict = st.load_file(str(weights_path))
model.load_state_dict(state_dict, strict=False) model.load_state_dict(state_dict, strict=strict)
return model return model
+27 -72
View File
@@ -5,17 +5,11 @@ import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from torch import Tensor from torch import Tensor
from astrai.inference.cache import CacheView
def repeat_kv(x: Tensor, n_rep: int) -> Tensor: def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
""" """Repeat KV heads n_rep times for GQA."""
Repeat k times along the dimension for attention heads.
Args:
x (Tensor): The input tensor.
n_rep (int): The number of repetitions.
Returns:
Tensor: The repeated tensor.
"""
bs, slen, n_heads, head_dim = x.shape bs, slen, n_heads, head_dim = x.shape
if n_rep == 1: if n_rep == 1:
return x return x
@@ -32,49 +26,25 @@ def get_rotary_emb(
base: float = 10000, base: float = 10000,
device: Optional[torch.device] = None, device: Optional[torch.device] = None,
) -> Tuple[Tensor, Tensor]: ) -> Tuple[Tensor, Tensor]:
""" """Precompute cos/sin for RoPE."""
Get the rotary embedding for the given dimension and maximum length.
Args:
dim (int): The dimension of the input.
max_len (int): The maximum length of the input.
base (float, optional): The base for the frequency. Defaults to 10000.
device (optional): The device to create tensors on. Defaults to None.
Returns:
Tensor: The rotary embedding tensor.
"""
theta = base ** (-torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim) theta = base ** (-torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim)
t = torch.arange(0, max_len, dtype=torch.float64, device=device) t = torch.arange(0, max_len, dtype=torch.float64, device=device)
freqs = torch.outer(t, theta) freqs = torch.outer(t, theta)
return torch.cos(freqs).float(), torch.sin(freqs).float() return torch.cos(freqs).float(), torch.sin(freqs).float()
def apply_rotary_emb(x: torch.Tensor, rotary_emb: Tuple[Tensor, Tensor]) -> Tensor: def apply_rotary_emb(x: torch.Tensor, rotary_emb: Tuple[Tensor, Tensor]) -> Tensor:
""" """Apply rotary embedding via cos/sin (shape-preserving)."""
Apply rotary embedding to the input tensor using cos/sin form.
Args:
x (Tensor): The input tensor (shape [..., seq_len, dim]).
rotary_emb (Tuple[Tensor, Tensor]): The rotary embedding (shape [seq_len, dim//2]).
Returns:
Tensor: The output tensor (rotated, same shape as input).
"""
dtype = x.dtype dtype = x.dtype
cos, sin = rotary_emb cos, sin = rotary_emb
cos = cos.unsqueeze(0).unsqueeze(2)
cos = cos.unsqueeze(0).unsqueeze(2) # [1, seq_len, 1, dim//2] sin = sin.unsqueeze(0).unsqueeze(2)
sin = sin.unsqueeze(0).unsqueeze(2) # [1, seq_len, 1, dim//2] x_real = x[..., 0::2]
x_imag = x[..., 1::2]
x_real = x[..., 0::2] # [batch, seq_len, dim//2]
x_imag = x[..., 1::2] # [batch, seq_len, dim//2]
x_real_rot = x_real * cos - x_imag * sin x_real_rot = x_real * cos - x_imag * sin
x_imag_rot = x_real * sin + x_imag * cos x_imag_rot = x_real * sin + x_imag * cos
x_out = torch.stack([x_real_rot, x_imag_rot], dim=-1)
x_out = torch.stack([x_real_rot, x_imag_rot], dim=-1) # [batch, seq_len, dim//2, 2] x_out = x_out.view(*x_out.shape[:-2], -1)
x_out = x_out.view(*x_out.shape[:-2], -1) # [batch, seq_len, dim]
return x_out.to(dtype) return x_out.to(dtype)
@@ -95,13 +65,10 @@ class RotaryEmbedding(nn.Module):
def forward(self, x: Tensor, start_pos: int = 0) -> Tuple[Tensor, Tensor]: def forward(self, x: Tensor, start_pos: int = 0) -> Tuple[Tensor, Tensor]:
seq_len = x.size(1) seq_len = x.size(1)
if self.max_len_cached < seq_len + start_pos: if self.max_len_cached < seq_len + start_pos:
self._set_rotary_buffer(self.max_len_cached * 2, x.device) self._set_rotary_buffer(self.max_len_cached * 2, x.device)
cos = self.cos_cached[start_pos : start_pos + seq_len] cos = self.cos_cached[start_pos : start_pos + seq_len]
sin = self.sin_cached[start_pos : start_pos + seq_len] sin = self.sin_cached[start_pos : start_pos + seq_len]
return (cos, sin) return (cos, sin)
@@ -185,13 +152,13 @@ class GQA(nn.Module):
x: Tensor, x: Tensor,
rotary_emb: Tuple[Tensor, Tensor], rotary_emb: Tuple[Tensor, Tensor],
mask: Tensor = None, mask: Tensor = None,
kv_cache: Optional[Tuple[Tensor, Tensor]] = None, paged_cache: Optional[CacheView] = None,
start_pos: int = 0, start_pos: int = 0,
) -> Tensor: ) -> Tensor:
bsz, seq_len, _ = x.size() bsz, seq_len, _ = x.size()
is_causal = mask is None is_causal = mask is None
# x(bsz, seq_len, n_heads * head_dim) -> (bsz, seq_len, n_heads, head_dim) # (bsz, seq_len, dim) -> (bsz, seq_len, n_heads, head_dim)
q = self._split_heads(self.q_proj(x), self.n_heads) q = self._split_heads(self.q_proj(x), self.n_heads)
k = self._split_heads(self.k_proj(x), self.n_kv_heads) k = self._split_heads(self.k_proj(x), self.n_kv_heads)
v = self._split_heads(self.v_proj(x), self.n_kv_heads) v = self._split_heads(self.v_proj(x), self.n_kv_heads)
@@ -200,22 +167,14 @@ class GQA(nn.Module):
if self.use_qk_norm: if self.use_qk_norm:
q, k = self.q_norm(q), self.k_norm(k) q, k = self.q_norm(q), self.k_norm(k)
if kv_cache is not None: if paged_cache is not None:
k_cache, v_cache = kv_cache paged_cache.write(self.layer_id, start_pos, k, v)
k, v = paged_cache.gather(self.layer_id)
# copy to cache
k_cache[:bsz, start_pos : start_pos + seq_len, self.layer_id] = k
v_cache[:bsz, start_pos : start_pos + seq_len, self.layer_id] = v
# get cache
k = k_cache[:bsz, : start_pos + seq_len, self.layer_id]
v = v_cache[:bsz, : start_pos + seq_len, self.layer_id]
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep) k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)
# (bsz, seq_len, n_heads, head_dim) -> (bsz, n_heads, seq_len, head_dim) # (bsz, seq_len, n_heads, head_dim) -> (bsz, n_heads, seq_len, head_dim)
q, k, v = q.permute(0, 2, 1, 3), k.permute(0, 2, 1, 3), v.permute(0, 2, 1, 3) q, k, v = q.permute(0, 2, 1, 3), k.permute(0, 2, 1, 3), v.permute(0, 2, 1, 3)
# (bsz, n_heads, seq_len, head_dim) - > (bsz, seq_len, n_heads*head_dim)
sdqa_out = ( sdqa_out = (
F.scaled_dot_product_attention(q, k, v, mask, is_causal=is_causal) F.scaled_dot_product_attention(q, k, v, mask, is_causal=is_causal)
.permute(0, 2, 1, 3) .permute(0, 2, 1, 3)
@@ -227,7 +186,6 @@ class GQA(nn.Module):
sdqa_out = sdqa_out * F.sigmoid(self.gate(x)) sdqa_out = sdqa_out * F.sigmoid(self.gate(x))
out = self.o_proj(sdqa_out) out = self.o_proj(sdqa_out)
return out return out
@@ -260,7 +218,7 @@ class MLA(nn.Module):
self.kv_a_proj = Linear(dim, kv_lora_rank, bias=False) self.kv_a_proj = Linear(dim, kv_lora_rank, bias=False)
self.kv_norm = RMSNorm(kv_lora_rank, norm_eps) self.kv_norm = RMSNorm(kv_lora_rank, norm_eps)
# KV (k_nope, k_rope, v) # fused KV: (k_nope, k_rope, v)
self.kv_b_proj = Linear( self.kv_b_proj = Linear(
kv_lora_rank, kv_lora_rank,
n_kv_heads * (self.head_dim + qk_rope_head_dim + self.head_dim), n_kv_heads * (self.head_dim + qk_rope_head_dim + self.head_dim),
@@ -276,7 +234,7 @@ class MLA(nn.Module):
x: Tensor, x: Tensor,
rotary_emb: Tuple[Tensor, Tensor], rotary_emb: Tuple[Tensor, Tensor],
mask: Tensor = None, mask: Tensor = None,
kv_cache: Optional[Tuple[Tensor, Tensor]] = None, paged_cache: Optional[CacheView] = None,
start_pos: int = 0, start_pos: int = 0,
) -> Tensor: ) -> Tensor:
bsz, seq_len, _ = x.size() bsz, seq_len, _ = x.size()
@@ -305,12 +263,9 @@ class MLA(nn.Module):
q = torch.cat([q_nope, q_rope], dim=-1) q = torch.cat([q_nope, q_rope], dim=-1)
k = torch.cat([k_nope, k_rope], dim=-1) k = torch.cat([k_nope, k_rope], dim=-1)
if kv_cache is not None: if paged_cache is not None:
k_cache, v_cache = kv_cache paged_cache.write(self.layer_id, start_pos, k, v)
k_cache[:bsz, start_pos : start_pos + seq_len, self.layer_id] = k k, v = paged_cache.gather(self.layer_id)
v_cache[:bsz, start_pos : start_pos + seq_len, self.layer_id] = v
k = k_cache[:bsz, : start_pos + seq_len, self.layer_id]
v = v_cache[:bsz, : start_pos + seq_len, self.layer_id]
q = q.permute(0, 2, 1, 3) q = q.permute(0, 2, 1, 3)
k = k.permute(0, 2, 1, 3) k = k.permute(0, 2, 1, 3)
@@ -323,7 +278,6 @@ class MLA(nn.Module):
attn_out = attn_out * F.sigmoid(self.gate(x)) attn_out = attn_out * F.sigmoid(self.gate(x))
out = self.o_proj(attn_out) out = self.o_proj(attn_out)
return out return out
@@ -358,18 +312,19 @@ class DecoderBlock(nn.Module):
x: Tensor, x: Tensor,
rotary_emb: Tuple[Tensor, Tensor], rotary_emb: Tuple[Tensor, Tensor],
attention_mask: Optional[Tensor] = None, attention_mask: Optional[Tensor] = None,
kv_cache: Optional[Tuple[Tensor, Tensor]] = None, paged_cache: Optional[CacheView] = None,
start_pos: int = 0, start_pos: int = 0,
) -> Tensor: ) -> Tensor:
# attention
attn_output = self.attention( attn_output = self.attention(
self.input_norm(x), rotary_emb, attention_mask, kv_cache, start_pos self.input_norm(x),
rotary_emb,
attention_mask,
paged_cache,
start_pos,
) )
x = attn_output + x x = attn_output + x
# feed forward
x = self.mlp(self.post_attention_norm(x)) + x x = self.mlp(self.post_attention_norm(x)) + x
return x return x
+8 -27
View File
@@ -1,10 +1,11 @@
from typing import Any, Mapping, Optional, Tuple from typing import Any, Mapping, Optional
import torch import torch
import torch.nn as nn import torch.nn as nn
from torch import Tensor from torch import Tensor
from astrai.config.model_config import ModelConfig from astrai.config.model_config import ModelConfig
from astrai.inference.cache import CacheView
from astrai.model.automodel import AutoModel from astrai.model.automodel import AutoModel
from astrai.model.module import ( from astrai.model.module import (
DecoderBlock, DecoderBlock,
@@ -21,39 +22,25 @@ def process_attention_mask(
start_pos: int = 0, start_pos: int = 0,
is_causal: bool = False, is_causal: bool = False,
) -> Tensor: ) -> Tensor:
""" """Build 4D attention mask from 2D seq_mask, with optional causal masking."""
Create attention mask for GQA
Args:
seq_mask (Tensor): A tensor indicating whether each position is valid or not.
input_tensor (Tensor): The input tensor.
start_pos (int): The starting position of the sequence.
is_causal (bool): Whether the attention is causal or not.
Returns:
Tensor: The attention mask tensor.
"""
device = input_tensor.device device = input_tensor.device
dtype = input_tensor.dtype dtype = input_tensor.dtype
seq_len = input_tensor.size(1) seq_len = input_tensor.size(1)
if seq_mask is None: if seq_mask is None:
if start_pos != 0: if start_pos != 0:
# for single prompt chat
seq_mask = torch.ones((1, seq_len), dtype=torch.bool, device=device) seq_mask = torch.ones((1, seq_len), dtype=torch.bool, device=device)
else: else:
return None return None
if seq_mask.dim() > 2: if seq_mask.dim() > 2:
# shape (bsz, seq_len) or (bsz,n_heads, seq_len, seq_len + start_pos)
# if ndim > 2, it's 4D tensor
return seq_mask return seq_mask
batch_size = seq_mask.size(0) batch_size = seq_mask.size(0)
seq_mask = seq_mask[:, : start_pos + seq_len].to(device=device, dtype=torch.bool) seq_mask = seq_mask[:, : start_pos + seq_len].to(device=device, dtype=torch.bool)
# (bsz, start_pos + seq_len)
expanded_mask = seq_mask.unsqueeze(1).expand( expanded_mask = seq_mask.unsqueeze(1).expand(
batch_size, seq_len, start_pos + seq_len batch_size, seq_len, start_pos + seq_len
) )
# (bsz, seq_len, start_pos + seq_len)
if is_causal: if is_causal:
expanded_mask = torch.tril(expanded_mask, diagonal=start_pos) expanded_mask = torch.tril(expanded_mask, diagonal=start_pos)
@@ -62,16 +49,13 @@ def process_attention_mask(
attention_mask = attention_mask.masked_fill_( attention_mask = attention_mask.masked_fill_(
~expanded_mask, -torch.finfo(dtype).max / 2 ~expanded_mask, -torch.finfo(dtype).max / 2
).unsqueeze(1) ).unsqueeze(1)
# (bsz, 1, seq_len, seq_len + start_pos)
return attention_mask return attention_mask
@AutoModel.register("transformer") @AutoModel.register("transformer")
class Transformer(AutoModel): class Transformer(AutoModel):
""" """Transformer language model with paged KV cache."""
Transformer language model.
"""
def __init__(self, config: ModelConfig): def __init__(self, config: ModelConfig):
super().__init__(config) super().__init__(config)
@@ -114,18 +98,15 @@ class Transformer(AutoModel):
lm_head_key = "lm_head.weight" lm_head_key = "lm_head.weight"
embed_key = "embed_tokens.weight" embed_key = "embed_tokens.weight"
# Make a copy to avoid modifying the original state_dict
state_dict = dict(state_dict) state_dict = dict(state_dict)
if self.config.tie_weight: if self.config.tie_weight:
# same tensor # same tensor for embed and lm_head
if embed_key in state_dict: if embed_key in state_dict:
state_dict[lm_head_key] = state_dict[embed_key] state_dict[lm_head_key] = state_dict[embed_key]
else: else:
# If lm_head.weight exists in checkpoint, use it directly
# If not, copy from embed_tokens.weight
if lm_head_key not in state_dict and embed_key in state_dict: if lm_head_key not in state_dict and embed_key in state_dict:
# use clone to avoid sharing the same tensor # clone to avoid sharing gradients
state_dict[lm_head_key] = torch.clone(state_dict[embed_key]) state_dict[lm_head_key] = torch.clone(state_dict[embed_key])
return super().load_state_dict(state_dict, strict, assign) return super().load_state_dict(state_dict, strict, assign)
@@ -146,7 +127,7 @@ class Transformer(AutoModel):
self, self,
input_ids: Tensor, input_ids: Tensor,
input_mask: Optional[Tensor] = None, input_mask: Optional[Tensor] = None,
persistent_key_values: Optional[Tuple[Tensor, Tensor]] = None, paged_cache: Optional[CacheView] = None,
start_pos: int = 0, start_pos: int = 0,
) -> Tensor: ) -> Tensor:
assert input_ids.ndim == 2 assert input_ids.ndim == 2
@@ -157,7 +138,7 @@ class Transformer(AutoModel):
attn_mask = process_attention_mask(input_mask, x, start_pos, is_causal=True) attn_mask = process_attention_mask(input_mask, x, start_pos, is_causal=True)
for layer in self.layers: for layer in self.layers:
x = layer(x, rotary_emb, attn_mask, persistent_key_values, start_pos) x = layer(x, rotary_emb, attn_mask, paged_cache, start_pos)
hidden_states = self.norm(x) hidden_states = self.norm(x)
logits = self.lm_head(hidden_states) logits = self.lm_head(hidden_states)
+5 -14
View File
@@ -1,7 +1,7 @@
import os import os
from contextlib import contextmanager from contextlib import contextmanager
from functools import wraps from functools import wraps
from typing import Callable, List, Optional from typing import Callable
import torch import torch
import torch.distributed as dist import torch.distributed as dist
@@ -34,7 +34,6 @@ def setup_parallel(
master_addr: str = "localhost", master_addr: str = "localhost",
master_port: str = "29500", master_port: str = "29500",
device_type: str = "cuda", device_type: str = "cuda",
device_ids: Optional[List[int]] = None,
): ):
if dist.is_available() and dist.is_initialized(): if dist.is_available() and dist.is_initialized():
@@ -45,15 +44,10 @@ def setup_parallel(
yield None yield None
return return
if device_ids is None: device_id = torch.device(device_type, rank)
device_ids = [i for i in range(world_size)]
rank = device_ids[rank % len(device_ids)]
device_id = torch.device(device_type, device_ids[rank])
os.environ["MASTER_ADDR"] = master_addr os.environ["MASTER_ADDR"] = master_addr
os.environ["MASTER_PORT"] = master_port os.environ["MASTER_PORT"] = master_port
os.environ["LOCAL_RANK"] = str(rank) os.environ["LOCAL_RANK"] = str(rank)
os.environ["WORLD_SIZE"] = str(world_size) os.environ["WORLD_SIZE"] = str(world_size)
os.environ["LOCAL_DEVICE"] = str(device_id) os.environ["LOCAL_DEVICE"] = str(device_id)
@@ -103,7 +97,6 @@ def wrapper_spawn_func(
master_addr: str, master_addr: str,
master_port: str, master_port: str,
device_type: str, device_type: str,
device_ids: List[int],
func: Callable, func: Callable,
kwargs: dict, kwargs: dict,
): ):
@@ -115,7 +108,6 @@ def wrapper_spawn_func(
master_addr=master_addr, master_addr=master_addr,
master_port=master_port, master_port=master_port,
device_type=device_type, device_type=device_type,
device_ids=device_ids,
): ):
func(**kwargs) func(**kwargs)
@@ -131,7 +123,6 @@ def spawn_parallel_fn(
master_addr: str = "localhost", master_addr: str = "localhost",
master_port: str = "29500", master_port: str = "29500",
device_type: str = "cuda", device_type: str = "cuda",
device_ids: Optional[List[int]] = None,
**kwargs, **kwargs,
): ):
# clear environment variables # clear environment variables
@@ -147,8 +138,9 @@ def spawn_parallel_fn(
del os.environ[key] del os.environ[key]
if world_size == 1: if world_size == 1:
device_ids = device_ids or [0] device_id = torch.device(device_type, 0)
device_id = torch.device(device_type, device_ids[0]) os.environ["LOCAL_RANK"] = "0"
os.environ["WORLD_SIZE"] = "1"
os.environ["LOCAL_DEVICE"] = str(device_id) os.environ["LOCAL_DEVICE"] = str(device_id)
func(**kwargs) func(**kwargs)
@@ -160,7 +152,6 @@ def spawn_parallel_fn(
master_addr, master_addr,
master_port, master_port,
device_type, device_type,
device_ids,
func, func,
kwargs, kwargs,
) )
+11 -1
View File
@@ -1,7 +1,7 @@
import json import json
import os import os
from pathlib import Path from pathlib import Path
from typing import Any, Dict, List from typing import Any, Dict, List, Optional
import h5py import h5py
import safetensors.torch as st import safetensors.torch as st
@@ -54,10 +54,12 @@ class Checkpoint:
state_dict: Dict[str, Any], state_dict: Dict[str, Any],
epoch: int = 0, epoch: int = 0,
iteration: int = 0, iteration: int = 0,
extra: Optional[Dict[str, Any]] = None,
): ):
self.state_dict = state_dict self.state_dict = state_dict
self.epoch = epoch self.epoch = epoch
self.iteration = iteration self.iteration = iteration
self.extra = extra or {}
def save( def save(
self, self,
@@ -77,6 +79,8 @@ class Checkpoint:
json.dump(meta, f, indent=2) json.dump(meta, f, indent=2)
st.save_file(self.state_dict, save_path / "state_dict.safetensors") st.save_file(self.state_dict, save_path / "state_dict.safetensors")
if self.extra:
torch.save(self.extra, save_path / "extra.pt")
@classmethod @classmethod
def load( def load(
@@ -99,8 +103,14 @@ class Checkpoint:
state_dict = st.load_file(save_path / "state_dict.safetensors") state_dict = st.load_file(save_path / "state_dict.safetensors")
extra = None
extra_path = save_path / "extra.pt"
if extra_path.exists():
extra = torch.load(extra_path, map_location="cpu", weights_only=False)
return cls( return cls(
state_dict=state_dict, state_dict=state_dict,
epoch=meta["epoch"], epoch=meta["epoch"],
iteration=meta["iteration"], iteration=meta["iteration"],
extra=extra,
) )
+5
View File
@@ -64,6 +64,11 @@ class AutoTokenizer:
save_path: Path to save the tokenizer save_path: Path to save the tokenizer
""" """
if self._tokenizer is None:
raise RuntimeError(
"Tokenizer not initialized. Load or create a tokenizer first."
)
save_path = Path(save_path) save_path = Path(save_path)
save_path.mkdir(parents=True, exist_ok=True) save_path.mkdir(parents=True, exist_ok=True)
+17 -5
View File
@@ -265,7 +265,9 @@ class DPOStrategy(BaseStrategy):
class GRPOStrategy(BaseStrategy): class GRPOStrategy(BaseStrategy):
"""Group Relative Policy Optimization strategy. """Group Relative Policy Optimization strategy.
Implements GRPO with clipping and KL penalty. On-policy GRPO following DeepSeek-R1: the policy model is updated while
a frozen ref_model stores the old-policy log-probs. ratio = exp(logπ_θ - logπ_ref),
clipped PPO objective. Call ``sync_ref_model()`` after each data-generation round.
""" """
def __init__( def __init__(
@@ -276,6 +278,7 @@ class GRPOStrategy(BaseStrategy):
kl_coef: float = 0.01, kl_coef: float = 0.01,
group_size: int = 4, group_size: int = 4,
reduction: str = "mean", reduction: str = "mean",
sync_interval: int = 200,
**kwargs, **kwargs,
): ):
super().__init__(model, device, **kwargs) super().__init__(model, device, **kwargs)
@@ -284,8 +287,19 @@ class GRPOStrategy(BaseStrategy):
self.kl_coef = kl_coef self.kl_coef = kl_coef
self.group_size = group_size self.group_size = group_size
self.reduction = reduction self.reduction = reduction
self.sync_interval = sync_interval
self._step = 0
def sync_ref_model(self):
"""Copy current model weights to ref model."""
ref_state = self.model.state_dict()
self.ref_model.load_state_dict(ref_state)
def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor: def compute_loss(self, batch: Dict[str, Tensor]) -> Tensor:
self._step += 1
if self._step % self.sync_interval == 0:
self.sync_ref_model()
batch = move_to_device(batch, self.device) batch = move_to_device(batch, self.device)
prompts = batch["prompts"] prompts = batch["prompts"]
responses = batch["responses"] responses = batch["responses"]
@@ -297,7 +311,6 @@ class GRPOStrategy(BaseStrategy):
masks_flat = masks.view(-1, response_len) masks_flat = masks.view(-1, response_len)
prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1) prompt_expanded = prompts.unsqueeze(1).repeat(1, group_size, 1).flatten(0, 1)
# Shape: (batch_size * group_size, seq_len + response_len)
full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1) full_sequences = torch.cat([prompt_expanded, responses_flat], dim=-1)
full_masks = torch.cat([torch.ones_like(prompt_expanded), masks_flat], dim=-1) full_masks = torch.cat([torch.ones_like(prompt_expanded), masks_flat], dim=-1)
@@ -312,14 +325,13 @@ class GRPOStrategy(BaseStrategy):
) )
log_probs_ref = log_probs_ref.view(batch_size, group_size) log_probs_ref = log_probs_ref.view(batch_size, group_size)
# Compute advantages from rewards with normalization
eps = torch.finfo(log_probs_policy.dtype).eps eps = torch.finfo(log_probs_policy.dtype).eps
mean = rewards.mean(dim=-1, keepdim=True) mean = rewards.mean(dim=-1, keepdim=True)
std = rewards.std(dim=-1, keepdim=True) std = rewards.std(dim=-1, keepdim=True)
advantages = (rewards - mean) / (std + eps) advantages = (rewards - mean) / (std + eps)
# PPO-style clipped surrogate objective ratio = torch.exp(log_probs_policy - log_probs_ref)
ratio = torch.exp(0) # Off-policy: policy_model = old_model
surr1 = ratio * advantages surr1 = ratio * advantages
surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * advantages surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * advantages
+7 -1
View File
@@ -121,11 +121,13 @@ class CheckpointCallback(TrainCallback):
interval: int, interval: int,
weight_only: bool = False, weight_only: bool = False,
state_dict_fn: Optional[Callable[[nn.Module], dict]] = None, state_dict_fn: Optional[Callable[[nn.Module], dict]] = None,
save_extra_fn: Optional[Callable[["TrainContext"], dict]] = None,
): ):
self.save_dir = save_dir self.save_dir = save_dir
self.interval = interval self.interval = interval
self.weight_only = weight_only self.weight_only = weight_only
self.state_dict_fn = state_dict_fn self.state_dict_fn = state_dict_fn
self.save_extra_fn = save_extra_fn
self.last_ckpt_iter = 0 self.last_ckpt_iter = 0
@only_on_rank(0) @only_on_rank(0)
@@ -139,8 +141,12 @@ class CheckpointCallback(TrainCallback):
else context.model.state_dict() else context.model.state_dict()
) )
extra = self.save_extra_fn(context) if self.save_extra_fn else None
context.checkpoint = Checkpoint( context.checkpoint = Checkpoint(
state_dict=state_dict, epoch=context.epoch, iteration=context.iteration state_dict=state_dict,
epoch=context.epoch,
iteration=context.iteration,
extra=extra,
) )
context.checkpoint.save(save_path) context.checkpoint.save(save_path)
+50 -48
View File
@@ -1,5 +1,5 @@
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Optional, Self from typing import Callable, Optional, Self
import torch.nn as nn import torch.nn as nn
from torch.optim import Optimizer from torch.optim import Optimizer
@@ -32,68 +32,70 @@ class TrainContext:
class TrainContextBuilder: class TrainContextBuilder:
def __init__(self, config: TrainConfig): def __init__(
self,
config: TrainConfig,
load_extra_fn: Optional[Callable[[dict, "TrainContext"], None]] = None,
):
self.config = config self.config = config
self._context = TrainContext( self._checkpoint: Optional[Checkpoint] = None
model=config.model, self._load_extra_fn = load_extra_fn
def with_checkpoint(self, checkpoint: Optional[Checkpoint]) -> Self:
self._checkpoint = checkpoint
return self
def build(self) -> TrainContext:
context = TrainContext(
model=self.config.model,
world_size=get_world_size(), world_size=get_world_size(),
rank=get_rank(), rank=get_rank(),
) )
device = get_current_device() device = get_current_device()
self._context.model = self._context.model.to(device=device) context.model = context.model.to(device=device)
if self.config.nprocs > 1: if self.config.nprocs > 1 and self.config.parallel_wrapper:
fn = self.config.parallel_wrapper context.model = self.config.parallel_wrapper(context.model)
self._context.model = fn(self._context.model)
self._context.optimizer = self.config.optimizer_fn(self._context.model) if self._checkpoint is not None:
self._context.scheduler = self.config.scheduler_fn(self._context.optimizer) context.epoch = max(self._checkpoint.epoch, self.config.start_epoch)
context.iteration = max(self._checkpoint.iteration, self.config.start_batch)
def with_checkpoint(self, checkpoint: Optional[Checkpoint]) -> Self: context.model.load_state_dict(self._checkpoint.state_dict)
if checkpoint is None: context.checkpoint = self._checkpoint
checkpoint = Checkpoint(
state_dict=self._context.model.state_dict(),
)
else: else:
# resume from the assigned checkpoint or assigned iteration context.checkpoint = Checkpoint(
self._context.epoch = max(checkpoint.epoch, self.config.start_epoch) state_dict=context.model.state_dict(),
self._context.iteration = max(checkpoint.iteration, self.config.start_batch) )
self._context.model.load_state_dict(checkpoint.state_dict)
self._context.checkpoint = checkpoint context.optimizer = self.config.optimizer_fn(context.model)
return self context.scheduler = self.config.scheduler_fn(context.optimizer)
def with_dataloader(self) -> Self: if self._checkpoint and self._checkpoint.extra and self._load_extra_fn:
# fix: change batch level iteration to sample level offset self._load_extra_fn(self._checkpoint.extra, context)
config = self.config
sampler_offset = self._context.iteration * config.batch_size cfg = self.config
resumeable_sampler = ResumableDistributedSampler( sampler_offset = context.iteration * cfg.batch_size
data_source=config.dataset, sampler = ResumableDistributedSampler(
start_epoch=self._context.epoch, data_source=cfg.dataset,
start_epoch=context.epoch,
start_iter=sampler_offset, start_iter=sampler_offset,
seed=config.random_seed, seed=cfg.random_seed,
)
context.dataloader = DataLoader(
cfg.dataset,
batch_size=cfg.batch_size,
sampler=sampler,
num_workers=cfg.num_workers,
pin_memory=cfg.pin_memory,
prefetch_factor=cfg.prefetch_factor,
) )
dataloader = DataLoader( context.strategy = StrategyFactory.create(
config.dataset, model=context.model,
batch_size=config.batch_size,
sampler=resumeable_sampler,
num_workers=config.num_workers,
pin_memory=config.pin_memory,
prefetch_factor=config.prefetch_factor,
)
self._context.dataloader = dataloader
return self
def with_strategy(self) -> Self:
self._context.strategy = StrategyFactory.create(
model=self._context.model,
train_type=self.config.strategy, train_type=self.config.strategy,
device=get_current_device(), device=device,
**self.config.extra_kwargs, **self.config.extra_kwargs,
) )
return self
def build(self) -> TrainContext: return context
return self._context
+4 -8
View File
@@ -35,11 +35,7 @@ class Trainer:
def _build_context(self, checkpoint: Optional[Checkpoint]) -> TrainContext: def _build_context(self, checkpoint: Optional[Checkpoint]) -> TrainContext:
return ( return (
TrainContextBuilder(self.train_config) TrainContextBuilder(self.train_config).with_checkpoint(checkpoint).build()
.with_checkpoint(checkpoint)
.with_dataloader()
.with_strategy()
.build()
) )
def _call_callbacks(self, method_name: str, context: TrainContext): def _call_callbacks(self, method_name: str, context: TrainContext):
@@ -57,7 +53,6 @@ class Trainer:
master_addr=config.master_addr, master_addr=config.master_addr,
master_port=config.master_port, master_port=config.master_port,
device_type=config.device_type, device_type=config.device_type,
device_ids=config.device_ids,
checkpoint=checkpoint, checkpoint=checkpoint,
) )
@@ -72,8 +67,9 @@ class Trainer:
context.epoch = epoch context.epoch = epoch
self._call_callbacks("on_epoch_begin", context) self._call_callbacks("on_epoch_begin", context)
accumulation_steps = max(self.train_config.accumulation_steps, 1)
for batch in context.dataloader: for batch in context.dataloader:
if context.iteration % self.train_config.accumulation_steps == 0: if context.iteration % accumulation_steps == 0:
# 2. step # 2. step
self._call_callbacks("on_step_begin", context) self._call_callbacks("on_step_begin", context)
context.optimizer.step() context.optimizer.step()
@@ -87,7 +83,7 @@ class Trainer:
context.iteration += 1 context.iteration += 1
# to make the loss normalized by accumulation steps # to make the loss normalized by accumulation steps
stand_loss = loss / self.train_config.accumulation_steps stand_loss = loss / accumulation_steps
stand_loss.backward() stand_loss.backward()
self._call_callbacks("on_batch_end", context) self._call_callbacks("on_batch_end", context)
+42
View File
@@ -0,0 +1,42 @@
services:
server:
build: .
image: astrai:latest
ports:
- "8000:8000"
volumes:
- ./params:/app/params:ro
- ./checkpoints:/app/checkpoints
command: python -m scripts.tools.server --port 8000 --device cuda
deploy:
resources:
reservations:
devices:
- driver: nvidia
count: 1
capabilities: [gpu]
healthcheck:
test: ["CMD", "curl", "-f", "http://localhost:8000/health"]
interval: 30s
timeout: 10s
retries: 3
start_period: 60s
restart: unless-stopped
server-cpu:
profiles: [cpu]
build: .
image: astrai:latest
ports:
- "8000:8000"
volumes:
- ./params:/app/params:ro
- ./checkpoints:/app/checkpoints
command: python -m scripts.tools.server --port 8000 --device cpu
healthcheck:
test: ["CMD", "curl", "-f", "http://localhost:8000/health"]
interval: 30s
timeout: 10s
retries: 3
start_period: 120s
restart: unless-stopped
+1 -1
View File
@@ -15,7 +15,7 @@ def chat():
tokenizer = AutoTokenizer.from_pretrained(PARAMETER_ROOT) tokenizer = AutoTokenizer.from_pretrained(PARAMETER_ROOT)
model.to(device="cuda", dtype=torch.bfloat16) model.to(device="cuda", dtype=torch.bfloat16)
messages = [] messages = [{"role": "system", "content": "You are a helpful assistant."}]
engine = InferenceEngine(model=model, tokenizer=tokenizer) engine = InferenceEngine(model=model, tokenizer=tokenizer)
while True: while True:
+75 -57
View File
@@ -1,9 +1,14 @@
"""Benchmark Transformer with PagedCache (replaces old persistent_key_values)."""
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, Dict from typing import Any, Dict
import torch import torch
from torch import Tensor
from astrai.model.transformer import ModelConfig, Transformer from astrai.config import ModelConfig
from astrai.inference.cache import PagedCache
from astrai.model.transformer import Transformer
@dataclass @dataclass
@@ -19,27 +24,25 @@ class GenerationBenchmark:
self, self,
config: ModelConfig, config: ModelConfig,
device: str = "cuda", device: str = "cuda",
dtype: torch.dtype = torch.float16, dtype: torch.dtype = torch.bfloat16,
page_size: int = 128,
): ):
self.config = config self.config = config
self.device = device self.device = device
self.dtype = dtype self.dtype = dtype
self.model = Transformer(config).to(device=device, dtype=dtype) self.model = Transformer(config).to(device=device, dtype=dtype)
self.model.eval() self.model.eval()
head_dim = config.dim // config.n_heads
def _initialize_kv_cache(self, batch_size: int) -> list: n_pages = (config.max_len * 4 + page_size - 1) // page_size
"""初始化KV缓存""" self._page_cache = PagedCache(
config = self.config
shape = (
batch_size,
config.max_len,
config.n_layers, config.n_layers,
n_pages,
page_size,
config.n_kv_heads, config.n_kv_heads,
config.dim // config.n_heads, head_dim,
device,
dtype,
) )
k_cache = torch.zeros(shape, device=self.device, dtype=self.dtype)
v_cache = torch.zeros(shape, device=self.device, dtype=self.dtype)
return (k_cache, v_cache)
def _prepare_inputs(self, batch_size: int, prompt_length: int, total_length: int): def _prepare_inputs(self, batch_size: int, prompt_length: int, total_length: int):
prompt_ids = torch.randint( prompt_ids = torch.randint(
@@ -49,7 +52,6 @@ class GenerationBenchmark:
device=self.device, device=self.device,
dtype=torch.long, dtype=torch.long,
) )
gen_ids = torch.randint( gen_ids = torch.randint(
low=0, low=0,
high=self.config.vocab_size, high=self.config.vocab_size,
@@ -57,9 +59,11 @@ class GenerationBenchmark:
device=self.device, device=self.device,
dtype=torch.long, dtype=torch.long,
) )
return prompt_ids, gen_ids return prompt_ids, gen_ids
def _make_mask(self, batch_size: int, seq_len: int) -> Tensor:
return torch.ones(batch_size, seq_len, dtype=torch.bool, device=self.device)
@torch.inference_mode() @torch.inference_mode()
def run_prefill_benchmark( def run_prefill_benchmark(
self, self,
@@ -67,13 +71,11 @@ class GenerationBenchmark:
prompt_length: int = 512, prompt_length: int = 512,
num_trials: int = 10, num_trials: int = 10,
) -> BenchmarkResult: ) -> BenchmarkResult:
for _ in range(3): for _ in range(3):
prompt_ids, _ = self._prepare_inputs( prompt_ids, _ = self._prepare_inputs(
batch_size, prompt_length, prompt_length batch_size, prompt_length, prompt_length
) )
_ = self.model(prompt_ids) _ = self.model(prompt_ids)
torch.cuda.synchronize() torch.cuda.synchronize()
total_time = 0.0 total_time = 0.0
@@ -83,20 +85,20 @@ class GenerationBenchmark:
prompt_ids, _ = self._prepare_inputs( prompt_ids, _ = self._prepare_inputs(
batch_size, prompt_length, prompt_length batch_size, prompt_length, prompt_length
) )
start_event = torch.cuda.Event(enable_timing=True) start = torch.cuda.Event(enable_timing=True)
end_event = torch.cuda.Event(enable_timing=True) end = torch.cuda.Event(enable_timing=True)
start_event.record() start.record()
_ = self.model(prompt_ids) _ = self.model(prompt_ids)
end_event.record() end.record()
torch.cuda.synchronize() torch.cuda.synchronize()
trial_time = start_event.elapsed_time(end_event) / 1000 trial_time = start.elapsed_time(end) / 1000
total_time += trial_time total_time += trial_time
print( print(
f" Trial {trial + 1}/{num_trials}: {prompt_length} tokens in {trial_time:.3f}s " f" Trial {trial + 1}/{num_trials}: {prompt_length} tokens in {trial_time:.3f}s "
f"({prompt_length / trial_time:.1f} tokens/s)" f"({prompt_length / trial_time:.1f} tok/s)"
) )
return BenchmarkResult( return BenchmarkResult(
@@ -107,7 +109,7 @@ class GenerationBenchmark:
"benchmark_type": "prefill", "benchmark_type": "prefill",
"batch_size": batch_size, "batch_size": batch_size,
"prompt_length": prompt_length, "prompt_length": prompt_length,
"dtype": self.dtype, "dtype": str(self.dtype),
"device": self.device, "device": self.device,
}, },
) )
@@ -120,41 +122,62 @@ class GenerationBenchmark:
gen_length: int = 128, gen_length: int = 128,
num_trials: int = 5, num_trials: int = 5,
) -> BenchmarkResult: ) -> BenchmarkResult:
total_time = 0.0 total_time = 0.0
total_tokens = batch_size * gen_length * num_trials total_tokens = batch_size * gen_length * num_trials
page_size = self._page_cache.page_size
for trial in range(num_trials): for trial in range(num_trials):
prompt_ids, gen_ids = self._prepare_inputs( prompt_ids, gen_ids = self._prepare_inputs(
batch_size, prompt_length, prompt_length + gen_length batch_size,
prompt_length,
prompt_length + gen_length,
)
n_pages = (prompt_length + gen_length + page_size - 1) // page_size
pages = self._page_cache.alloc_n(n_pages * batch_size)
page_table = torch.tensor(
[pages[i * n_pages : (i + 1) * n_pages] for i in range(batch_size)],
dtype=torch.long,
device=self.device,
)
cv = self._page_cache.bind(page_table, total_len=prompt_length)
_ = self.model(
prompt_ids,
paged_cache=cv,
start_pos=0,
input_mask=self._make_mask(batch_size, prompt_length),
) )
kv_cache = self._initialize_kv_cache(batch_size)
_ = self.model(prompt_ids, persistent_key_values=kv_cache, start_pos=0)
torch.cuda.synchronize() torch.cuda.synchronize()
start_event = torch.cuda.Event(enable_timing=True) start = torch.cuda.Event(enable_timing=True)
end_event = torch.cuda.Event(enable_timing=True) end = torch.cuda.Event(enable_timing=True)
start_event.record()
start.record()
current_pos = prompt_length current_pos = prompt_length
for i in range(gen_length): for i in range(gen_length):
input_token = gen_ids[:, i : i + 1] input_token = gen_ids[:, i : i + 1]
cv = self._page_cache.bind(page_table, total_len=current_pos + 1)
_ = self.model( _ = self.model(
input_token, persistent_key_values=kv_cache, start_pos=current_pos input_token,
paged_cache=cv,
start_pos=current_pos,
input_mask=self._make_mask(batch_size, 1),
) )
current_pos += 1 current_pos += 1
end.record()
end_event.record()
torch.cuda.synchronize() torch.cuda.synchronize()
trial_time = start_event.elapsed_time(end_event) / 1000 trial_time = start.elapsed_time(end) / 1000
total_time += trial_time total_time += trial_time
for idx in pages:
self._page_cache.free(idx)
print( print(
f" Trial {trial + 1}/{num_trials}: {gen_length} tokens in {trial_time:.3f}s " f" Trial {trial + 1}/{num_trials}: {gen_length} tokens in {trial_time:.3f}s "
f"({gen_length / trial_time:.1f} tokens/s)" f"({gen_length / trial_time:.1f} tok/s)"
) )
return BenchmarkResult( return BenchmarkResult(
@@ -166,31 +189,21 @@ class GenerationBenchmark:
"batch_size": batch_size, "batch_size": batch_size,
"prompt_length": prompt_length, "prompt_length": prompt_length,
"gen_length": gen_length, "gen_length": gen_length,
"dtype": self.dtype, "dtype": str(self.dtype),
"device": self.device, "device": self.device,
}, },
) )
def print_benchmark_result(result: BenchmarkResult): def print_benchmark_result(result: BenchmarkResult):
"""打印基准测试结果""" btype = result.metadata["benchmark_type"]
benchmark_type = result.metadata["benchmark_type"] print(f"\n{' ' + btype.upper() + ' Benchmark ':-^80}")
print(f"\n{' ' + benchmark_type.upper().replace('_', ' ') + ' Benchmark ':-^80}")
print(f"Total Tokens Processed: {result.total_tokens:,}") print(f"Total Tokens Processed: {result.total_tokens:,}")
print(f"Time Consumed: {result.total_time:.3f}s") print(f"Time Consumed: {result.total_time:.3f}s")
print(f"Throughput: {result.tokens_per_second:,.1f} tokens/s") print(f"Throughput: {result.tokens_per_second:,.1f} tok/s")
for k, v in result.metadata.items():
if benchmark_type == "prefill": if k != "benchmark_type":
print( print(f"{k.replace('_', ' ').title()}: {v}")
f"Batch Size: {result.metadata['batch_size']} | Prompt Length: {result.metadata['prompt_length']}"
)
elif benchmark_type == "decoding":
print(
f"Batch Size: {result.metadata['batch_size']} | Gen Length: {result.metadata['gen_length']}"
)
print(f"Device: {result.metadata['device']} | Dtype: {result.metadata['dtype']}")
print("-" * 80) print("-" * 80)
@@ -209,15 +222,20 @@ if __name__ == "__main__":
benchmark = GenerationBenchmark(config) benchmark = GenerationBenchmark(config)
print("=" * 80) print("=" * 80)
print("Running Transformer Generation Benchmark") print("Running Transformer Generation Benchmark (PagedCache)")
print("=" * 80) print("=" * 80)
prefill_result = benchmark.run_prefill_benchmark( prefill_result = benchmark.run_prefill_benchmark(
batch_size=4, prompt_length=512, num_trials=5 batch_size=4,
prompt_length=512,
num_trials=5,
) )
print_benchmark_result(prefill_result) print_benchmark_result(prefill_result)
gen_result = benchmark.run_decoding_benchmark( gen_result = benchmark.run_decoding_benchmark(
batch_size=4, prompt_length=512, gen_length=128, num_trials=5 batch_size=4,
prompt_length=512,
gen_length=128,
num_trials=5,
) )
print_benchmark_result(gen_result) print_benchmark_result(gen_result)
+4 -4
View File
@@ -9,7 +9,7 @@ from astrai.tokenize import AutoTokenizer
def processor( def processor(
model_dir: str, param_path: str,
input_json_file: str, input_json_file: str,
output_json_file: str, output_json_file: str,
temperature: float, temperature: float,
@@ -20,8 +20,8 @@ def processor(
max_tokens: int, max_tokens: int,
): ):
# Load model and tokenizer # Load model and tokenizer
model = AutoModel.from_pretrained(model_dir) model = AutoModel.from_pretrained(param_path)
tokenizer = AutoTokenizer.from_pretrained(model_dir) tokenizer = AutoTokenizer.from_pretrained(param_path)
model.to(device="cuda", dtype=torch.bfloat16) model.to(device="cuda", dtype=torch.bfloat16)
# Create inference engine # Create inference engine
@@ -72,7 +72,7 @@ if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Run generate with a Khaosz model.") parser = argparse.ArgumentParser(description="Run generate with a Khaosz model.")
parser.add_argument( parser.add_argument(
"--model_dir", type=str, required=True, help="Path to the model directory." "--param_path", type=str, required=True, help="Path to the model directory."
) )
parser.add_argument( parser.add_argument(
"--input_json_file", "--input_json_file",
+32 -10
View File
@@ -23,7 +23,7 @@ def parse_args() -> argparse.Namespace:
"--train_type", "--train_type",
type=str, type=str,
required=True, required=True,
choices=["seq", "sft", "dpo"], choices=["seq", "sft", "dpo", "grpo"],
help="Train type.", help="Train type.",
) )
parser.add_argument( parser.add_argument(
@@ -42,9 +42,7 @@ def parse_args() -> argparse.Namespace:
parser.add_argument( parser.add_argument(
"--n_epoch", type=int, default=1, help="Number of epochs to train." "--n_epoch", type=int, default=1, help="Number of epochs to train."
) )
parser.add_argument( parser.add_argument("--batch_size", type=int, default=1, help="Batch size per GPU.")
"--batch_size", type=int, default=1, help="Batch size for training."
)
parser.add_argument( parser.add_argument(
"--accumulation_steps", "--accumulation_steps",
type=int, type=int,
@@ -55,7 +53,7 @@ def parse_args() -> argparse.Namespace:
"--warmup_steps", "--warmup_steps",
type=int, type=int,
default=1000, default=1000,
help="Number of iters between warnings.", help="Number of warmup steps for LR scheduler.",
) )
parser.add_argument( parser.add_argument(
"--max_lr", type=float, default=3e-4, help="Max learning rate for training." "--max_lr", type=float, default=3e-4, help="Max learning rate for training."
@@ -100,12 +98,19 @@ def parse_args() -> argparse.Namespace:
"--window_size", "--window_size",
type=int, type=int,
default=None, default=None,
help="the max length of the input sequence.", help="Max length of the input sequence.",
) )
parser.add_argument( parser.add_argument(
"--stride", type=int, default=None, help="the step size of the input sequence." "--stride", type=int, default=None, help="Step size of the input sequence."
) )
parser.add_argument("--dpo_beta", type=float, default=0.1, help="DPO beta value.") parser.add_argument("--dpo_beta", type=float, default=0.1, help="DPO beta value.")
parser.add_argument("--group_size", type=int, default=4, help="GRPO group size.")
parser.add_argument(
"--grpo_clip_eps", type=float, default=0.2, help="GRPO clipping epsilon."
)
parser.add_argument(
"--grpo_kl_coef", type=float, default=0.01, help="GRPO KL penalty coefficient."
)
parser.add_argument( parser.add_argument(
"--label_smoothing", "--label_smoothing",
type=float, type=float,
@@ -125,6 +130,12 @@ def parse_args() -> argparse.Namespace:
default="checkpoint", default="checkpoint",
help="Directory to save checkpoints.", help="Directory to save checkpoints.",
) )
parser.add_argument(
"--grpo_sync_interval",
type=int,
default=200,
help="GRPO ref model sync interval (steps).",
)
parser.add_argument( parser.add_argument(
"--start_epoch", type=int, default=0, help="Start epoch for training." "--start_epoch", type=int, default=0, help="Start epoch for training."
) )
@@ -144,7 +155,7 @@ def parse_args() -> argparse.Namespace:
def ddp_wrap(model: nn.Module): def ddp_wrap(model: nn.Module):
local_rank = get_rank() local_rank = get_rank()
model = model.to(device=f"cuda:{local_rank}", dtype=torch.bfloat16) model = model.to(dtype=torch.bfloat16)
ddp_model = DDP( ddp_model = DDP(
model, model,
device_ids=[local_rank], device_ids=[local_rank],
@@ -182,6 +193,10 @@ def train(
ckpt_interval: int, ckpt_interval: int,
ckpt_dir: str, ckpt_dir: str,
dpo_beta: float, dpo_beta: float,
grpo_clip_eps: float,
grpo_kl_coef: float,
group_size: int,
grpo_sync_interval: int,
adamw_beta1: float, adamw_beta1: float,
adamw_beta2: float, adamw_beta2: float,
adamw_weight_decay: float, adamw_weight_decay: float,
@@ -195,7 +210,7 @@ def train(
nprocs: int, nprocs: int,
device_type: str, device_type: str,
): ):
assert train_type in ["seq", "sft", "dpo"] assert train_type in ["seq", "sft", "dpo", "grpo"]
assert os.path.exists(param_path) assert os.path.exists(param_path)
# Load config # Load config
@@ -216,7 +231,14 @@ def train(
state_dict = st.load_file(weights_path) state_dict = st.load_file(weights_path)
model.load_state_dict(state_dict, strict=False) model.load_state_dict(state_dict, strict=False)
strategy_kwargs = {"dpo_beta": dpo_beta, "label_smoothing": label_smoothing} strategy_kwargs = {
"dpo_beta": dpo_beta,
"label_smoothing": label_smoothing,
"clip_eps": grpo_clip_eps,
"kl_coef": grpo_kl_coef,
"group_size": group_size,
"sync_interval": grpo_sync_interval,
}
dataset = DatasetFactory.load( dataset = DatasetFactory.load(
train_type=train_type, train_type=train_type,
+14 -19
View File
@@ -14,37 +14,32 @@ def client():
return TestClient(app) return TestClient(app)
@pytest.fixture
def mock_model_param():
"""Create a mock ModelParameter."""
mock_param = MagicMock()
mock_param.model = MagicMock()
mock_param.tokenizer = MagicMock()
mock_param.config = MagicMock()
mock_param.config.max_len = 100
mock_param.tokenizer.encode = MagicMock(return_value=[1, 2, 3])
mock_param.tokenizer.decode = MagicMock(return_value="mock response")
mock_param.tokenizer.stop_ids = []
mock_param.tokenizer.pad_id = 0
return mock_param
@pytest.fixture @pytest.fixture
def mock_engine(): def mock_engine():
"""Create a mock InferenceEngine.""" """Create a mock InferenceEngine."""
async def _async_gen():
yield "chunk1"
yield "chunk2"
yield "[DONE]"
mock = MagicMock() mock = MagicMock()
mock.generate.return_value = "mock response" mock.generate.return_value = "mock response"
mock.generate_async.return_value = _async_gen()
mock.get_stats.return_value = { mock.get_stats.return_value = {
"total_tasks": 0, "total_tasks": 0,
"total_tokens": 0, "total_tokens": 0,
"active_tasks": 0, "active_tasks": 0,
"waiting_queue": 0, "waiting_queue": 0,
} }
mock.tokenizer.encode.return_value = [1, 2, 3]
mock.tokenizer.decode.return_value = "mock response"
mock.tokenizer.apply_chat_template.return_value = "mock prompt"
return mock return mock
@pytest.fixture @pytest.fixture
def loaded_model(mock_model_param, monkeypatch): def loaded_model(mock_engine, monkeypatch):
"""Simulate that the model is loaded.""" """Simulate that the engine is loaded."""
monkeypatch.setattr("astrai.inference.server._model_param", mock_model_param) monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
return mock_model_param return mock_engine
+5 -148
View File
@@ -5,103 +5,9 @@ import time
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
import pytest import pytest
import torch
from astrai.inference.scheduler import ( from astrai.inference.scheduler import InferenceScheduler
InferenceScheduler,
PrefixCacheManager,
)
def test_prefix_cache_concurrent_insert_find():
"""Test concurrent insert and find operations."""
cache = PrefixCacheManager(max_capacity=100)
results = {"errors": [], "inserts": 0, "finds": 0}
def insert_worker():
try:
for i in range(50):
cache.insert((i,), slot=i % 10)
results["inserts"] += 1
except Exception as e:
results["errors"].append(str(e))
def find_worker():
try:
for i in range(50):
cache.find_longest_prefix([i])
results["finds"] += 1
except Exception as e:
results["errors"].append(str(e))
threads = [threading.Thread(target=insert_worker) for _ in range(3)]
threads += [threading.Thread(target=find_worker) for _ in range(3)]
for t in threads:
t.start()
for t in threads:
t.join()
assert len(results["errors"]) == 0, f"Errors: {results['errors']}"
assert results["inserts"] == 150
assert results["finds"] == 150
def test_prefix_cache_concurrent_release():
"""Test concurrent release operations."""
cache = PrefixCacheManager(max_capacity=100)
# Insert some prefixes
for i in range(10):
cache.insert((i,), slot=i)
results = {"errors": []}
def release_worker():
try:
for i in range(10):
cache.release((i,))
except Exception as e:
results["errors"].append(str(e))
threads = [threading.Thread(target=release_worker) for _ in range(3)]
for t in threads:
t.start()
for t in threads:
t.join()
assert len(results["errors"]) == 0, f"Errors: {results['errors']}"
def test_prefix_cache_concurrent_insert_release_find():
"""Test mixed concurrent operations."""
cache = PrefixCacheManager(max_capacity=50)
results = {"errors": []}
def worker(worker_id):
try:
for i in range(20):
token_ids = (worker_id * 100 + i,)
cache.insert(token_ids, slot=worker_id)
# Find after insert
cache.find_longest_prefix(list(token_ids))
# Release
cache.release(token_ids)
except Exception as e:
results["errors"].append(f"Worker {worker_id}: {str(e)}")
threads = [threading.Thread(target=worker, args=(i,)) for i in range(5)]
for t in threads:
t.start()
for t in threads:
t.join()
assert len(results["errors"]) == 0, f"Errors: {results['errors']}"
@pytest.fixture @pytest.fixture
@@ -114,6 +20,9 @@ def mock_model_and_tokenizer():
mock_model.config.dim = 128 mock_model.config.dim = 128
mock_model.config.n_layers = 2 mock_model.config.n_layers = 2
mock_model.config.max_len = 100 mock_model.config.max_len = 100
mock_model.parameters.return_value = iter(
[MagicMock(dtype=torch.float32, device=torch.device("cpu"))]
)
mock_tokenizer = MagicMock() mock_tokenizer = MagicMock()
mock_tokenizer.encode.return_value = [1, 2, 3, 4, 5] mock_tokenizer.encode.return_value = [1, 2, 3, 4, 5]
@@ -266,55 +175,3 @@ def test_scheduler_concurrent_get_stats(mock_model_and_tokenizer):
for stats in results["stats"]: for stats in results["stats"]:
assert "total_tasks" in stats assert "total_tasks" in stats
assert stats["total_tasks"] >= 0 assert stats["total_tasks"] >= 0
def test_prefix_cache_insert_same_prefix_concurrently():
"""Test inserting the same prefix concurrently."""
cache = PrefixCacheManager(max_capacity=100)
results = {"slot_values": [], "errors": []}
def insert_worker():
try:
# All workers try to insert the same prefix
cache.insert((1, 2, 3), slot=threading.current_thread().name)
node = cache.root.children.get(1)
if node:
node = node.children.get(2)
if node:
node = node.children.get(3)
if node:
results["slot_values"].append(node.slot)
except Exception as e:
results["errors"].append(str(e))
threads = [threading.Thread(target=insert_worker) for _ in range(10)]
for t in threads:
t.start()
for t in threads:
t.join()
# All inserts should succeed, final slot should be one of the values
assert len(results["errors"]) == 0, f"Errors: {results['errors']}"
# Check ref_count is correct (should be 10)
node = cache.root.children.get(1).children.get(2).children.get(3)
assert node.ref_count == 10, f"Expected ref_count=10, got {node.ref_count}"
def test_prefix_cache_ref_count_underflow_prevention():
"""Test that ref_count doesn't go negative."""
cache = PrefixCacheManager(max_capacity=100)
# Insert a prefix
cache.insert((1, 2, 3), slot=0)
# Release multiple times
for _ in range(5):
cache.release((1, 2, 3))
# Try to find it - should return None since ref_count would be negative
# or handle it gracefully
node = cache.root.children.get(1).children.get(2).children.get(3)
# The ref_count should be 0, not negative
assert node.ref_count >= 0, f"ref_count went negative: {node.ref_count}"
+95 -82
View File
@@ -4,88 +4,38 @@ import pytest
def test_health_no_model(client, monkeypatch): def test_health_no_model(client, monkeypatch):
"""GET /health should return 200 even when model not loaded.""" """GET /health should return 200 even when engine not loaded."""
monkeypatch.setattr("astrai.inference.server._model_param", None) monkeypatch.setattr("astrai.inference.server._state.engine", None)
monkeypatch.setattr("astrai.inference.server._engine", None)
response = client.get("/health") response = client.get("/health")
assert response.status_code == 200 assert response.status_code == 200
data = response.json() data = response.json()
assert data["status"] == "ok" assert data["status"] == "ok"
assert not data["model_loaded"] assert not data["model_loaded"]
assert not data["engine_ready"]
def test_health_with_model(client, loaded_model, mock_engine, monkeypatch): def test_health_with_model(client, loaded_model):
"""GET /health should return 200 when model is loaded.""" """GET /health should return 200 when engine is loaded."""
monkeypatch.setattr("astrai.inference.server._engine", mock_engine)
response = client.get("/health") response = client.get("/health")
assert response.status_code == 200 assert response.status_code == 200
data = response.json() data = response.json()
assert data["status"] == "ok" assert data["status"] == "ok"
assert data["model_loaded"] is True assert data["model_loaded"] is True
assert data["engine_ready"] is True
def test_generate_non_stream(client, loaded_model, mock_engine, monkeypatch): def test_chat_completions_non_stream(client, loaded_model, monkeypatch):
"""POST /generate with stream=false should return JSON response.""" """POST /v1/chat/completions with stream=false returns OpenAI-style JSON."""
monkeypatch.setattr("astrai.inference.server._engine", mock_engine)
response = client.post(
"/generate",
params={
"query": "Hello",
"temperature": 0.8,
"top_p": 0.95,
"top_k": 50,
"max_len": 100,
"stream": False,
},
)
assert response.status_code == 200
data = response.json()
assert data["response"] == "mock response"
async def async_gen():
yield "Assistant reply"
def test_generate_stream(client, loaded_model, mock_engine, monkeypatch): mock_engine = loaded_model
"""POST /generate with stream=true should return plain text stream.""" mock_engine.generate_async.return_value = async_gen()
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
# Create a streaming mock
def stream_gen():
yield "chunk1"
yield "chunk2"
mock_engine.generate.return_value = stream_gen()
monkeypatch.setattr("astrai.inference.server._engine", mock_engine)
response = client.post(
"/generate",
params={
"query": "Hello",
"temperature": 0.8,
"top_p": 0.95,
"top_k": 50,
"max_len": 100,
"stream": True,
},
headers={"Accept": "text/plain"},
)
assert response.status_code == 200
assert response.headers["content-type"] == "text/plain; charset=utf-8"
# The stream yields lines ending with newline
content = response.content.decode("utf-8")
assert "chunk1" in content
assert "chunk2" in content
def test_chat_completions_non_stream(client, loaded_model, mock_engine, monkeypatch):
"""POST /v1/chat/completions with stream=false returns OpenAIstyle JSON."""
mock_engine.generate.return_value = "Assistant reply"
monkeypatch.setattr("astrai.inference.server._engine", mock_engine)
response = client.post( response = client.post(
"/v1/chat/completions", "/v1/chat/completions",
json={ json={
"messages": [{"role": "user", "content": "Hello"}], "messages": [{"role": "user", "content": "Hello"}],
"temperature": 0.8, "temperature": 0.8,
"top_p": 0.95,
"top_k": 50,
"max_tokens": 100, "max_tokens": 100,
"stream": False, "stream": False,
}, },
@@ -94,57 +44,120 @@ def test_chat_completions_non_stream(client, loaded_model, mock_engine, monkeypa
data = response.json() data = response.json()
assert data["object"] == "chat.completion" assert data["object"] == "chat.completion"
assert len(data["choices"]) == 1 assert len(data["choices"]) == 1
assert data["choices"][0]["message"]["content"] == "Assistant reply" assert "usage" in data
assert "prompt_tokens" in data["usage"]
def test_chat_completions_stream(client, loaded_model, mock_engine, monkeypatch): def test_chat_completions_stream(client, loaded_model, monkeypatch):
"""POST /v1/chat/completions with stream=true returns SSE stream.""" """POST /v1/chat/completions with stream=true returns SSE stream."""
# Simulate a streaming generator that yields cumulative responses async def async_gen():
def stream_gen():
yield "cumulative1" yield "cumulative1"
yield "cumulative2" yield "cumulative2"
yield "[DONE]"
mock_engine.generate.return_value = stream_gen() mock_engine = loaded_model
monkeypatch.setattr("astrai.inference.server._engine", mock_engine) mock_engine.generate_async.return_value = async_gen()
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
response = client.post( response = client.post(
"/v1/chat/completions", "/v1/chat/completions",
json={ json={
"messages": [{"role": "user", "content": "Hello"}], "messages": [{"role": "user", "content": "Hello"}],
"temperature": 0.8, "temperature": 0.8,
"top_p": 0.95,
"top_k": 50,
"max_tokens": 100, "max_tokens": 100,
"stream": True, "stream": True,
}, },
headers={"Accept": "text/event-stream"}, headers={"Accept": "text/event-stream"},
) )
assert response.status_code == 200 assert response.status_code == 200
assert response.headers["content-type"] == "text/event-stream; charset=utf-8"
# Parse SSE lines
lines = [ lines = [
line.strip() for line in response.content.decode("utf-8").split("\n") if line line.strip() for line in response.content.decode("utf-8").split("\n") if line
] ]
# Should contain data lines and a final [DONE]
assert any("cumulative1" in line for line in lines) assert any("cumulative1" in line for line in lines)
assert any("cumulative2" in line for line in lines) assert any("cumulative2" in line for line in lines)
assert any("[DONE]" in line for line in lines)
def test_generate_with_history(client, loaded_model, mock_engine, monkeypatch): def test_messages_non_stream(client, loaded_model, monkeypatch):
"""POST /generate with history parameter.""" """POST /v1/messages with stream=false returns Anthropic-style JSON."""
monkeypatch.setattr("astrai.inference.server._engine", mock_engine)
async def async_gen():
yield "Assistant reply"
mock_engine = loaded_model
mock_engine.generate_async.return_value = async_gen()
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
response = client.post( response = client.post(
"/generate", "/v1/messages",
params={ json={
"query": "Hi", "messages": [{"role": "user", "content": "Hello"}],
"history": [["user1", "assistant1"], ["user2", "assistant2"]], "temperature": 0.8,
"max_tokens": 100,
"stream": False, "stream": False,
}, },
) )
assert response.status_code == 200 assert response.status_code == 200
# Verify the engine.generate was called data = response.json()
mock_engine.generate.assert_called_once() assert data["type"] == "message"
assert data["role"] == "assistant"
assert len(data["content"]) == 1
assert data["content"][0]["type"] == "text"
assert "usage" in data
assert "input_tokens" in data["usage"]
def test_messages_stream(client, loaded_model, monkeypatch):
"""POST /v1/messages with stream=true returns Anthropic SSE stream."""
async def async_gen():
yield "cumulative1"
yield "cumulative2"
mock_engine = loaded_model
mock_engine.generate_async.return_value = async_gen()
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
response = client.post(
"/v1/messages",
json={
"messages": [{"role": "user", "content": "Hello"}],
"temperature": 0.8,
"max_tokens": 100,
"stream": True,
},
headers={"Accept": "text/event-stream"},
)
assert response.status_code == 200
content = response.content.decode("utf-8")
assert "message_start" in content
assert "content_block_start" in content
assert "content_block_delta" in content
assert "cumulative1" in content
assert "cumulative2" in content
assert "content_block_stop" in content
assert "message_delta" in content
assert "message_stop" in content
def test_messages_with_system(client, loaded_model, monkeypatch):
"""POST /v1/messages with system prompt."""
async def async_gen():
yield "Reply"
mock_engine = loaded_model
mock_engine.generate_async.return_value = async_gen()
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
response = client.post(
"/v1/messages",
json={
"messages": [{"role": "user", "content": "Hello"}],
"system": "You are a helpful assistant.",
"max_tokens": 100,
"stream": False,
},
)
assert response.status_code == 200
data = response.json()
assert data["type"] == "message"
if __name__ == "__main__": if __name__ == "__main__":
+3
View File
@@ -72,6 +72,7 @@ def test_schedule_factory_random_configs():
# Test scheduler step functionality # Test scheduler step functionality
initial_lr = scheduler.get_last_lr() initial_lr = scheduler.get_last_lr()
optimizer.step()
scheduler.step() scheduler.step()
new_lr = scheduler.get_last_lr() new_lr = scheduler.get_last_lr()
@@ -112,6 +113,7 @@ def test_schedule_factory_edge_cases():
# Test multiple steps # Test multiple steps
for _ in range(10): for _ in range(10):
optimizer.step()
scheduler.step() scheduler.step()
@@ -136,6 +138,7 @@ def test_schedule_factory_state_persistence():
# Take a few steps # Take a few steps
for _ in range(5): for _ in range(5):
optimizer.step()
scheduler.step() scheduler.step()
# Save state # Save state