Compare commits
4
Commits
v1.3.5
...
3d12a03909
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3d12a03909 | ||
|
|
c169659611 | ||
|
|
e12f1a7ee5 | ||
|
|
ef25efffa2 |
@@ -2,7 +2,7 @@
|
|||||||
name: Bug report
|
name: Bug report
|
||||||
about: Create a report to help us improve
|
about: Create a report to help us improve
|
||||||
title: "[BUG]"
|
title: "[BUG]"
|
||||||
labels: enhancement
|
labels: bug
|
||||||
assignees: ''
|
assignees: ''
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|||||||
@@ -16,9 +16,9 @@ Please delete options that are not relevant.
|
|||||||
Please describe the tests that you ran to verify your changes. Provide instructions so we can reproduce.
|
Please describe the tests that you ran to verify your changes. Provide instructions so we can reproduce.
|
||||||
|
|
||||||
## Checklist:
|
## Checklist:
|
||||||
- [ ] My code follows the style guidelines of this project (run `ruff format .` and `ruff check --fix .`)
|
- [ ] My code follows the style guidelines of this project (run `ruff format .` and `ruff check . --select I`)
|
||||||
- [ ] I have performed a self-review of my own code
|
- [ ] I have performed a self-review of my own code
|
||||||
- [ ] I have commented my code, particularly in hard-to-understand areas
|
- [ ] Code is self-documenting (no unnecessary comments)
|
||||||
- [ ] I have made corresponding changes to the documentation
|
- [ ] I have made corresponding changes to the documentation
|
||||||
- [ ] My changes generate no new warnings
|
- [ ] My changes generate no new warnings
|
||||||
- [ ] I have added tests that prove my fix is effective or that my feature works
|
- [ ] I have added tests that prove my fix is effective or that my feature works
|
||||||
|
|||||||
+68
-36
@@ -1,68 +1,100 @@
|
|||||||
# Contributing to AstrAI
|
# Contributing to AstrAI
|
||||||
|
|
||||||
Thank you for your interest in contributing to AstrAI! This document provides guidelines and steps for contributing.
|
Thank you for your interest in contributing! This document provides step-by-step guidelines.
|
||||||
|
|
||||||
## How to Contribute
|
## Quick Start
|
||||||
|
|
||||||
### Reporting Issues
|
|
||||||
If you encounter a bug or have a feature request, please open an issue on GitHub. Include as much detail as possible:
|
|
||||||
- A clear description of the problem or request.
|
|
||||||
- Steps to reproduce (for bugs).
|
|
||||||
- Your environment (Python version, OS, etc.).
|
|
||||||
|
|
||||||
### Submitting Changes
|
|
||||||
1. **Fork** the repository.
|
|
||||||
2. **Clone** your fork:
|
|
||||||
```bash
|
```bash
|
||||||
git clone https://github.com/your-username/AstrAI.git
|
git clone https://github.com/your-username/AstrAI.git
|
||||||
cd AstrAI
|
cd AstrAI
|
||||||
|
pip install -e ".[dev]" # install with dev dependencies (pytest, ruff)
|
||||||
```
|
```
|
||||||
3. **Create a feature branch**:
|
|
||||||
|
## Before You Commit
|
||||||
|
|
||||||
|
Run the following checks **in order** — CI will reject if any fail.
|
||||||
|
|
||||||
|
### 1. Format
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
git checkout -b feature/your-feature-name
|
ruff format .
|
||||||
```
|
```
|
||||||
4. **Make your changes**. Follow the code style guidelines below.
|
|
||||||
5. **Commit your changes** with a descriptive commit message:
|
> **Note**: `ruff format` may rename parameters (e.g. `mask` → `attn_mask`).
|
||||||
|
> Always review the diff after formatting.
|
||||||
|
|
||||||
|
### 2. Import sorting
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
git commit -m "Add: brief description of the change"
|
ruff check . --select I
|
||||||
```
|
```
|
||||||
6. **Push** to your fork:
|
|
||||||
|
If this fails, **manually fix** import ordering (ruff does not auto-fix in this project's CI):
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
git push origin feature/your-feature-name
|
ruff check . --select I --fix .
|
||||||
|
ruff format . # re-format after fix
|
||||||
```
|
```
|
||||||
7. **Open a Pull Request** (PR) against the `main` branch of the upstream repository.
|
|
||||||
|
|
||||||
## Code Style
|
### 3. Run tests
|
||||||
|
|
||||||
AstrAI uses [Ruff](https://docs.astral.sh/ruff/) for code formatting and linting. Please ensure your code is formatted before submitting.
|
|
||||||
|
|
||||||
- Run Ruff to format and lint (requires conda environment `nlp`):
|
|
||||||
```bash
|
```bash
|
||||||
conda run -n nlp ruff format .
|
python -u -m pytest tests/ -v
|
||||||
conda run -n nlp ruff check --fix .
|
|
||||||
```
|
```
|
||||||
- The project uses **double quotes** for strings and **4‑space indentation** (as configured in `pyproject.toml`).
|
|
||||||
|
|
||||||
## Testing
|
> Failed tests may leave orphan tempdirs under `%TEMP%`. Clean them manually if needed.
|
||||||
|
|
||||||
If you add or modify functionality, please include appropriate tests.
|
### 4. (Optional) Full pre-commit check
|
||||||
|
|
||||||
|
If you have Git Bash available:
|
||||||
|
|
||||||
- Run the test suite with:
|
|
||||||
```bash
|
```bash
|
||||||
conda run -n nlp python -u -m pytest
|
bash scripts/pre_commit.sh
|
||||||
```
|
```
|
||||||
- Ensure all tests pass before submitting your PR.
|
|
||||||
|
This runs format check, import sort check, and tests in one go.
|
||||||
|
|
||||||
|
## Commit Style
|
||||||
|
|
||||||
|
```
|
||||||
|
fix/feat/chore/docs/refactor/perf/test/style/ci/build/revert : short description (~50 chars)
|
||||||
|
|
||||||
|
- bullet point body (each ~60 chars)
|
||||||
|
```
|
||||||
|
|
||||||
|
- **Type** must be one of: `fix`, `feat`, `chore`, `docs`, `refactor`, `perf`, `test`, `style`, `ci`, `build`, `revert`.
|
||||||
|
- **Subject line** ends with no period.
|
||||||
|
- **Body** uses bullet points starting with `-`.
|
||||||
|
- No `(scope)` parentheses.
|
||||||
|
|
||||||
|
## Common Issues
|
||||||
|
|
||||||
|
| Problem | Cause | Fix |
|
||||||
|
|---------|-------|-----|
|
||||||
|
| `ruff check --select I` fails | Wrong import order | `ruff check . --select I --fix .` then `ruff format .` |
|
||||||
|
| `ruff format` changed many files | Not formatted before commit | Review diff carefully before staging |
|
||||||
|
| Pre-commit hook rejects | Tests or lint failed | Fix individually, do not `--no-verify` |
|
||||||
|
| Tests fail with tempdir left | Test crash | Clean `%TEMP%` manually |
|
||||||
|
|
||||||
|
## Submitting Changes
|
||||||
|
|
||||||
|
1. Fork the repo.
|
||||||
|
2. Create a feature branch: `git checkout -b feat/my-feature`
|
||||||
|
3. Make changes following the steps above.
|
||||||
|
4. Commit with the commit style above.
|
||||||
|
5. Push: `git push origin feat/my-feature`
|
||||||
|
6. Open a Pull Request against `main`.
|
||||||
|
|
||||||
## Code Review
|
## Code Review
|
||||||
|
|
||||||
All submissions will be reviewed. We may request changes or discuss alternatives. Please be responsive to feedback.
|
- All PRs are reviewed. We may request changes.
|
||||||
|
- CI runs `ruff format --check .` then `ruff check . --select I` (no `--fix` in CI).
|
||||||
|
- Ensure all tests pass.
|
||||||
|
|
||||||
## License
|
## License
|
||||||
|
|
||||||
By contributing, you agree that your contributions will be licensed under the same [GPL-3.0 License](LICENSE) that covers the project.
|
By contributing, you agree that your contributions will be licensed under the [GPL-3.0 License](LICENSE).
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
If you have any questions, feel free to ask in the [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions) or open an issue.
|
Questions? Ask in [GitHub Discussions](https://github.com/ViperEkura/AstrAI/discussions) or open an issue.
|
||||||
|
|
||||||
Happy contributing!
|
|
||||||
|
|||||||
@@ -208,9 +208,10 @@ Watch a video walkthrough on [bilibili](https://www.bilibili.com/video/BV1z5RPYH
|
|||||||
| Document | Description |
|
| Document | Description |
|
||||||
|----------|-------------|
|
|----------|-------------|
|
||||||
| [Parameter Guide](./assets/docs/params.md) | Training & inference parameters |
|
| [Parameter Guide](./assets/docs/params.md) | Training & inference parameters |
|
||||||
| [Design Document](./assets/docs/design.md) | Framework architecture & module design |
|
| [Architecture](./assets/docs/architecture.md) | System architecture, class diagram & design patterns |
|
||||||
| [Data Flow](./assets/docs/dataflow.md) | Data processing pipeline details |
|
| [Training](./assets/docs/training.md) | Training loop, strategies & formulas |
|
||||||
| [Model Introduction](./assets/docs/introduction.md) | Model architecture & technical details |
|
| [Inference](./assets/docs/inference.md) | KVCache, continuous batching, sampling & HTTP API |
|
||||||
|
| [Data Flow](./assets/docs/dataflow.md) | Data pipeline, storage backends & dataset architecture |
|
||||||
|
|
||||||
### Contributing
|
### Contributing
|
||||||
|
|
||||||
|
|||||||
@@ -214,9 +214,10 @@ python scripts/demo/generate_ar.py
|
|||||||
| 文档 | 说明 |
|
| 文档 | 说明 |
|
||||||
|------|------|
|
|------|------|
|
||||||
| [参数说明](./params.md) | 训练与推理参数配置 |
|
| [参数说明](./params.md) | 训练与推理参数配置 |
|
||||||
| [设计文档](./design.md) | 系统架构与模块设计 |
|
| [架构文档](./architecture.md) | 系统架构、类图与设计模式 |
|
||||||
| [数据流程](./dataflow.md) | 数据处理管道详解 |
|
| [训练文档](./training.md) | 训练循环、策略与公式 |
|
||||||
| [模型介绍](./introduction.md) | 模型架构与技术细节 |
|
| [推理文档](./inference.md) | KVCache、连续批处理、采样与 HTTP API |
|
||||||
|
| [数据流程](./dataflow.md) | 数据管道、存储后端与数据集架构 |
|
||||||
|
|
||||||
### 贡献
|
### 贡献
|
||||||
|
|
||||||
|
|||||||
@@ -1,14 +1,16 @@
|
|||||||
## 1. Why I Created This Project
|
# AstrAI Architecture
|
||||||
|
|
||||||
There are many large language models on the market today, such as GPT, LLaMA, and others, with tens of billions or even hundreds of billions of parameters. But honestly, these models have extremely high hardware requirements, making them inaccessible for ordinary developers. I thought: **Can we create a model that is both useful and can run on ordinary computers?** This is also what most people currently hope for - a locally deployable AI project that achieves complete privatization while maintaining some level of intelligence.
|
## Class Diagram
|
||||||
|
|
||||||
Thus, the AstrAI project was born - 1B parameters, Chinese-English bilingual, supporting dialogue, text generation, and the training code is open source!
|
|
||||||
|
|
||||||
## 2. System Architecture
|
|
||||||
|
|
||||||
```mermaid
|
```mermaid
|
||||||
classDiagram
|
classDiagram
|
||||||
namespace config {
|
namespace config {
|
||||||
|
class BaseModelConfig {
|
||||||
|
+Optional[str] model_type
|
||||||
|
+load(config_path) Self
|
||||||
|
+save(config_path)
|
||||||
|
}
|
||||||
|
|
||||||
class ModelConfig {
|
class ModelConfig {
|
||||||
+int vocab_size
|
+int vocab_size
|
||||||
+int dim
|
+int dim
|
||||||
@@ -22,6 +24,12 @@ classDiagram
|
|||||||
+int n_kv_heads
|
+int n_kv_heads
|
||||||
+bool use_qk_norm
|
+bool use_qk_norm
|
||||||
+bool use_gated_attention
|
+bool use_gated_attention
|
||||||
|
+str attn_type
|
||||||
|
+str ffn_type
|
||||||
|
+int n_routed_experts
|
||||||
|
+int n_shared_experts
|
||||||
|
+int n_activated_experts
|
||||||
|
+str moe_topk_method
|
||||||
+load(config_path) ModelConfig
|
+load(config_path) ModelConfig
|
||||||
+save(config_path)
|
+save(config_path)
|
||||||
}
|
}
|
||||||
@@ -42,7 +50,7 @@ classDiagram
|
|||||||
+int ckpt_interval
|
+int ckpt_interval
|
||||||
+int random_seed
|
+int random_seed
|
||||||
+int num_workers
|
+int num_workers
|
||||||
+int prefetch_factor
|
+Optional[int] prefetch_factor
|
||||||
+bool pin_memory
|
+bool pin_memory
|
||||||
+int nprocs
|
+int nprocs
|
||||||
+str backend
|
+str backend
|
||||||
@@ -118,8 +126,8 @@ classDiagram
|
|||||||
}
|
}
|
||||||
|
|
||||||
class ResumableDistributedSampler {
|
class ResumableDistributedSampler {
|
||||||
+int epoch
|
+int start_epoch
|
||||||
+int iter
|
+int start_iter
|
||||||
}
|
}
|
||||||
|
|
||||||
class DatasetFactory {
|
class DatasetFactory {
|
||||||
@@ -135,6 +143,7 @@ classDiagram
|
|||||||
+dict state_dict
|
+dict state_dict
|
||||||
+int epoch
|
+int epoch
|
||||||
+int iteration
|
+int iteration
|
||||||
|
+dict extra
|
||||||
+save(save_dir)
|
+save(save_dir)
|
||||||
+load(save_dir) Checkpoint
|
+load(save_dir) Checkpoint
|
||||||
}
|
}
|
||||||
@@ -158,15 +167,15 @@ classDiagram
|
|||||||
+ModuleList layers
|
+ModuleList layers
|
||||||
+RMSNorm norm
|
+RMSNorm norm
|
||||||
+Linear lm_head
|
+Linear lm_head
|
||||||
+forward(input_ids, input_mask, paged_cache, position_ids) Tensor
|
+forward(input_ids, input_mask, paged_cache, position_ids) Dict
|
||||||
+load_state_dict(state_dict)
|
+load_state_dict(state_dict)
|
||||||
+state_dict()
|
+state_dict()
|
||||||
}
|
}
|
||||||
|
|
||||||
class DecoderBlock {
|
class DecoderBlock {
|
||||||
+GQA attention
|
+nn.Module attention # GQA or MLA via AttnFactory
|
||||||
+RMSNorm input_norm
|
+RMSNorm input_norm
|
||||||
+MLP mlp
|
+nn.Module mlp # MLP or DeepSeekMoE via FFNFactory
|
||||||
+RMSNorm post_attention_norm
|
+RMSNorm post_attention_norm
|
||||||
+forward(x, rotary_emb, attention_mask, paged_cache) Tensor
|
+forward(x, rotary_emb, attention_mask, paged_cache) Tensor
|
||||||
}
|
}
|
||||||
@@ -175,8 +184,12 @@ classDiagram
|
|||||||
+int n_heads
|
+int n_heads
|
||||||
+int n_kv_heads
|
+int n_kv_heads
|
||||||
+int head_dim
|
+int head_dim
|
||||||
|
+int n_rep
|
||||||
|
+bool use_qk_norm
|
||||||
|
+bool use_gated_attention
|
||||||
+Linear q_proj, k_proj, v_proj, o_proj
|
+Linear q_proj, k_proj, v_proj, o_proj
|
||||||
+RMSNorm q_norm, k_norm
|
+Linear gate # only if use_gated_attention
|
||||||
|
+RMSNorm q_norm, k_norm # only if use_qk_norm
|
||||||
+forward(x, rotary_emb, attn_mask, paged_cache) Tensor
|
+forward(x, rotary_emb, attn_mask, paged_cache) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -187,8 +200,11 @@ classDiagram
|
|||||||
+int kv_lora_rank
|
+int kv_lora_rank
|
||||||
+int qk_nope_head_dim
|
+int qk_nope_head_dim
|
||||||
+int qk_rope_head_dim
|
+int qk_rope_head_dim
|
||||||
|
+int n_rep
|
||||||
|
+bool use_gated_attention
|
||||||
+Linear q_proj, kv_a_proj, kv_b_proj
|
+Linear q_proj, kv_a_proj, kv_b_proj
|
||||||
+Linear o_proj
|
+Linear o_proj
|
||||||
|
+Linear gate # only if use_gated_attention
|
||||||
+RMSNorm kv_norm
|
+RMSNorm kv_norm
|
||||||
+forward(x, rotary_emb, attn_mask, paged_cache) Tensor
|
+forward(x, rotary_emb, attn_mask, paged_cache) Tensor
|
||||||
}
|
}
|
||||||
@@ -198,6 +214,25 @@ classDiagram
|
|||||||
+forward(x) Tensor
|
+forward(x) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
|
class DeepSeekMoE {
|
||||||
|
+int n_routed_experts
|
||||||
|
+int n_shared_experts
|
||||||
|
+int n_activated_experts
|
||||||
|
+str topk_method
|
||||||
|
+Linear router
|
||||||
|
+ModuleList shared_experts
|
||||||
|
+ModuleList routed_experts
|
||||||
|
+forward(x) Tensor
|
||||||
|
}
|
||||||
|
|
||||||
|
class AttnFactory {
|
||||||
|
+create(attn_type, **kwargs) nn.Module
|
||||||
|
}
|
||||||
|
|
||||||
|
class FFNFactory {
|
||||||
|
+create(ffn_type, dim, dim_ffn, **kwargs) nn.Module
|
||||||
|
}
|
||||||
|
|
||||||
class RMSNorm {
|
class RMSNorm {
|
||||||
+Parameter weight
|
+Parameter weight
|
||||||
+float norm_eps
|
+float norm_eps
|
||||||
@@ -206,7 +241,7 @@ classDiagram
|
|||||||
|
|
||||||
class Linear {
|
class Linear {
|
||||||
+Parameter weight
|
+Parameter weight
|
||||||
+Parameter bias
|
+Optional[Parameter] bias # only if bias=True
|
||||||
+forward(x) Tensor
|
+forward(x) Tensor
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -365,7 +400,7 @@ classDiagram
|
|||||||
|
|
||||||
class GradientClippingCallback {
|
class GradientClippingCallback {
|
||||||
+float max_grad_norm
|
+float max_grad_norm
|
||||||
+on_step_begin(context)
|
+on_step_end(context)
|
||||||
}
|
}
|
||||||
|
|
||||||
class CheckpointCallback {
|
class CheckpointCallback {
|
||||||
@@ -410,15 +445,24 @@ classDiagram
|
|||||||
+shutdown()
|
+shutdown()
|
||||||
}
|
}
|
||||||
|
|
||||||
class InferenceScheduler {
|
class Executor {
|
||||||
+nn.Module model
|
+AutoModel model
|
||||||
+AutoTokenizer tokenizer
|
+AutoTokenizer tokenizer
|
||||||
|
+KVCache page_cache
|
||||||
|
+execute_prefill(tasks, prompt_len, start_pos)
|
||||||
|
+execute_decode(tasks) List[int]
|
||||||
|
}
|
||||||
|
|
||||||
|
class InferenceScheduler {
|
||||||
+KVCache _page_cache
|
+KVCache _page_cache
|
||||||
|
+Executor _executor
|
||||||
|
+TaskManager _task_mgr
|
||||||
|
+bool _running
|
||||||
|
+Thread _loop_thread
|
||||||
+int max_batch_size
|
+int max_batch_size
|
||||||
+int max_seq_len
|
+int max_seq_len
|
||||||
+int max_prompt_len
|
+int max_prompt_len
|
||||||
+int page_size
|
+int page_size
|
||||||
+TaskManager _task_mgr
|
|
||||||
+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
|
||||||
+remove_task(task_id)
|
+remove_task(task_id)
|
||||||
+start()
|
+start()
|
||||||
@@ -428,8 +472,8 @@ classDiagram
|
|||||||
|
|
||||||
class Allocator {
|
class Allocator {
|
||||||
+int _free_mask
|
+int _free_mask
|
||||||
+int refs_count
|
+List[int] _refs
|
||||||
+LRU _lru
|
+OrderedDict _lru
|
||||||
+alloc() int
|
+alloc() int
|
||||||
+free(idx, keep_cached)
|
+free(idx, keep_cached)
|
||||||
+inc_ref(idx)
|
+inc_ref(idx)
|
||||||
@@ -523,6 +567,19 @@ classDiagram
|
|||||||
ABORTED
|
ABORTED
|
||||||
}
|
}
|
||||||
|
|
||||||
|
class TaskManager {
|
||||||
|
+AutoTokenizer tokenizer
|
||||||
|
+Deque waiting_queue
|
||||||
|
+List active_tasks
|
||||||
|
+add_task(prompt, **kwargs) str
|
||||||
|
+remove_task(task_id) List[Task]
|
||||||
|
+remove_finished_tasks(stop_ids) List[Task]
|
||||||
|
+pull_candidates(n) List[Task]
|
||||||
|
+activate(task)
|
||||||
|
+return_to_waiting(tasks)
|
||||||
|
+get_active_tasks() List[Task]
|
||||||
|
}
|
||||||
|
|
||||||
class GenerationRequest {
|
class GenerationRequest {
|
||||||
+List[Dict] messages
|
+List[Dict] messages
|
||||||
+int top_k
|
+int top_k
|
||||||
@@ -564,9 +621,9 @@ classDiagram
|
|||||||
+List[bool] _done
|
+List[bool] _done
|
||||||
+append(token, idx)
|
+append(token, idx)
|
||||||
+get_results() List[str]
|
+get_results() List[str]
|
||||||
+pop_all() List[str]
|
+pop_all() List[Tuple[int, str]]
|
||||||
+wait(timeout) bool
|
+wait(timeout) bool
|
||||||
+wait_completion()
|
+wait_completion(timeout)
|
||||||
}
|
}
|
||||||
|
|
||||||
class ChatMessage {
|
class ChatMessage {
|
||||||
@@ -584,6 +641,65 @@ classDiagram
|
|||||||
+Optional[str] stop
|
+Optional[str] stop
|
||||||
+Optional[int] n
|
+Optional[int] n
|
||||||
}
|
}
|
||||||
|
|
||||||
|
class AnthropicMessage {
|
||||||
|
+str role
|
||||||
|
+Union[str, List[Dict]] content
|
||||||
|
}
|
||||||
|
|
||||||
|
class MessagesRequest {
|
||||||
|
+List[AnthropicMessage] messages
|
||||||
|
+Optional[str] system
|
||||||
|
+float temperature
|
||||||
|
+float top_p
|
||||||
|
+int top_k
|
||||||
|
+int max_tokens
|
||||||
|
+bool stream
|
||||||
|
+Optional[List[str]] stop_sequences
|
||||||
|
}
|
||||||
|
|
||||||
|
class ProtocolHandler {
|
||||||
|
<<abstract>>
|
||||||
|
+build_prompt() str
|
||||||
|
+create_response_id() str
|
||||||
|
+format_stream_start(ctx) List[str]
|
||||||
|
+format_stream_token(ctx, token) str
|
||||||
|
+format_stream_end(ctx) List[str]
|
||||||
|
+format_non_stream_response(ctx, content) Dict
|
||||||
|
+handle() Union[StreamingResponse, Dict]
|
||||||
|
}
|
||||||
|
|
||||||
|
class OpenAIHandler {
|
||||||
|
+build_prompt() str
|
||||||
|
+create_response_id() str
|
||||||
|
}
|
||||||
|
|
||||||
|
class AnthropicHandler {
|
||||||
|
+List[str] stop_sequences
|
||||||
|
+build_prompt() str
|
||||||
|
+create_response_id() str
|
||||||
|
+on_token(ctx, token, stop_checker) Optional[str]
|
||||||
|
}
|
||||||
|
|
||||||
|
class StopChecker {
|
||||||
|
+check(text) Optional[str]
|
||||||
|
+trim(text, matched) str
|
||||||
|
}
|
||||||
|
|
||||||
|
class StreamContext {
|
||||||
|
+str resp_id
|
||||||
|
+int created
|
||||||
|
+str model
|
||||||
|
+int prompt_tokens
|
||||||
|
+int completion_tokens
|
||||||
|
+str accumulated
|
||||||
|
+Optional[str] stop_matched
|
||||||
|
}
|
||||||
|
|
||||||
|
class app {
|
||||||
|
<<singleton>>
|
||||||
|
+FastAPI app
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
namespace parallel {
|
namespace parallel {
|
||||||
@@ -610,170 +726,156 @@ classDiagram
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
%% Relationships
|
%% Relationships — UML notation: <|-- generalization, *-- composition, o-- aggregation, --> association, ..> dependency
|
||||||
TrainConfig --> BaseDataset : uses
|
|
||||||
TrainConfig ..> BaseStrategy : selects
|
%% --- Generalization (inheritance) ---
|
||||||
StrategyFactory ..> BaseStrategy : creates
|
|
||||||
BaseStrategy <|-- SEQStrategy
|
BaseStrategy <|-- SEQStrategy
|
||||||
BaseStrategy <|-- SFTStrategy
|
BaseStrategy <|-- SFTStrategy
|
||||||
BaseStrategy <|-- DPOStrategy
|
BaseStrategy <|-- DPOStrategy
|
||||||
BaseStrategy <|-- GRPOStrategy
|
BaseStrategy <|-- GRPOStrategy
|
||||||
DPOStrategy --> Transformer : uses
|
|
||||||
GRPOStrategy --> Transformer : uses
|
|
||||||
Trainer --> TrainConfig : uses
|
|
||||||
Trainer --> TrainContextBuilder : uses
|
|
||||||
Trainer --> TrainCallback : manages
|
|
||||||
TrainContextBuilder --> TrainContext : creates
|
|
||||||
TrainContextBuilder --> StrategyFactory : uses
|
|
||||||
Checkpoint ..> Checkpoint : serializes
|
|
||||||
TrainContext --> Checkpoint : manages
|
|
||||||
TrainContext --> BaseStrategy : uses
|
|
||||||
TrainContext --> BaseScheduler : uses
|
|
||||||
SchedulerFactory ..> BaseScheduler : creates
|
|
||||||
BaseScheduler <|-- CosineScheduler
|
BaseScheduler <|-- CosineScheduler
|
||||||
BaseScheduler <|-- SGDRScheduler
|
BaseScheduler <|-- SGDRScheduler
|
||||||
CallbackFactory ..> TrainCallback : creates
|
|
||||||
TrainCallback <|-- GradientClippingCallback
|
TrainCallback <|-- GradientClippingCallback
|
||||||
TrainCallback <|-- CheckpointCallback
|
TrainCallback <|-- CheckpointCallback
|
||||||
TrainCallback <|-- ProgressBarCallback
|
TrainCallback <|-- ProgressBarCallback
|
||||||
TrainCallback <|-- MetricLoggerCallback
|
TrainCallback <|-- MetricLoggerCallback
|
||||||
PagePool --> Allocator : composes
|
|
||||||
PagePool --> PrefixCache : composes
|
|
||||||
KVCache --> PagePool : composes
|
|
||||||
KVCache --> Storage : composes
|
|
||||||
KVCache --> TaskTable : composes
|
|
||||||
KvcacheView --> Storage : wraps
|
|
||||||
InferenceEngine --> InferenceScheduler : uses
|
|
||||||
InferenceEngine --> GenerationRequest : uses
|
|
||||||
InferenceEngine --> GenerateResult : creates
|
|
||||||
InferenceScheduler --> Task : manages
|
|
||||||
InferenceScheduler --> TaskStatus : uses
|
|
||||||
InferenceScheduler --> KVCache : uses
|
|
||||||
InferenceScheduler --> Transformer : uses
|
|
||||||
Task --> TaskStatus : uses
|
|
||||||
InferenceEngine --> Transformer : uses
|
|
||||||
BaseSamplingStrategy <|-- TemperatureStrategy
|
|
||||||
BaseSamplingStrategy <|-- TopKStrategy
|
|
||||||
BaseSamplingStrategy <|-- TopPStrategy
|
|
||||||
SamplingPipeline --> BaseSamplingStrategy : composes
|
|
||||||
BaseDataset <|-- SEQDataset
|
BaseDataset <|-- SEQDataset
|
||||||
BaseDataset <|-- SFTDataset
|
BaseDataset <|-- SFTDataset
|
||||||
BaseDataset <|-- DPODataset
|
BaseDataset <|-- DPODataset
|
||||||
BaseDataset <|-- GRPODataset
|
BaseDataset <|-- GRPODataset
|
||||||
DatasetFactory ..> BaseDataset : creates
|
|
||||||
BaseStorage <|-- H5Storage
|
BaseStorage <|-- H5Storage
|
||||||
BaseStorage <|-- JSONStorage
|
BaseStorage <|-- JSONStorage
|
||||||
BaseDataset --> BaseStorage : uses
|
BaseSamplingStrategy <|-- TemperatureStrategy
|
||||||
MultiSegmentFetcher --> BaseSegmentFetcher : uses
|
BaseSamplingStrategy <|-- TopKStrategy
|
||||||
AutoModel <|-- Transformer
|
BaseSamplingStrategy <|-- TopPStrategy
|
||||||
AutoModel --> ModelConfig : contains
|
|
||||||
Transformer --> DecoderBlock : uses
|
|
||||||
Transformer --> RotaryEmbedding : uses
|
|
||||||
Transformer --> Embedding : uses
|
|
||||||
DecoderBlock --> GQA : uses
|
|
||||||
DecoderBlock --> MLP : uses
|
|
||||||
DecoderBlock --> RMSNorm : uses
|
|
||||||
TrainContextBuilder --> ResumableDistributedSampler : creates
|
|
||||||
ResumableDistributedSampler --> BaseDataset : samples
|
|
||||||
ParallelModel <|-- RowParallelLinear
|
ParallelModel <|-- RowParallelLinear
|
||||||
ParallelModel <|-- ColumnParallelLinear
|
ParallelModel <|-- ColumnParallelLinear
|
||||||
AutoTokenizer --> ChatTemplate : uses
|
AutoModel <|-- Transformer
|
||||||
|
BaseModelConfig <|-- ModelConfig
|
||||||
BaseFactory <|-- AutoModel
|
BaseFactory <|-- AutoModel
|
||||||
|
BaseFactory <|-- AttnFactory
|
||||||
|
BaseFactory <|-- FFNFactory
|
||||||
BaseFactory <|-- DatasetFactory
|
BaseFactory <|-- DatasetFactory
|
||||||
BaseFactory <|-- StrategyFactory
|
BaseFactory <|-- StrategyFactory
|
||||||
BaseFactory <|-- SchedulerFactory
|
BaseFactory <|-- SchedulerFactory
|
||||||
BaseFactory <|-- CallbackFactory
|
BaseFactory <|-- CallbackFactory
|
||||||
|
ProtocolHandler <|-- OpenAIHandler
|
||||||
|
ProtocolHandler <|-- AnthropicHandler
|
||||||
|
|
||||||
|
%% --- Composition (strong ownership, part destroyed with whole) ---
|
||||||
|
KVCache *-- PagePool
|
||||||
|
KVCache *-- Storage
|
||||||
|
KVCache *-- TaskTable
|
||||||
|
KVCache *-- Allocator
|
||||||
|
KVCache *-- PrefixCache
|
||||||
|
InferenceEngine *-- InferenceScheduler
|
||||||
|
InferenceScheduler *-- KVCache
|
||||||
|
InferenceScheduler *-- Executor
|
||||||
|
InferenceScheduler *-- TaskManager
|
||||||
|
SamplingPipeline *-- BaseSamplingStrategy
|
||||||
|
TrainContextBuilder *-- TrainContext
|
||||||
|
Transformer *-- DecoderBlock
|
||||||
|
Transformer *-- RotaryEmbedding
|
||||||
|
Transformer *-- Embedding
|
||||||
|
DecoderBlock *-- RMSNorm
|
||||||
|
BaseDataset *-- BaseStorage
|
||||||
|
ChatCompletionRequest *-- ChatMessage
|
||||||
|
MessagesRequest *-- AnthropicMessage
|
||||||
|
|
||||||
|
%% --- Aggregation (weak ownership) ---
|
||||||
|
AutoModel o-- ModelConfig
|
||||||
|
Trainer o-- TrainCallback
|
||||||
|
TrainContext o-- BaseStrategy
|
||||||
|
TrainContext o-- BaseScheduler
|
||||||
|
TrainContext o-- Checkpoint
|
||||||
|
AutoTokenizer o-- ChatTemplate
|
||||||
|
KvcacheView o-- Storage
|
||||||
|
BaseFactory o-- Registry
|
||||||
|
|
||||||
|
%% --- Dependency (uses temporarily) ---
|
||||||
|
TrainConfig ..> BaseStrategy : selects
|
||||||
|
StrategyFactory ..> BaseStrategy : creates
|
||||||
|
SchedulerFactory ..> BaseScheduler : creates
|
||||||
|
DatasetFactory ..> BaseDataset : creates
|
||||||
|
CallbackFactory ..> TrainCallback : creates
|
||||||
|
AttnFactory ..> GQA : creates
|
||||||
|
AttnFactory ..> MLA : creates
|
||||||
|
FFNFactory ..> MLP : creates
|
||||||
|
FFNFactory ..> DeepSeekMoE : creates
|
||||||
|
DecoderBlock ..> AttnFactory : uses
|
||||||
|
DecoderBlock ..> FFNFactory : uses
|
||||||
|
Trainer ..> TrainContextBuilder : uses
|
||||||
|
Trainer ..> Functions : spawns
|
||||||
|
TrainContextBuilder ..> StrategyFactory : uses
|
||||||
|
TrainContextBuilder ..> ResumableDistributedSampler : creates
|
||||||
|
Checkpoint ..> Checkpoint : serializes
|
||||||
|
CheckpointCallback ..> Checkpoint : creates
|
||||||
|
KVCache ..> KvcacheView : binds
|
||||||
|
InferenceEngine ..> GenerationRequest : uses
|
||||||
|
InferenceEngine ..> GenerateResult : creates
|
||||||
|
OpenAIHandler ..> ChatCompletionRequest : receives
|
||||||
|
AnthropicHandler ..> MessagesRequest : receives
|
||||||
|
ProtocolHandler ..> StopChecker : creates
|
||||||
|
ProtocolHandler ..> StreamContext : creates
|
||||||
|
|
||||||
|
%% --- Association (general usage) ---
|
||||||
|
Trainer --> TrainConfig
|
||||||
|
DPOStrategy --> Transformer
|
||||||
|
GRPOStrategy --> Transformer
|
||||||
|
InferenceScheduler --> Task
|
||||||
|
InferenceScheduler --> TaskStatus
|
||||||
|
Task --> TaskStatus
|
||||||
|
InferenceEngine --> Transformer
|
||||||
|
Executor --> Transformer
|
||||||
|
Executor --> AutoTokenizer
|
||||||
|
TaskManager --> AutoTokenizer
|
||||||
|
MultiSegmentFetcher --> BaseSegmentFetcher
|
||||||
|
ResumableDistributedSampler --> BaseDataset
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### Module Overview
|
|
||||||
|
## Module Overview
|
||||||
|
|
||||||
| 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, BaseStorage, H5Storage, JSONStorage, BaseSegmentFetcher, MultiSegmentFetcher, ResumableDistributedSampler, DatasetFactory, save_h5, load_h5 | Dataset loading and management |
|
| **astrai.dataset** | BaseDataset–GRPODataset, BaseStorage–JSONStorage, BaseSegmentFetcher, MultiSegmentFetcher, ResumableDistributedSampler, DatasetFactory | Dataset loading and management |
|
||||||
| **astrai.serialization** | Checkpoint | Model serialization and checkpoint management |
|
| **astrai.serialization** | Checkpoint | Model serialization |
|
||||||
| **astrai.model** | AutoModel, Transformer, DecoderBlock, GQA, MLA, MLP, RMSNorm, Linear, RotaryEmbedding, Embedding | Neural network model |
|
| **astrai.model** | AutoModel, Transformer, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, 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–GRPOStrategy, StrategyFactory, BaseScheduler–SGDRScheduler, SchedulerFactory, TrainCallback–MetricLoggerCallback, CallbackFactory | Training workflow |
|
||||||
| **astrai.inference** | InferenceEngine, InferenceScheduler, KVCache, KvcacheView, Allocator, PrefixCache, PagePool, Storage, TaskTable, Task, TaskStatus, GenerationRequest, BaseSamplingStrategy, TemperatureStrategy, TopKStrategy, TopPStrategy, SamplingPipeline, ChatMessage, ChatCompletionRequest | Inference service with continuous batching and paged KV cache |
|
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, KVCache–KvcacheView, Allocator–Storage, Task, TaskManager, TaskStatus, GenerationRequest, BaseSamplingStrategy–SamplingPipeline, ProtocolHandler–AnthropicHandler, ChatMessage–MessagesRequest, app | Inference service |
|
||||||
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank, get_world_size, get_current_device, ParallelModel, ColumnParallelLinear, RowParallelLinear | Distributed parallel |
|
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, ParallelModel, RowParallelLinear, ColumnParallelLinear | Distributed parallel |
|
||||||
| **astrai.factory** | Registry, BaseFactory | Generic component registration |
|
| **astrai.factory** | Registry, BaseFactory[T] | Component registration |
|
||||||
|
|
||||||
### Design Patterns
|
## Design Patterns
|
||||||
|
|
||||||
| Pattern | Classes | Purpose |
|
| Pattern | Classes | Purpose |
|
||||||
|---------|---------|---------|
|
|---------|---------|---------|
|
||||||
| **Strategy** | `BaseStrategy`, `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy`, `StrategyFactory` | Flexible training strategy switching, supports SEQ/SFT/DPO/GRPO |
|
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory` | Decorator-based component creation |
|
||||||
| **Builder** | `TrainContextBuilder` | Chain-building training context, step-by-step initialization of components |
|
| **Registry** | `BaseFactory`, `Registry` | Component registration with category/priority |
|
||||||
| **Factory** | `StrategyFactory`, `SchedulerFactory`, `DatasetFactory`, `CallbackFactory`, `BaseFactory` | Decorator registration mechanism, dynamically create training strategies, schedulers, datasets, and callbacks |
|
| **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching |
|
||||||
| **Observer** | `TrainCallback`, `CallbackFactory` | Callback mechanism for training process monitoring (checkpoint, early stopping, metrics) |
|
| **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations |
|
||||||
| **Context** | `TrainContext` | Training process state container with model, optimizer, scheduler and checkpoint |
|
| **Template Method** | `ProtocolHandler`, `OpenAIHandler`, `AnthropicHandler` | HTTP API handler with format hooks |
|
||||||
| **Registry** | `BaseFactory`, `Registry` | Generic component registration with category and priority support |
|
| **Builder** | `TrainContextBuilder` | Chain-building training context |
|
||||||
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with O(1) alloc/free via bitmask + LRU eviction |
|
| **Observer** | `TrainCallback`, callback implementations | Training process monitoring |
|
||||||
| **Strategy (Sampling)** | `BaseSamplingStrategy`, `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations with temperature, top-k, top-p |
|
| **Context** | `TrainContext` | Unified training state bag |
|
||||||
| **Producer-Consumer** | `InferenceScheduler`, `Task`, `waiting_queue`, `active_tasks` | Continuous batching with dynamic task queue management |
|
| **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction |
|
||||||
| **Event-Driven** | `threading.Event`, `_task_event` | Non-blocking wait mechanism for task scheduling using Python's `threading` module |
|
| **Storage** | `BaseStorage`, `H5Storage`, `JSONStorage` | Format-agnostic data access |
|
||||||
| **AutoModel Registry** | `AutoModel`, `Transformer` | Model type registration and dynamic loading via decorator pattern |
|
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
|
||||||
| **Generator Pattern** | `GenerateResult`, `GenerationRequest` | Event-based result notification for streaming/non-streaming generation |
|
| **AutoModel Registry** | `AutoModel`, `Transformer` | Model-type dynamic loading |
|
||||||
|
|
||||||
### Core Relationships
|
## Core Relationships
|
||||||
|
|
||||||
1. **Configuration → Training**: `TrainConfig` holds model, dataset, optimizer_fn, scheduler_fn and other training configuration references
|
1. **Config → Training**: `TrainConfig` holds model, dataset, optimizer_fn, scheduler_fn
|
||||||
2. **Training Flow**: `Trainer` → `TrainContextBuilder` → `TrainContext`, uses `BaseStrategy` to compute loss
|
2. **Training Flow**: `Trainer` → `TrainContextBuilder` → `TrainContext`, uses `BaseStrategy` for loss
|
||||||
3. **Strategy Selection**: `StrategyFactory` creates corresponding strategy instance based on `train_type`
|
3. **Strategy Selection**: `StrategyFactory` creates strategy by `train_type`
|
||||||
4. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `Transformer`, uses `KVCache` (backed by `Allocator` + `PrefixCache` + `PagePool` + `Storage`) for paged KV cache management and `SamplingPipeline` for efficient continuous batching with streaming/non-streaming
|
4. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `Transformer`, backed by `KVCache` + `SamplingPipeline`
|
||||||
5. **Distributed Support**: `spawn_parallel_fn` and `setup_parallel` provide multi-process training capability for `Trainer`
|
5. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
|
||||||
6. **Dataset Loading**: `DatasetFactory` creates datasets (SEQDataset, SFTDataset, DPODataset, GRPODataset), supports HDF5 loading via `BaseSegmentFetcher` and `MultiSegmentFetcher`
|
6. **Dataset Loading**: `DatasetFactory` creates datasets, `BaseStorage` (H5Storage/JSONStorage) loads via `BaseSegmentFetcher` + `MultiSegmentFetcher`
|
||||||
7. **Checkpoint Management**: `Checkpoint` handles model state serialization/deserialization with safetensors
|
7. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only)
|
||||||
8. **Scheduler Support**: `SchedulerFactory` creates learning rate schedulers (CosineScheduler, SGDRScheduler)
|
8. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`
|
||||||
9. **AutoModel Loading**: `AutoModel.from_pretrained()` dynamically loads model based on `config.json` model_type, uses `Registry` pattern for model type registration
|
9. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
|
||||||
|
|
||||||
## 3. Training Process
|
> Document Update Time: 2026-05-15
|
||||||
|
|
||||||
The common training process for large language models (LLM) typically includes three stages: **Pre-training (SEQ)**, **Supervised Fine-Tuning (SFT)**, and **Reinforcement Learning from Human Feedback (DPO/GRPO)**. This system is designed to support seamless end-to-end flow, achieving efficient switching and state management of different training stages through modular strategies.
|
|
||||||
|
|
||||||
### Core Formulas
|
|
||||||
|
|
||||||
**Pre-training (SEQ):**
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{PT}} = - \sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta)
|
|
||||||
$$
|
|
||||||
|
|
||||||
**SFT:**
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{SFT}} = - \sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
|
|
||||||
$$
|
|
||||||
|
|
||||||
**DPO:**
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{DPO}} = -\mathbb{E}_{(x, y_w, y_l) \sim D} \left[ \log \sigma\left( \beta \log \frac{\pi_\theta(y_w \mid x)}{\pi_{\text{ref}}(y_w \mid x)} - \beta \log \frac{\pi_\theta(y_l \mid x)}{\pi_{\text{ref}}(y_l \mid x)} \right) \right]
|
|
||||||
$$
|
|
||||||
|
|
||||||
**GRPO:**
|
|
||||||
|
|
||||||
GRPO (Group Relative Policy Optimization) computes advantages from multiple responses to the same prompt, then optimizes using a PPO-style clipped objective:
|
|
||||||
|
|
||||||
$$
|
|
||||||
\text{Advantage}_i = \frac{r_i - \mu}{\sigma + \epsilon}
|
|
||||||
$$
|
|
||||||
|
|
||||||
Where $r_i$ is the reward for the $i$-th response, $\mu$ and $\sigma$ are the mean and standard deviation of group rewards.
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{\text{GRPO}} = -\mathbb{E} \left[ \min\left( \frac{\pi_\theta(a|s)}{\pi_{\text{ref}}(a|s)} \cdot A, \text{clip}\left(\frac{\pi_\theta(a|s)}{\pi_{\text{ref}}(a|s)}, 1-\epsilon, 1+\epsilon\right) \cdot A \right) \right] + \lambda \cdot D_{KL}
|
|
||||||
$$
|
|
||||||
|
|
||||||
The KL divergence term uses mean squared error approximation:
|
|
||||||
|
|
||||||
$$
|
|
||||||
L_{KL} = \lambda \cdot \mathbb{E} \left[ (\log \pi_\theta - \log \pi_{\text{ref}})^2 \right]
|
|
||||||
$$
|
|
||||||
|
|
||||||
The final loss is the sum of both: $L = L_{\text{policy}} + L_{KL}$
|
|
||||||
|
|
||||||
Through the above three-stage progressive training, the model completes its evolution from a general language foundation to a specialized, highly-aligned dialogue intelligence.
|
|
||||||
|
|
||||||
> Document Update Time: 2026-05-14
|
|
||||||
+32
-212
@@ -1,237 +1,57 @@
|
|||||||
# AstrAI Data Flow Documentation
|
# Data Flow
|
||||||
|
|
||||||
This document describes the data flow of the AstrAI project (a training and inference framework for autoregressive Transformer language models). It covers the complete flow from raw data to model training and inference.
|
This document describes the data pipeline: from raw text to model input tensors.
|
||||||
|
|
||||||
## Overview
|
## Overview
|
||||||
|
|
||||||
AstrAI adopts a modular design with the following main components:
|
```
|
||||||
- **Dataset Module** (`astrai/dataset/`): Dataset, sampler, serialization tools
|
Raw Text → AutoTokenizer → Token IDs → .h5/.json → Dataset → Sampler → DataLoader → Training/Inference
|
||||||
- **Model Module** (`astrai/model/`): AutoModel, Transformer model and its submodules
|
|
||||||
- **Training Module** (`astrai/trainer/`): Trainer, training context, strategies, schedulers, callbacks, metric utilities
|
|
||||||
- **Inference Module** (`astrai/inference/`): Inference engine with continuous batching, streaming generation
|
|
||||||
- **Config Module** (`astrai/config/`): ModelConfig, TrainConfig
|
|
||||||
- **Factory Module** (`astrai/factory/`): Registry, BaseFactory for component registration
|
|
||||||
- **Parallel Module** (`astrai/parallel/`): Distributed training support
|
|
||||||
- **Serialization** (`astrai/serialization.py`): Checkpoint management with safetensors
|
|
||||||
|
|
||||||
## Data Flow Diagram
|
|
||||||
|
|
||||||
```mermaid
|
|
||||||
flowchart LR
|
|
||||||
subgraph A[Data Preparation]
|
|
||||||
direction TB
|
|
||||||
A1[Raw Text] --> A2[AutoTokenizer]
|
|
||||||
A2 --> A3[Tokenized .h5 files]
|
|
||||||
A3 --> A4[BaseDataset]
|
|
||||||
A4 --> A5[ResumableDistributedSampler]
|
|
||||||
A5 --> A6[DataLoader]
|
|
||||||
end
|
|
||||||
|
|
||||||
subgraph B[Training]
|
|
||||||
direction TB
|
|
||||||
B1[DataLoader] --> B2[BaseStrategy]
|
|
||||||
B2 --> B3[Transformer Forward]
|
|
||||||
B3 --> B4[Loss + Backward]
|
|
||||||
B4 --> B5[Gradient Accumulation]
|
|
||||||
B5 -->|every accum_steps| B6[Optimizer Step]
|
|
||||||
B6 --> B7[LR Scheduler]
|
|
||||||
B7 -->|next batch| B2
|
|
||||||
B6 --> B8[CheckpointCallback]
|
|
||||||
end
|
|
||||||
|
|
||||||
subgraph C[Inference]
|
|
||||||
direction TB
|
|
||||||
C1[Checkpoint] --> C2[AutoModel]
|
|
||||||
C1 --> C3[AutoTokenizer]
|
|
||||||
C2 --> C4[InferenceEngine]
|
|
||||||
C3 --> C4
|
|
||||||
C4 --> C5[InferenceScheduler]
|
|
||||||
C5 --> C6[Transformer Forward]
|
|
||||||
C6 --> C7[sample]
|
|
||||||
C7 --> C8{End?}
|
|
||||||
C8 -->|No| C6
|
|
||||||
C8 -->|Yes| C9[Generated Text]
|
|
||||||
end
|
|
||||||
|
|
||||||
A --> B
|
|
||||||
B --> C
|
|
||||||
```
|
```
|
||||||
|
|
||||||
## Detailed Module Descriptions
|
## Data Preparation
|
||||||
|
|
||||||
### 1. Data Serialization (`astrai/dataset/storage.py` & `astrai/serialization.py`)
|
Raw text is tokenized via `AutoTokenizer.encode()` and saved as HDF5 (`.h5`) or JSON (`.json`/`.jsonl`) files with keyed tensor groups.
|
||||||
|
|
||||||
- **`save_h5`**: Saves tensors by groups as HDF5 files (`.h5`), each key maps to a list of tensors
|
Storage format is auto-detected by `detect_format()`; backends are dispatched via registry:
|
||||||
- **`load_h5`**: Loads `.h5` files, returns `Dict[str, List[Tensor]]`, supports shared memory
|
|
||||||
- **`Checkpoint`**: Encapsulates model state dict + epoch + iteration; uses safetensors
|
|
||||||
|
|
||||||
### 2. Dataset Module
|
|
||||||
|
|
||||||
#### 2.1 Dataset (`dataset.py`)
|
|
||||||
- **`BaseDataset`**: Abstract base class for windowed sequence sampling
|
|
||||||
- **`BaseSegmentFetcher` / `MultiSegmentFetcher`**: Fetch tensor segments by index range
|
|
||||||
- **`DatasetFactory`**: Creates dataset instances by `train_type` (`seq`, `sft`, `dpo`, `grpo`)
|
|
||||||
- Data keys: `"sequence"` (SEQ), `"loss_mask"` (SFT), `"chosen_mask"/"rejected_mask"` (DPO), `"masks"` (GRPO)
|
|
||||||
|
|
||||||
#### 2.2 Sampler (`sampler.py`)
|
|
||||||
- **`ResumableDistributedSampler`**: Tracks `epoch` and `iter` for breakpoint resume; supports shuffle and drop_last
|
|
||||||
|
|
||||||
### 3. Model Module
|
|
||||||
|
|
||||||
#### 3.1 Transformer / AutoModel
|
|
||||||
- **`AutoModel`**: Base class with `from_pretrained()` / `save_pretrained()`
|
|
||||||
- **`Transformer`**: Decoder-only architecture, registered via `@AutoModel.register('transformer')`
|
|
||||||
- Embedding → N×DecoderBlock → RMSNorm → Linear lm_head
|
|
||||||
- RoPE position encoding, optional weight tying
|
|
||||||
|
|
||||||
#### 3.2 Submodules (`module.py`)
|
|
||||||
- **`DecoderBlock`**: GQA attention + residual + MLP + RMSNorm
|
|
||||||
- **`GQA`**: Grouped Query Attention (also `MLA` for multi-latent attention)
|
|
||||||
- **`MLP`**: `SiLU(gate(x)) * up(x)` → down projection
|
|
||||||
- **`RotaryEmbedding`**: RoPE complex cache (freqs_cis)
|
|
||||||
- **`RMSNorm`**: Layer normalization
|
|
||||||
|
|
||||||
### 4. Training Module
|
|
||||||
|
|
||||||
#### 4.1 Training Context (`train_context.py`)
|
|
||||||
- **`TrainContext`**: Dataclass holding model, optimizer, dataloader, strategy, scheduler, checkpoint state
|
|
||||||
- **`TrainContextBuilder`**: Builder pattern — takes checkpoint for resume, builds all components
|
|
||||||
|
|
||||||
#### 4.2 Trainer (`trainer.py`)
|
|
||||||
|
|
||||||
The training loop is nested: **epoch** → **batch** (with step phase interspersed):
|
|
||||||
|
|
||||||
```
|
```
|
||||||
on_train_begin
|
create_storage("h5") → H5Storage
|
||||||
on_epoch_begin
|
create_storage("json") → JSONStorage
|
||||||
for each accumulation window of batches: ← step phase
|
|
||||||
on_step_begin
|
|
||||||
for each batch in window: ← batch phase
|
|
||||||
on_batch_begin → strategy(batch) → loss → backward → on_batch_end
|
|
||||||
iteration += 1
|
|
||||||
on_step_end
|
|
||||||
optimizer.step() → zero_grad
|
|
||||||
|
|
||||||
on_epoch_end
|
|
||||||
on_train_end
|
|
||||||
```
|
```
|
||||||
|
|
||||||
Key points:
|
Both support shared memory via `.share_memory_()`.
|
||||||
- `on_step_*` fires every `accumulation_steps` batches, wrapping optimizer step AFTER the hook
|
|
||||||
- `on_batch_*` fires every batch, wrapping loss computation
|
|
||||||
- `GradientClippingCallback` fires on `on_step_end`
|
|
||||||
- LR scheduler steps inline (no `SchedulerCallback` class)
|
|
||||||
|
|
||||||
#### 4.3 Strategy (`strategy.py`)
|
## Data Keys by Training Type
|
||||||
- **`SEQStrategy`**: Next-token prediction, cross-entropy with label smoothing
|
|
||||||
- **`SFTStrategy`**: Supervised fine-tuning with loss masking
|
|
||||||
- **`DPOStrategy`**: Direct Preference Optimization with reference model
|
|
||||||
- **`GRPOStrategy`**: Group Relative Policy Optimization with clipped ratio
|
|
||||||
|
|
||||||
#### 4.4 Scheduler (`schedule.py`)
|
| Type | Storage Keys |
|
||||||
- **`CosineScheduler`**: Cosine decay + linear warmup
|
|------|-------------|
|
||||||
- **`SGDRScheduler`**: Cosine annealing with warm restarts
|
| `seq` | `sequence` (→ input_ids, target_ids via offset-by-1) |
|
||||||
- Created by `SchedulerFactory` and bound to optimizer
|
| `sft` | `sequence`, `loss_mask` |
|
||||||
|
| `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` |
|
||||||
|
| `grpo` | `prompts`, `responses`, `masks`, `rewards` |
|
||||||
|
|
||||||
#### 4.5 Callbacks
|
## Dataset Architecture
|
||||||
- **`CheckpointCallback`**: Saves safetensors at `ckpt_interval` iterations
|
|
||||||
- **`ProgressBarCallback`**: tqdm progress display
|
|
||||||
- **`MetricLoggerCallback`**: Writes JSONL metrics to `{ckpt_dir}/logs/`
|
|
||||||
- **`GradientClippingCallback`**: `clip_grad_norm_` on `on_step_end`
|
|
||||||
|
|
||||||
### 5. Inference Module
|
|
||||||
|
|
||||||
#### 5.1 Inference Engine (`engine.py`)
|
|
||||||
- **`InferenceEngine`**: Facade over scheduler; provides `generate()`, `generate_with_request()`, `generate_async()`
|
|
||||||
- Accepts `prompt: str | List[str]`, returns generator (stream) or string (non-stream)
|
|
||||||
|
|
||||||
#### 5.2 Scheduler 4-Phase Loop (`scheduler.py`)
|
|
||||||
|
|
||||||
Background thread runs continuously:
|
|
||||||
|
|
||||||
```
|
```
|
||||||
1. Cleanup → Remove finished tasks, free KV cache pages
|
DatasetFactory.load(train_type, path, window_size, stride)
|
||||||
2. Refill → Pop from waiting_queue, alloc pages, add to active
|
→ create_storage(detect_format(path))
|
||||||
3. Prefill → Group active tasks by prompt_len, run full forward pass
|
→ MultiSegmentFetcher(BaseSegmentFetcher per key)
|
||||||
4. Decode → Pick largest same-position group, run single-token forward
|
→ BaseDataset.__getitem__(idx)
|
||||||
|
→ sliding window [begin, end) via get_index(idx)
|
||||||
```
|
```
|
||||||
|
|
||||||
- **`Task`**: Tracks prompt_ids, output_ids, status (PENDING/RUNNING/FINISHED/ABORTED)
|
`window_size` = max input length, `stride` = step between consecutive samples.
|
||||||
- **`KVCache`**: Facade over `Allocator` + `PrefixCache` + `PagePool` + `Storage` for paged KV cache
|
|
||||||
- **`KvcacheView`**: Batch view bundling cache + page table for attention layers
|
|
||||||
- **`sample()`**: Temperature → top-k → top-p → multinomial
|
|
||||||
|
|
||||||
#### 5.3 Server (`server.py`)
|
## Sampler
|
||||||
- FastAPI with OpenAI `/v1/chat/completions` and Anthropic `/v1/messages` endpoints
|
|
||||||
- Streaming via SSE, health check at `/health`, stats at `/stats`
|
|
||||||
|
|
||||||
### 6. Tokenizer Module
|
`ResumableDistributedSampler` supports checkpoint-aware distributed sampling:
|
||||||
|
|
||||||
- **`AutoTokenizer`**: Wraps HuggingFace tokenizers (BBPE); `encode`/`decode`/`apply_chat_template`
|
- Tracks `start_epoch` / `start_iter` for resume
|
||||||
- **`ChatTemplate`**: Jinja2-based template rendering for multi-turn chat
|
- Shuffle via `torch.Generator(seed + epoch)`
|
||||||
|
- Per-replica index slicing for DDP
|
||||||
|
|
||||||
### 7. Factory & Parallel
|
## DataLoader
|
||||||
|
|
||||||
- **`Registry` / `BaseFactory`**: Decorator-based component registration
|
Standard PyTorch `DataLoader` with configurable `batch_size`, `num_workers`, `pin_memory`, `prefetch_factor`. Sampler produces indices; dataloader fetches tensor batches via `__getitem__`.
|
||||||
- **`spawn_parallel_fn`**: Multi-process DDP launcher with NCCL backend
|
|
||||||
- **`ParallelModel` / `ColumnParallelLinear` / `RowParallelLinear`**: Tensor model parallelism
|
|
||||||
|
|
||||||
## Training Data Flow — Detailed Steps
|
> Document Update Time: 2026-05-15
|
||||||
|
|
||||||
1. **Data Preparation**
|
|
||||||
- Raw text → token IDs via `AutoTokenizer.encode()`
|
|
||||||
- Save as `.h5` files (groups of tensor lists per data key)
|
|
||||||
|
|
||||||
2. **Dataset Loading**
|
|
||||||
- `BaseDataset.load()` calls `load_h5()`, builds `MultiSegmentFetcher`
|
|
||||||
- Sliding window of `window_size` with `stride` determines sample boundaries
|
|
||||||
|
|
||||||
3. **Sampling & Batching**
|
|
||||||
- `ResumableDistributedSampler` produces shuffled index sequences
|
|
||||||
- `DataLoader` fetches `[batch_size, window_size]` tensors via `__getitem__`
|
|
||||||
|
|
||||||
4. **Strategy Forward**
|
|
||||||
- Strategy receives batch, calls `Transformer.forward()` for logits
|
|
||||||
- Computes task-specific loss (cross-entropy, DPO, GRPO)
|
|
||||||
|
|
||||||
5. **Backward & Accumulation**
|
|
||||||
- `loss = raw_loss / accumulation_steps`
|
|
||||||
- `loss.backward()` accumulates gradients
|
|
||||||
- Every `accumulation_steps` batches: `optimizer.step()` → `zero_grad()`
|
|
||||||
- Every batch: `scheduler.step()` updates learning rate
|
|
||||||
|
|
||||||
6. **Checkpoint**
|
|
||||||
- `CheckpointCallback` saves `model.state_dict()` + metadata to safetensors at `ckpt_interval` iterations
|
|
||||||
- Does NOT save optimizer/scheduler state (resume resets those)
|
|
||||||
|
|
||||||
## Inference Data Flow — Detailed Steps
|
|
||||||
|
|
||||||
1. **Model Loading**
|
|
||||||
- `AutoModel.from_pretrained(path)` loads weights from safetensors
|
|
||||||
- `torch.inference_mode()` wraps generation
|
|
||||||
|
|
||||||
2. **Prompt Construction**
|
|
||||||
- Messages → `apply_chat_template(messages, tokenize=False)` → prompt string
|
|
||||||
- `tokenizer.encode(prompt)` → token IDs (truncated to `max_prompt_len`)
|
|
||||||
|
|
||||||
3. **Continuous Batching Loop**
|
|
||||||
- **Cleanup**: Finished tasks → `stream_callback(STOP)`, free KV pages
|
|
||||||
- **Refill**: Pop from waiting queue, `PagePool.task_alloc()` for prompt pages
|
|
||||||
- **Prefill**: Group by prompt length, run full forward with `start_pos=0`
|
|
||||||
- **Decode**: Pick position group with most tasks, single-token forward:
|
|
||||||
- Model forward → `logits` → `sample()` → next token ID
|
|
||||||
- Append to `output_ids`, update `output_tokens`
|
|
||||||
- `PagePool.task_alloc()` allocates pages as needed
|
|
||||||
- `stream_callback(token)` for streaming clients
|
|
||||||
|
|
||||||
4. **Output**
|
|
||||||
- `tokenizer.decode(output_ids)` → text
|
|
||||||
- Return to caller (streaming: token-by-token; non-streaming: complete string)
|
|
||||||
|
|
||||||
## Checkpoint & Serialization
|
|
||||||
|
|
||||||
- **Training Checkpoint**: safetensors weights + epoch/iteration metadata. Optimizer/scheduler state is NOT persisted.
|
|
||||||
- **Inference Loading**: `AutoModel.from_pretrained()` loads from the same safetensors format.
|
|
||||||
- **Dataset Serialization**: HDF5 with shared memory support for large-scale pre-training data.
|
|
||||||
|
|
||||||
> Document Update Time: 2026-05-14
|
|
||||||
|
|||||||
@@ -0,0 +1,140 @@
|
|||||||
|
# Inference
|
||||||
|
|
||||||
|
## KV Cache
|
||||||
|
|
||||||
|
At decode time, only the last query token matters. All previous K/V are cached to avoid recomputation:
|
||||||
|
|
||||||
|
$$
|
||||||
|
o_n = \sum_j \text{softmax}\left(\frac{q_n k_j}{\sqrt{d_k}}\right) v_j
|
||||||
|
$$
|
||||||
|
|
||||||
|
RoPE is applied **before** KV cache write, not after — otherwise position encoding drift occurs.
|
||||||
|
|
||||||
|
## KVCache System
|
||||||
|
|
||||||
|
Six classes working together:
|
||||||
|
|
||||||
|
```
|
||||||
|
KVCache (facade)
|
||||||
|
├── Allocator bitmask-based page allocator + ref-count + LRU eviction
|
||||||
|
├── PrefixCache hash-based prefix matching (page_hash via rolling hash)
|
||||||
|
├── PagePool orchestrates Allocator + PrefixCache
|
||||||
|
├── TaskTable maps task_id → page_table + cached token count
|
||||||
|
├── Storage k_cache / v_cache tensors (n_layers × n_pages × page_size × n_kv_heads × head_dim)
|
||||||
|
└── KvcacheView bundles Storage + page_table + total_len for attention layers
|
||||||
|
```
|
||||||
|
|
||||||
|
`KVCache.bind(page_table, total_len)` returns a `KvcacheView` used by attention layers via `write()` / `gather()`.
|
||||||
|
|
||||||
|
## Continuous Batching
|
||||||
|
|
||||||
|
`InferenceScheduler` runs a daemon thread with a 4-phase loop:
|
||||||
|
|
||||||
|
```
|
||||||
|
1. Cleanup → Remove finished tasks, free KV pages
|
||||||
|
2. Refill → Pop from waiting_queue, task_alloc pages, activate
|
||||||
|
3. Prefill → Group by (prompt_len, start_pos), run full forward
|
||||||
|
4. Decode → Pick largest same-position group, single-token forward
|
||||||
|
```
|
||||||
|
|
||||||
|
## Sampling (Strategy Pattern)
|
||||||
|
|
||||||
|
```
|
||||||
|
BaseSamplingStrategy → TemperatureStrategy → TopKStrategy → TopPStrategy
|
||||||
|
```
|
||||||
|
|
||||||
|
`SamplingPipeline` composes them: Temperature → Top-K → Top-P → softmax → multinomial.
|
||||||
|
`sample()` is a convenience shortcut for one-shot usage.
|
||||||
|
|
||||||
|
## Protocol Handlers (Template Method)
|
||||||
|
|
||||||
|
```python
|
||||||
|
class ProtocolHandler(ABC):
|
||||||
|
def handle(self):
|
||||||
|
ctx = StreamContext(...)
|
||||||
|
agen = engine.generate_async(prompt, ...)
|
||||||
|
if stream: self._handle_stream(agen, ctx)
|
||||||
|
else: self._handle_non_stream(agen, ctx)
|
||||||
|
```
|
||||||
|
|
||||||
|
Subclass hooks: `build_prompt()`, `create_response_id()`, `format_stream_start/token/end()`, `format_non_stream_response()`.
|
||||||
|
|
||||||
|
`OpenAIHandler` → `/v1/chat/completions`, `AnthropicHandler` → `/v1/messages`.
|
||||||
|
|
||||||
|
## Engine & GenerateResult
|
||||||
|
|
||||||
|
```
|
||||||
|
InferenceEngine
|
||||||
|
├── generate(prompt, stream, ...) → str | List[str] | Generator
|
||||||
|
├── generate_with_request(req) → same
|
||||||
|
└── generate_async(prompt, ...) → AsyncGenerator
|
||||||
|
```
|
||||||
|
|
||||||
|
`GenerateResult` uses `Condition` for non-streaming (`wait_completion()`) and `Event` for streaming (`wait()`). Stream callback is `cb(token)`.
|
||||||
|
|
||||||
|
## HTTP API
|
||||||
|
|
||||||
|
```
|
||||||
|
POST /v1/chat/completions OpenAI
|
||||||
|
POST /v1/messages Anthropic
|
||||||
|
GET /health {"status":"ok","model_loaded":true}
|
||||||
|
GET /stats scheduler statistics
|
||||||
|
```
|
||||||
|
|
||||||
|
### OpenAI
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl -X POST http://localhost:8000/v1/chat/completions \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{"messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||||
|
```
|
||||||
|
|
||||||
|
Response:
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"id": "chatcmpl-abc123",
|
||||||
|
"object": "chat.completion",
|
||||||
|
"choices": [{"message": {"role": "assistant", "content": "Hello!"}, "finish_reason": "stop"}],
|
||||||
|
"usage": {"prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Streaming SSE: `data: {"choices":[{"delta":{"role":"assistant"}}]}` → token chunks → `data: [DONE]`
|
||||||
|
|
||||||
|
### Anthropic
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl -X POST http://localhost:8000/v1/messages \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{"model":"astrai","system":"You are helpful.","messages":[{"role":"user","content":"Hello"}],"max_tokens":512}'
|
||||||
|
```
|
||||||
|
|
||||||
|
Supports `stop_sequences` and streaming via `event: content_block_delta`.
|
||||||
|
|
||||||
|
### GenerationRequest Parameters
|
||||||
|
|
||||||
|
| Param | Type | Default | Description |
|
||||||
|
|-------|------|---------|-------------|
|
||||||
|
| `messages` | List[dict] | required | Chat messages (role, content) |
|
||||||
|
| `temperature` | float | 1.0 | Sampling temperature (0.0–2.0) |
|
||||||
|
| `top_p` | float | 1.0 | Nucleus threshold |
|
||||||
|
| `top_k` | int | 50 | Top-k count |
|
||||||
|
| `max_tokens` | int | None | Max generation length |
|
||||||
|
| `stream` | bool | False | Stream output |
|
||||||
|
|
||||||
|
## Engine API
|
||||||
|
|
||||||
|
```python
|
||||||
|
# Non-streaming
|
||||||
|
engine.generate("Hello", stream=False) # -> str
|
||||||
|
engine.generate(["A", "B"], stream=False) # -> List[str]
|
||||||
|
|
||||||
|
# Streaming
|
||||||
|
engine.generate("Hello", stream=True) # -> Generator[str]
|
||||||
|
engine.generate(["A", "B"], stream=True) # -> Generator[Tuple[int, str]]
|
||||||
|
|
||||||
|
# Async
|
||||||
|
await engine.generate_async("Hello", ...) # -> AsyncGenerator[str]
|
||||||
|
```
|
||||||
|
|
||||||
|
> Document Update Time: 2026-05-15
|
||||||
@@ -1,334 +0,0 @@
|
|||||||
## Model Introduction
|
|
||||||
|
|
||||||
### 1. Model Architecture
|
|
||||||
|
|
||||||
This model uses the Transformer architecture with GQA mechanism (q_head=24, kv_head=4), which saves KV cache memory compared to traditional MHA. The model is built by stacking multiple layers of Transformer blocks, with 1.0 billion parameters. Transformer is an autoregressive model that calculates the relationship between all previous tokens to obtain the probability distribution of the next token.
|
|
||||||
|
|
||||||
The model now uses the **AutoModel** base class for flexible loading and saving:
|
|
||||||
|
|
||||||
```python
|
|
||||||
from astrai.model import AutoModel
|
|
||||||
|
|
||||||
# Load model from checkpoint
|
|
||||||
model = AutoModel.from_pretrained("path/to/model")
|
|
||||||
|
|
||||||
# Save model to new directory
|
|
||||||
model.save_pretrained("path/to/save")
|
|
||||||
```
|
|
||||||
|
|
||||||
The Transformer model is registered via `@AutoModel.register('transformer')` decorator, allowing easy extension for new model types.
|
|
||||||
|
|
||||||
```mermaid
|
|
||||||
flowchart TB
|
|
||||||
subgraph Layers["Transformer Layers"]
|
|
||||||
direction TB
|
|
||||||
A[Input Embedding] --> B[Transformer Block\nLayer 1]
|
|
||||||
B --> C[Transformer Block\nLayer ...]
|
|
||||||
C --> D[Transformer Block\nLayer ...]
|
|
||||||
D --> E[RMSNorm]
|
|
||||||
E --> F[Linear]
|
|
||||||
F --> G[SoftMax]
|
|
||||||
end
|
|
||||||
|
|
||||||
subgraph TransformerBlock["Transformer Block"]
|
|
||||||
direction TB
|
|
||||||
H[x] --> I[RMSNorm]
|
|
||||||
I --> J[Linear → Q/K/V]
|
|
||||||
J --> K[Q]
|
|
||||||
J --> L[K]
|
|
||||||
J --> M[V]
|
|
||||||
K --> N[RoPE]
|
|
||||||
L --> O[RoPE]
|
|
||||||
N --> P["Q @ K^T / sqrt(d)"]
|
|
||||||
O --> P
|
|
||||||
P --> Q[Masked SoftMax]
|
|
||||||
Q --> R[S @ V]
|
|
||||||
M --> R
|
|
||||||
R --> S[Linear]
|
|
||||||
S --> T[+]
|
|
||||||
H --> T
|
|
||||||
T --> U[RMSNorm]
|
|
||||||
U --> V["Linear (gate)"]
|
|
||||||
U --> W["Linear (up)"]
|
|
||||||
V --> X[SiLU]
|
|
||||||
X --> Y[×]
|
|
||||||
W --> Y
|
|
||||||
Y --> Z["Linear (down)"]
|
|
||||||
Z --> AA[+]
|
|
||||||
T --> AA
|
|
||||||
AA --> BB[x']
|
|
||||||
end
|
|
||||||
|
|
||||||
classDef main fill:#e6f3ff,stroke:#0066cc;
|
|
||||||
classDef block fill:#fff2e6,stroke:#cc6600;
|
|
||||||
class Layers main;
|
|
||||||
class TransformerBlock block;
|
|
||||||
```
|
|
||||||
|
|
||||||
What is an autoregressive model? After splitting a sentence into tokens, the model predicts the probability distribution of the next token. This means the model calculates the probability of the next possible token and its corresponding probability based on the given context (the sequence of tokens that have already appeared).
|
|
||||||
|
|
||||||
#### 1. Autoregression
|
|
||||||
|
|
||||||
In autoregressive modeling, when a sentence is tokenized into a sequence of tokens, the model learns to predict what comes next. Given a sequence of tokens as input, the model calculates a probability distribution over all possible next tokens. This distribution tells us how likely each potential next token is, given the current context.
|
|
||||||
|
|
||||||
For instance, if the input sequence contains tokens representing a question, the model might predict that certain response tokens have higher probabilities than others. The sampling process then selects one token from this distribution—controlled by parameters like top_k, top_p, and temperature—to serve as the next token in the sequence.
|
|
||||||
|
|
||||||
Once a token is selected, it is appended to the input sequence, and the model repeats this process. The updated sequence is then fed back into the model to predict the next token. This iterative process continues until either a special end-of-sequence token is generated, or the maximum sequence length is reached. These control tokens are essential because without them, the model would continue generating tokens indefinitely, eventually exhausting available memory.
|
|
||||||
|
|
||||||
#### 2. Causal Mask
|
|
||||||
|
|
||||||
Transformers use attention mechanism. The input shape is generally [bsz, seq_len], and the output is [bsz, seq_len, n_dim]. To predict the next token, the model's input and output must be offset by one position. The target predicted by the model must be offset by one position, and during training we also use the offset-by-one method:
|
|
||||||
|
|
||||||
```
|
|
||||||
sequence : [[1, 2, 3, 4, 5, 6]]
|
|
||||||
input_ids: [[1, 2, 3, 4, 5]]
|
|
||||||
target_ids: [[2, 3, 4, 5, 6]]
|
|
||||||
```
|
|
||||||
|
|
||||||
The attention score calculation formula is:
|
|
||||||
|
|
||||||
$$ s_{ij} = softmax(\frac{q_i^Tk_j}{\sqrt{d_k}}) $$
|
|
||||||
$$ s_{ij} := s_{ij} + mask_{ij} $$
|
|
||||||
|
|
||||||
Here, the attention score represents the degree to which the model attends to the similarity between two tokens.
|
|
||||||
|
|
||||||
For decoder-only structure models, to prevent the model from "stealing" information from future positions, a mask needs to be added during attention calculation. We need to apply a mask before attention score calculation. This mask is typically a lower triangular matrix, and for a sequence of length n, its shape is [n, n]. Below is an example of how to create such a causal mask matrix for a sequence of length 5:
|
|
||||||
|
|
||||||
```
|
|
||||||
[[0, -inf, -inf, -inf, -inf],
|
|
||||||
[0, 0, -inf, -inf, -inf],
|
|
||||||
[0, 0, 0, -inf, -inf],
|
|
||||||
[0, 0, 0, 0, -inf],
|
|
||||||
[0, 0, 0, 0, 0]]
|
|
||||||
```
|
|
||||||
|
|
||||||
In this matrix, 0 represents positions that can be attended to, while -inf represents positions that should be masked (i.e., should not be attended to). Because this matrix ensures that after the softmax, the parts of the attention scores where $j > i$ change from `inf` to 0, meaning the model cannot see future information.
|
|
||||||
|
|
||||||
#### 3. Rotary Position Embedding
|
|
||||||
|
|
||||||
Rotary Position Embedding (RoPE) is a position encoding method designed to solve the problem of lacking direct modeling of sequence position information in Transformer models. Unlike traditional position encodings (such as sine and cosine function position encodings), RoPE embeds position information directly into the Query (Q) and Key (K) vectors, allowing the model to more naturally handle relative position relationships in sequences.
|
|
||||||
|
|
||||||
$$ q_i = R_i W_q x_i $$
|
|
||||||
$$ k_j = R_j W_k x_j $$
|
|
||||||
$$ q_i^T k_j = (R_i W_q x_i)^T( R_j W_k x_j) = x_i^T W_q^T R_{i-j} W_k x_j $$
|
|
||||||
|
|
||||||
The $R_{i-j}$ controls the attenuation of attention for different tokens at different relative distances. When the absolute value of $i - j$ is larger, the degree of attenuation is stronger. This approach allows the model to learn relative position relationships, enabling the model to scale and adapt to longer sequences.
|
|
||||||
|
|
||||||
## KV Cache Implementation
|
|
||||||
|
|
||||||
According to the attention calculation formula:
|
|
||||||
|
|
||||||
$$
|
|
||||||
\begin{align*}
|
|
||||||
o_i &= \sum_j s_{ij} v_{j} \newline
|
|
||||||
s_{ij} &= \text{softmax}\left( \frac{q_{i} k_{j}}{\sqrt{d_k}} \right)
|
|
||||||
\end{align*}
|
|
||||||
$$
|
|
||||||
|
|
||||||
Since the model is an autoregressive model, we only need to calculate for the last part of the sequence, meaning the index $i$ is fixed as the last element of the sequence, and we compute $o_{n}$:
|
|
||||||
|
|
||||||
$$
|
|
||||||
\begin{align*}
|
|
||||||
o_n &= \sum_j s_{j}v_{j} \newline
|
|
||||||
s_j &= \text{softmax}\left(\frac{q_n k_{j}}{\sqrt{d_k}} \right)
|
|
||||||
\end{align*}
|
|
||||||
$$
|
|
||||||
|
|
||||||
If we expand the expression:
|
|
||||||
|
|
||||||
$$
|
|
||||||
o_n = \sum_j \text{softmax}\left(\frac{q_n k_{j}}{\sqrt{d_k}}\right)v_{j}
|
|
||||||
$$
|
|
||||||
|
|
||||||
In the above expression, only k and v have length indices, while $q$ does not. Therefore, during the calculation process, the input of $q$ is fixed as the last token from the previous input, while $k$ and $v$ need to be cached for parts of different lengths. Also, when caching, note that position encoding calculation should be performed before KV cache computation, otherwise there will be position encoding calculation errors.
|
|
||||||
|
|
||||||
### 4. AutoModel Loading
|
|
||||||
|
|
||||||
The project now uses the **AutoModel** base class for flexible model loading and saving:
|
|
||||||
|
|
||||||
```python
|
|
||||||
from astrai.model import AutoModel
|
|
||||||
|
|
||||||
# Load model from checkpoint
|
|
||||||
model = AutoModel.from_pretrained("path/to/model")
|
|
||||||
|
|
||||||
# Save model to new directory
|
|
||||||
model.save_pretrained("path/to/save")
|
|
||||||
```
|
|
||||||
|
|
||||||
The Transformer model is registered via `@AutoModel.register('transformer')` decorator, allowing easy extension for new model types. The `from_pretrained` method automatically loads the `config.json` to determine the model type and uses safetensors format for weights.
|
|
||||||
|
|
||||||
### 5. Continuous Batching Inference
|
|
||||||
|
|
||||||
The inference engine supports **continuous batching** for efficient batch processing:
|
|
||||||
|
|
||||||
```python
|
|
||||||
from astrai.inference import InferenceEngine, GenerationRequest
|
|
||||||
|
|
||||||
# Create inference engine with continuous batching
|
|
||||||
engine = InferenceEngine(
|
|
||||||
model=model,
|
|
||||||
tokenizer=tokenizer,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Use GenerationRequest with messages format
|
|
||||||
request = GenerationRequest(
|
|
||||||
messages=[
|
|
||||||
{"role": "system", "content": "You are a helpful assistant."},
|
|
||||||
{"role": "user", "content": "Hello"},
|
|
||||||
],
|
|
||||||
temperature=0.8,
|
|
||||||
top_p=0.95,
|
|
||||||
top_k=50,
|
|
||||||
max_tokens=None,
|
|
||||||
stream=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Generate with streaming
|
|
||||||
for token in engine.generate_with_request(request):
|
|
||||||
print(token, end="", flush=True)
|
|
||||||
```
|
|
||||||
|
|
||||||
The continuous batching feature allows dynamic batch composition where new requests can join at any time and completed requests are released immediately.
|
|
||||||
|
|
||||||
## HTTP API Usage
|
|
||||||
|
|
||||||
The inference server provides HTTP endpoints for remote inference. Start the server first:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python -m scripts.tools.server --port 8000
|
|
||||||
```
|
|
||||||
|
|
||||||
### OpenAI-Compatible Endpoint
|
|
||||||
|
|
||||||
The server provides an OpenAI-compatible chat completion endpoint at `/v1/chat/completions`:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
|
||||||
-H "Content-Type: application/json" \
|
|
||||||
-d '{
|
|
||||||
"messages": [
|
|
||||||
{"role": "system", "content": "You are a helpful assistant."},
|
|
||||||
{"role": "user", "content": "Hello, how are you?"}
|
|
||||||
],
|
|
||||||
"temperature": 0.8,
|
|
||||||
"max_tokens": 2048,
|
|
||||||
"stream": false
|
|
||||||
}'
|
|
||||||
```
|
|
||||||
|
|
||||||
**Request Parameters:**
|
|
||||||
| Parameter | Type | Default | Description |
|
|
||||||
|-----------|------|---------|-------------|
|
|
||||||
| `messages` | List[dict] | Required | Chat messages with role and content |
|
|
||||||
| `temperature` | float | 1.0 | Sampling temperature (0.0-2.0) |
|
|
||||||
| `top_p` | float | 1.0 | Nucleus sampling threshold |
|
|
||||||
| `top_k` | int | 50 | Top-k sampling parameter |
|
|
||||||
| `max_tokens` | int | 1024 | Maximum tokens to generate |
|
|
||||||
| `stream` | bool | false | Enable streaming response |
|
|
||||||
|
|
||||||
**Response (non-streaming):**
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"id": "chatcmpl-1234567890",
|
|
||||||
"object": "chat.completion",
|
|
||||||
"created": 1234567890,
|
|
||||||
"model": "astrai",
|
|
||||||
"choices": [
|
|
||||||
{
|
|
||||||
"index": 0,
|
|
||||||
"message": {"role": "assistant", "content": "Hello! I'm doing well..."},
|
|
||||||
"finish_reason": "stop"
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"usage": {
|
|
||||||
"prompt_tokens": 20,
|
|
||||||
"completion_tokens": 15,
|
|
||||||
"total_tokens": 35
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
### Streaming Response
|
|
||||||
|
|
||||||
Enable streaming for real-time token-by-token output:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
curl -X POST http://localhost:8000/v1/chat/completions \
|
|
||||||
-H "Content-Type: application/json" \
|
|
||||||
-d '{
|
|
||||||
"messages": [{"role": "user", "content": "Write a story"}],
|
|
||||||
"stream": true,
|
|
||||||
"max_tokens": 500
|
|
||||||
}'
|
|
||||||
```
|
|
||||||
|
|
||||||
The server uses Server-Sent Events (SSE) with content type `text/event-stream`.
|
|
||||||
|
|
||||||
### Anthropic-Compatible Endpoint
|
|
||||||
|
|
||||||
The server also provides an Anthropic-compatible endpoint at `/v1/messages`:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
curl -X POST http://localhost:8000/v1/messages \
|
|
||||||
-H "Content-Type: application/json" \
|
|
||||||
-d '{
|
|
||||||
"model": "astrai",
|
|
||||||
"system": "You are a helpful assistant.",
|
|
||||||
"messages": [{"role": "user", "content": "Hello, how are you?"}],
|
|
||||||
"max_tokens": 2048
|
|
||||||
}'
|
|
||||||
```
|
|
||||||
|
|
||||||
Response:
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"id": "msg_abc123...",
|
|
||||||
"type": "message",
|
|
||||||
"role": "assistant",
|
|
||||||
"model": "astrai",
|
|
||||||
"content": [{"type": "text", "text": "Hello! I am doing well..."}],
|
|
||||||
"stop_reason": "end_turn",
|
|
||||||
"stop_sequence": null,
|
|
||||||
"usage": {"input_tokens": 20, "output_tokens": 15}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
Streaming:
|
|
||||||
```bash
|
|
||||||
curl -X POST http://localhost:8000/v1/messages \
|
|
||||||
-H "Content-Type: application/json" \
|
|
||||||
-d '{
|
|
||||||
"model": "astrai",
|
|
||||||
"system": "You are a helpful assistant.",
|
|
||||||
"messages": [{"role": "user", "content": "Write a short poem"}],
|
|
||||||
"max_tokens": 500,
|
|
||||||
"stream": true
|
|
||||||
}'
|
|
||||||
```
|
|
||||||
|
|
||||||
Supports `stop_sequences` for early termination:
|
|
||||||
```bash
|
|
||||||
curl -X POST http://localhost:8000/v1/messages \
|
|
||||||
-H "Content-Type: application/json" \
|
|
||||||
-d '{
|
|
||||||
"model": "astrai",
|
|
||||||
"messages": [{"role": "user", "content": "Write a story"}],
|
|
||||||
"max_tokens": 500,
|
|
||||||
"stop_sequences": ["The end", "THE END"]
|
|
||||||
}'
|
|
||||||
```
|
|
||||||
|
|
||||||
### Health Check
|
|
||||||
|
|
||||||
Monitor server and model status:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
curl http://localhost:8000/health
|
|
||||||
# {"status": "ok", "model_loaded": true}
|
|
||||||
|
|
||||||
curl http://localhost:8000/stats
|
|
||||||
# {"total_tasks": 10, "total_tokens": 5000, "active_tasks": 1, "waiting_queue": 0}
|
|
||||||
```
|
|
||||||
|
|
||||||
> Document Update Time: 2026-05-14
|
|
||||||
@@ -60,7 +60,7 @@
|
|||||||
| Parameter | Description | Default | Used by |
|
| Parameter | Description | Default | Used by |
|
||||||
|-----------|-------------|---------|---------|
|
|-----------|-------------|---------|---------|
|
||||||
| `--dpo_beta` | DPO beta value | 0.1 | `dpo` |
|
| `--dpo_beta` | DPO beta value | 0.1 | `dpo` |
|
||||||
| `--label_smoothing` | Label smoothing for cross-entropy loss | 0.1 | `seq`, `sft` |
|
| `--label_smoothing` | Label smoothing for cross-entropy loss | 0.1 (CLI) / 0.0 (strategy default) | `seq`, `sft` |
|
||||||
| `--group_size` | GRPO group size | 4 | `grpo` |
|
| `--group_size` | GRPO group size | 4 | `grpo` |
|
||||||
| `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo` |
|
| `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo` |
|
||||||
| `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo` |
|
| `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo` |
|
||||||
@@ -98,7 +98,7 @@ python scripts/tools/train.py \
|
|||||||
| `temperature` | Sampling temperature (higher = more random) | 1.0 |
|
| `temperature` | Sampling temperature (higher = more random) | 1.0 |
|
||||||
| `top_p` | Nucleus sampling threshold | 1.0 |
|
| `top_p` | Nucleus sampling threshold | 1.0 |
|
||||||
| `top_k` | Top-k sampling count | 50 |
|
| `top_k` | Top-k sampling count | 50 |
|
||||||
| `max_tokens` | Maximum generation length | None (unlimited) |
|
| `max_tokens` | Maximum generation length | None (defaults to max_seq_len - prompt_len) |
|
||||||
| `stream` | Whether to stream output | False |
|
| `stream` | Whether to stream output | False |
|
||||||
|
|
||||||
### Usage Example
|
### Usage Example
|
||||||
@@ -155,4 +155,4 @@ result = engine.generate(
|
|||||||
| `stream=True` | Streaming output, yields token by token |
|
| `stream=True` | Streaming output, yields token by token |
|
||||||
| `stream=False` | Non-streaming output, returns complete result |
|
| `stream=False` | Non-streaming output, returns complete result |
|
||||||
|
|
||||||
> Document Update Time: 2026-05-14
|
> Document Update Time: 2026-05-15
|
||||||
@@ -0,0 +1,199 @@
|
|||||||
|
# Training
|
||||||
|
|
||||||
|
## Model Architecture
|
||||||
|
|
||||||
|
The model uses a decoder-only Transformer with **GQA** (Grouped Query Attention) and optional **MLA** (Multi-head Latent Attention). 1.0 billion parameters, Chinese–English bilingual.
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
flowchart TB
|
||||||
|
subgraph Layers["Transformer Layers"]
|
||||||
|
direction TB
|
||||||
|
A[Input Embedding] --> B[Transformer Block\nLayer 1]
|
||||||
|
B --> C[Transformer Block\nLayer ...]
|
||||||
|
C --> D[Transformer Block\nLayer ...]
|
||||||
|
D --> E[RMSNorm]
|
||||||
|
E --> F[Linear]
|
||||||
|
F --> G[SoftMax]
|
||||||
|
end
|
||||||
|
|
||||||
|
subgraph TransformerBlock["Transformer Block"]
|
||||||
|
direction TB
|
||||||
|
H[x] --> I[RMSNorm]
|
||||||
|
I --> J[Linear → Q/K/V]
|
||||||
|
J --> K[Q]; J --> L[K]; J --> M[V]
|
||||||
|
K --> N[RoPE]; L --> O[RoPE]
|
||||||
|
N --> P["Q @ K^T / sqrt(d)"]; O --> P
|
||||||
|
P --> Q[Masked SoftMax]; Q --> R[S @ V]; M --> R
|
||||||
|
R --> S[Linear]; S --> T[+]; H --> T
|
||||||
|
T --> U[RMSNorm]
|
||||||
|
U --> V["Linear (gate)"]; U --> W["Linear (up)"]
|
||||||
|
V --> X[SiLU]; X --> Y[×]; W --> Y
|
||||||
|
Y --> Z["Linear (down)"]; Z --> AA[+]; T --> AA
|
||||||
|
AA --> BB[x']
|
||||||
|
end
|
||||||
|
```
|
||||||
|
|
||||||
|
### Autoregression
|
||||||
|
|
||||||
|
Given a token sequence, the model predicts the probability of the next token. Each generated token is appended to the input and fed back, repeating until an end-of-sequence token or max length.
|
||||||
|
|
||||||
|
### Causal Mask
|
||||||
|
|
||||||
|
```
|
||||||
|
sequence : [[1, 2, 3, 4, 5, 6]]
|
||||||
|
input_ids: [[1, 2, 3, 4, 5]]
|
||||||
|
target_ids: [[2, 3, 4, 5, 6]]
|
||||||
|
```
|
||||||
|
|
||||||
|
Lower-triangular mask prevents attending to future positions:
|
||||||
|
|
||||||
|
```
|
||||||
|
[[0, -inf, -inf, -inf, -inf],
|
||||||
|
[0, 0, -inf, -inf, -inf],
|
||||||
|
[0, 0, 0, -inf, -inf],
|
||||||
|
[0, 0, 0, 0, -inf],
|
||||||
|
[0, 0, 0, 0, 0]]
|
||||||
|
```
|
||||||
|
|
||||||
|
### Rotary Position Embedding (RoPE)
|
||||||
|
|
||||||
|
RoPE embeds position into Q/K vectors via complex rotation:
|
||||||
|
|
||||||
|
$$ q_i = R_i W_q x_i, \quad k_j = R_j W_k x_j, \quad q_i^T k_j = x_i^T W_q^T R_{i-j} W_k x_j $$
|
||||||
|
|
||||||
|
The complex rotation `freqs_cis` is pre-computed once (`cos, sin` pairs per position). `apply_rotary_emb` multiplies Q/K as complex numbers.
|
||||||
|
|
||||||
|
## Training Loop
|
||||||
|
|
||||||
|
Nested loop: **epoch** → **step** (accumulation window) → **batch**.
|
||||||
|
|
||||||
|
```
|
||||||
|
on_train_begin
|
||||||
|
on_epoch_begin
|
||||||
|
for steps in batched(dataloader, accumulation_steps):
|
||||||
|
on_step_begin
|
||||||
|
step_batch_nums = len(steps)
|
||||||
|
for batch in steps:
|
||||||
|
on_batch_begin
|
||||||
|
loss = strategy(batch)
|
||||||
|
(loss / step_batch_nums).backward()
|
||||||
|
iteration += 1
|
||||||
|
on_batch_end
|
||||||
|
on_step_end
|
||||||
|
optimizer.step()
|
||||||
|
optimizer.zero_grad()
|
||||||
|
scheduler.step()
|
||||||
|
on_epoch_end
|
||||||
|
on_train_end
|
||||||
|
```
|
||||||
|
|
||||||
|
### Callback Lifecycle
|
||||||
|
|
||||||
|
| Hook | Fires | Default callback |
|
||||||
|
|------|-------|-----------------|
|
||||||
|
| `on_step_end` | Every accumulation window | `GradientClippingCallback` |
|
||||||
|
| `on_batch_end` | Every batch | `CheckpointCallback`, `MetricLoggerCallback`, `ProgressBarCallback` |
|
||||||
|
| `on_train_end` | Training ends | `CheckpointCallback` (final save) |
|
||||||
|
|
||||||
|
Default callbacks: `progress_bar` (tqdm), `checkpoint` (safetensors, rank-0), `metric_logger` (JSONL, rank-0), `gradient_clipping`.
|
||||||
|
|
||||||
|
## Strategies
|
||||||
|
|
||||||
|
### SEQ (Pre-training)
|
||||||
|
|
||||||
|
Next-token cross-entropy with optional label smoothing:
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{\text{PT}} = -\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta)
|
||||||
|
$$
|
||||||
|
|
||||||
|
Keys: `input_ids`, `target_ids`
|
||||||
|
|
||||||
|
### SFT (Supervised Fine-Tuning)
|
||||||
|
|
||||||
|
Masked cross-entropy (`ignore_index=-100`) over response tokens:
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{\text{SFT}} = -\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
|
||||||
|
$$
|
||||||
|
|
||||||
|
Keys: `input_ids`, `target_ids`, `loss_mask`
|
||||||
|
|
||||||
|
### DPO (Direct Preference Optimization)
|
||||||
|
|
||||||
|
Frozen reference model, preference margin via log-ratio:
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{\text{DPO}} = -\mathbb{E}\left[\log\sigma\left(\beta\log\frac{\pi_\theta(y_w\mid x)}{\pi_{\text{ref}}(y_w\mid x)} - \beta\log\frac{\pi_\theta(y_l\mid x)}{\pi_{\text{ref}}(y_l\mid x)}\right)\right]
|
||||||
|
$$
|
||||||
|
|
||||||
|
Parameters: `beta=0.1`. Keys: `chosen`, `rejected`, `chosen_mask`, `rejected_mask`.
|
||||||
|
|
||||||
|
### GRPO (Group Relative Policy Optimization)
|
||||||
|
|
||||||
|
On-policy PPO with group-normalized advantages:
|
||||||
|
|
||||||
|
$$
|
||||||
|
\text{Advantage}_i = \frac{r_i - \mu}{\sigma + \epsilon}
|
||||||
|
$$
|
||||||
|
|
||||||
|
$$
|
||||||
|
L_{\text{GRPO}} = -\mathbb{E}\left[\min\left(\frac{\pi_\theta}{\pi_{\text{ref}}}A,\; \text{clip}\left(\frac{\pi_\theta}{\pi_{\text{ref}}}, 1-\epsilon, 1+\epsilon\right)A\right)\right] + \lambda \cdot \mathbb{E}\left[(\log\pi_\theta - \log\pi_{\text{ref}})^2\right]
|
||||||
|
$$
|
||||||
|
|
||||||
|
Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`, `sync_interval=200`.
|
||||||
|
|
||||||
|
Keys: `prompts`, `responses`, `masks`, `rewards`.
|
||||||
|
|
||||||
|
## LR Schedulers
|
||||||
|
|
||||||
|
| Type | Class | Description |
|
||||||
|
|------|-------|-------------|
|
||||||
|
| Cosine | `CosineScheduler` | Linear warmup → cosine decay to `min_rate` |
|
||||||
|
| SGDR | `SGDRScheduler` | Cosine annealing with warm restarts (`t_mult=2`) |
|
||||||
|
|
||||||
|
Created by `SchedulerFactory.create(optimizer, schedule_type, **kwargs)`.
|
||||||
|
|
||||||
|
## Checkpoint
|
||||||
|
|
||||||
|
```
|
||||||
|
Checkpoint(state_dict, epoch, iteration, extra)
|
||||||
|
├── save(save_dir) rank-0 only: meta.json + state_dict.safetensors + optional extra.pt
|
||||||
|
└── load(save_dir) broadcasts metadata from rank-0
|
||||||
|
```
|
||||||
|
|
||||||
|
Optimizer/scheduler state NOT persisted by default; `Checkpoint.extra` can store arbitrary data.
|
||||||
|
|
||||||
|
## TrainContextBuilder (Builder Pattern)
|
||||||
|
|
||||||
|
```python
|
||||||
|
context = (
|
||||||
|
TrainContextBuilder(config)
|
||||||
|
.with_checkpoint(checkpoint)
|
||||||
|
.build()
|
||||||
|
)
|
||||||
|
# Returns TrainContext with model, strategy, optimizer, scheduler, dataloader, checkpoint
|
||||||
|
```
|
||||||
|
|
||||||
|
- Loads checkpoint weights if provided
|
||||||
|
- Wraps model with `parallel_wrapper` if `nprocs > 1`
|
||||||
|
- Creates `ResumableDistributedSampler` for shuffle+resume
|
||||||
|
- Builds strategy via `StrategyFactory.create(train_type, ...)`
|
||||||
|
|
||||||
|
## Training CLI
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python scripts/tools/train.py \
|
||||||
|
--train_type seq \
|
||||||
|
--data_root_path /path/to/data \
|
||||||
|
--param_path /path/to/model \
|
||||||
|
--batch_size 4 \
|
||||||
|
--accumulation_steps 8 \
|
||||||
|
--max_lr 3e-4 \
|
||||||
|
--warmup_steps 1000 \
|
||||||
|
--n_epoch 1
|
||||||
|
```
|
||||||
|
|
||||||
|
Full parameter reference at [params.md](params.md).
|
||||||
|
|
||||||
|
> Document Update Time: 2026-05-15
|
||||||
@@ -1,12 +1,92 @@
|
|||||||
import json
|
import json
|
||||||
from dataclasses import asdict, dataclass
|
import sys
|
||||||
from typing import Optional, Self
|
from dataclasses import dataclass, fields
|
||||||
|
from typing import Any, Dict, Optional, Self, get_type_hints
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ModelConfig:
|
class BaseModelConfig:
|
||||||
# basic config
|
"""Field-aware JSON load/save for dataclass configs.
|
||||||
|
|
||||||
|
Subclass with additional fields. The base ``model_type`` field
|
||||||
|
enables ``AutoModel`` to pick the correct subclass.
|
||||||
|
"""
|
||||||
|
|
||||||
model_type: Optional[str] = None
|
model_type: Optional[str] = None
|
||||||
|
|
||||||
|
def load(self, config_path: str) -> Self:
|
||||||
|
raw: Dict[str, Any] = {}
|
||||||
|
with open(config_path, "r") as f:
|
||||||
|
raw.update(json.load(f))
|
||||||
|
|
||||||
|
hints = get_type_hints(type(self))
|
||||||
|
valid = {fld.name for fld in fields(self)}
|
||||||
|
for key, value in raw.items():
|
||||||
|
if key not in valid:
|
||||||
|
sys.stderr.write(f"WARNING: unknown config key '{key}'\n")
|
||||||
|
continue
|
||||||
|
|
||||||
|
target_type = self._unwrap_optional(hints.get(key))
|
||||||
|
if target_type is None:
|
||||||
|
continue
|
||||||
|
|
||||||
|
try:
|
||||||
|
value = self._coerce(value, target_type)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
sys.stderr.write(
|
||||||
|
f"WARNING: cannot coerce '{key}' = {value!r} to {target_type}\n"
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
setattr(self, key, value)
|
||||||
|
|
||||||
|
return self
|
||||||
|
|
||||||
|
def save(self, config_path: str):
|
||||||
|
config_dict: Dict[str, Any] = {}
|
||||||
|
for fld in fields(self):
|
||||||
|
v = getattr(self, fld.name)
|
||||||
|
if v is not None:
|
||||||
|
config_dict[fld.name] = v
|
||||||
|
with open(config_path, "w") as f:
|
||||||
|
json.dump(config_dict, f, indent=4)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _unwrap_optional(tp: type) -> Optional[type]:
|
||||||
|
if tp is None:
|
||||||
|
return None
|
||||||
|
origin = getattr(tp, "__origin__", None)
|
||||||
|
if origin is not None:
|
||||||
|
args = getattr(tp, "__args__", ())
|
||||||
|
non_none = [a for a in args if a is not type(None)]
|
||||||
|
return non_none[0] if non_none else None
|
||||||
|
return tp
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _coerce(value: Any, target_type: type) -> Any:
|
||||||
|
if target_type is bool and isinstance(value, bool):
|
||||||
|
return value
|
||||||
|
if (
|
||||||
|
target_type is int
|
||||||
|
and isinstance(value, (int, float))
|
||||||
|
and not isinstance(value, bool)
|
||||||
|
):
|
||||||
|
return int(value)
|
||||||
|
if (
|
||||||
|
target_type is float
|
||||||
|
and isinstance(value, (int, float))
|
||||||
|
and not isinstance(value, bool)
|
||||||
|
):
|
||||||
|
return float(value)
|
||||||
|
if target_type is str and isinstance(value, str):
|
||||||
|
return value
|
||||||
|
if isinstance(value, target_type):
|
||||||
|
return value
|
||||||
|
raise TypeError
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ModelConfig(BaseModelConfig):
|
||||||
vocab_size: Optional[int] = None
|
vocab_size: Optional[int] = None
|
||||||
dim: Optional[int] = None
|
dim: Optional[int] = None
|
||||||
|
|
||||||
@@ -19,24 +99,16 @@ class ModelConfig:
|
|||||||
max_len: Optional[int] = None
|
max_len: Optional[int] = None
|
||||||
rope_theta: Optional[float] = None
|
rope_theta: Optional[float] = None
|
||||||
|
|
||||||
# GQA
|
# attention
|
||||||
|
attn_type: str = "gqa"
|
||||||
n_heads: Optional[int] = None
|
n_heads: Optional[int] = None
|
||||||
n_kv_heads: Optional[int] = None
|
n_kv_heads: Optional[int] = None
|
||||||
use_qk_norm: Optional[bool] = None
|
use_qk_norm: Optional[bool] = None
|
||||||
use_gated_attention: Optional[bool] = None
|
use_gated_attention: Optional[bool] = None
|
||||||
|
|
||||||
def load(self, config_path: str) -> Self:
|
# MoE
|
||||||
config = {}
|
ffn_type: str = "mlp"
|
||||||
with open(config_path, "r") as f:
|
n_routed_experts: Optional[int] = None
|
||||||
config.update(json.load(f))
|
n_shared_experts: Optional[int] = None
|
||||||
|
n_activated_experts: Optional[int] = None
|
||||||
for key, value in config.items():
|
moe_topk_method: Optional[str] = None
|
||||||
if hasattr(self, key):
|
|
||||||
setattr(self, key, value)
|
|
||||||
|
|
||||||
return self
|
|
||||||
|
|
||||||
def save(self, config_path: str):
|
|
||||||
config_dict = {k: v for k, v in asdict(self).items() if v is not None}
|
|
||||||
with open(config_path, "w") as f:
|
|
||||||
json.dump(config_dict, f, indent=4)
|
|
||||||
|
|||||||
@@ -1,11 +1,9 @@
|
|||||||
from astrai.model.automodel import AutoModel
|
from astrai.model.automodel import AutoModel
|
||||||
from astrai.model.module import (
|
from astrai.model.components.attention import GQA
|
||||||
GQA,
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
MLP,
|
from astrai.model.components.linear import Linear
|
||||||
DecoderBlock,
|
from astrai.model.components.mlp import MLP
|
||||||
Linear,
|
from astrai.model.components.norm import RMSNorm
|
||||||
RMSNorm,
|
|
||||||
)
|
|
||||||
from astrai.model.transformer import Transformer
|
from astrai.model.transformer import Transformer
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
|||||||
@@ -0,0 +1,25 @@
|
|||||||
|
from astrai.model.components.attention import GQA, MLA, repeat_kv
|
||||||
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
|
from astrai.model.components.embedding import Embedding
|
||||||
|
from astrai.model.components.linear import Linear
|
||||||
|
from astrai.model.components.mlp import MLP
|
||||||
|
from astrai.model.components.norm import RMSNorm
|
||||||
|
from astrai.model.components.rope import (
|
||||||
|
RotaryEmbedding,
|
||||||
|
apply_rotary_emb,
|
||||||
|
get_rotary_emb,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"Linear",
|
||||||
|
"RMSNorm",
|
||||||
|
"MLP",
|
||||||
|
"Embedding",
|
||||||
|
"GQA",
|
||||||
|
"MLA",
|
||||||
|
"DecoderBlock",
|
||||||
|
"RotaryEmbedding",
|
||||||
|
"apply_rotary_emb",
|
||||||
|
"get_rotary_emb",
|
||||||
|
"repeat_kv",
|
||||||
|
]
|
||||||
@@ -5,11 +5,14 @@ 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.factory import BaseFactory
|
||||||
from astrai.inference.core.cache import KvcacheView
|
from astrai.inference.core.cache import KvcacheView
|
||||||
|
from astrai.model.components.linear import Linear
|
||||||
|
from astrai.model.components.norm import RMSNorm
|
||||||
|
from astrai.model.components.rope import apply_rotary_emb
|
||||||
|
|
||||||
|
|
||||||
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
||||||
"""Repeat KV heads n_rep times for GQA."""
|
|
||||||
bs, slen, n_heads, head_dim = x.shape
|
bs, slen, n_heads, head_dim = x.shape
|
||||||
if n_rep == 1:
|
if n_rep == 1:
|
||||||
return x
|
return x
|
||||||
@@ -20,88 +23,13 @@ def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def get_rotary_emb(
|
class AttnFactory(BaseFactory[nn.Module]):
|
||||||
dim: int,
|
@classmethod
|
||||||
max_len: int,
|
def create(cls, attn_type: str, **kwargs) -> nn.Module:
|
||||||
base: float = 10000,
|
return super().create(attn_type, **kwargs)
|
||||||
device: Optional[torch.device] = None,
|
|
||||||
) -> Tensor:
|
|
||||||
theta = base ** (-torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim)
|
|
||||||
t = torch.arange(0, max_len, dtype=torch.float64, device=device)
|
|
||||||
freqs = torch.outer(t, theta).float()
|
|
||||||
cos = torch.cos(freqs)
|
|
||||||
sin = torch.sin(freqs)
|
|
||||||
return torch.complex(cos, sin)
|
|
||||||
|
|
||||||
|
|
||||||
def apply_rotary_emb(x: torch.Tensor, freqs_cis: Tensor) -> Tensor:
|
|
||||||
dtype = x.dtype
|
|
||||||
x_ = x.float().reshape(*x.shape[:-1], -1, 2)
|
|
||||||
x_complex = torch.view_as_complex(x_)
|
|
||||||
freqs_cis = freqs_cis.unsqueeze(2)
|
|
||||||
x_rotated = x_complex * freqs_cis
|
|
||||||
x_out = torch.view_as_real(x_rotated).flatten(-2)
|
|
||||||
return x_out.to(dtype)
|
|
||||||
|
|
||||||
|
|
||||||
class RotaryEmbedding(nn.Module):
|
|
||||||
def __init__(self, dim: int, max_len: int, base: int = 10000):
|
|
||||||
super().__init__()
|
|
||||||
self.dim = dim
|
|
||||||
self.max_len = max_len
|
|
||||||
self.base = base
|
|
||||||
self._set_rotary_buffer(self.max_len)
|
|
||||||
|
|
||||||
def _set_rotary_buffer(self, max_len: int):
|
|
||||||
rotary_emb = get_rotary_emb(self.dim, max_len, self.base)
|
|
||||||
freqs_cis = torch.view_as_real(rotary_emb)
|
|
||||||
self.register_buffer("freqs_cis", freqs_cis, persistent=False)
|
|
||||||
|
|
||||||
def forward(self, x: Tensor, position_ids: Optional[Tensor] = None) -> Tensor:
|
|
||||||
if position_ids is None:
|
|
||||||
position_ids = (
|
|
||||||
torch.arange(x.size(1), device=x.device)
|
|
||||||
.unsqueeze(0)
|
|
||||||
.expand(x.size(0), -1)
|
|
||||||
)
|
|
||||||
position_freq_cis = self.freqs_cis[position_ids].float()
|
|
||||||
return torch.view_as_complex(position_freq_cis)
|
|
||||||
|
|
||||||
|
|
||||||
class Linear(nn.Module):
|
|
||||||
def __init__(self, in_dim: int, out_dim: int, bias: bool = False):
|
|
||||||
super().__init__()
|
|
||||||
self.weight = nn.Parameter(torch.empty((out_dim, in_dim)))
|
|
||||||
self.bias = nn.Parameter(torch.zeros(out_dim)) if bias else None
|
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
|
||||||
return F.linear(x, self.weight, self.bias)
|
|
||||||
|
|
||||||
|
|
||||||
class RMSNorm(nn.Module):
|
|
||||||
def __init__(self, dim, norm_eps):
|
|
||||||
super().__init__()
|
|
||||||
self.weight = nn.Parameter(torch.ones(dim))
|
|
||||||
self.normalized_shape = (dim,)
|
|
||||||
self.norm_eps = norm_eps
|
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
|
||||||
return F.rms_norm(x, self.normalized_shape, self.weight, self.norm_eps)
|
|
||||||
|
|
||||||
|
|
||||||
class MLP(nn.Module):
|
|
||||||
def __init__(self, dim: int, dim_feed_forward: int):
|
|
||||||
super().__init__()
|
|
||||||
self.up = Linear(dim, dim_feed_forward)
|
|
||||||
self.gate = Linear(dim, dim_feed_forward)
|
|
||||||
self.down = Linear(dim_feed_forward, dim)
|
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
|
||||||
gated = self.up(x) * F.silu(self.gate(x))
|
|
||||||
out = self.down(gated)
|
|
||||||
return out
|
|
||||||
|
|
||||||
|
|
||||||
|
@AttnFactory.register("gqa")
|
||||||
class GQA(nn.Module):
|
class GQA(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -112,6 +40,7 @@ class GQA(nn.Module):
|
|||||||
norm_eps: float,
|
norm_eps: float,
|
||||||
use_gated_attention: bool,
|
use_gated_attention: bool,
|
||||||
layer_id: int,
|
layer_id: int,
|
||||||
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
assert dim % n_heads == 0
|
assert dim % n_heads == 0
|
||||||
@@ -152,7 +81,6 @@ class GQA(nn.Module):
|
|||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
is_causal = attn_mask is None
|
is_causal = attn_mask is None
|
||||||
|
|
||||||
# (bsz, seq_len, dim) -> (bsz, seq_len, n_heads, head_dim)
|
|
||||||
q = self._split_heads(self.q_proj(x), self.n_heads)
|
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)
|
||||||
@@ -167,7 +95,6 @@ class GQA(nn.Module):
|
|||||||
|
|
||||||
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)
|
|
||||||
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)
|
||||||
sdqa_out = (
|
sdqa_out = (
|
||||||
F.scaled_dot_product_attention(q, k, v, attn_mask, is_causal=is_causal)
|
F.scaled_dot_product_attention(q, k, v, attn_mask, is_causal=is_causal)
|
||||||
@@ -183,6 +110,7 @@ class GQA(nn.Module):
|
|||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
@AttnFactory.register("mla")
|
||||||
class MLA(nn.Module):
|
class MLA(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -195,6 +123,7 @@ class MLA(nn.Module):
|
|||||||
norm_eps: float,
|
norm_eps: float,
|
||||||
use_gated_attention: bool,
|
use_gated_attention: bool,
|
||||||
layer_id: int,
|
layer_id: int,
|
||||||
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.dim = dim
|
self.dim = dim
|
||||||
@@ -212,7 +141,6 @@ 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)
|
||||||
|
|
||||||
# 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),
|
||||||
@@ -274,57 +202,3 @@ class MLA(nn.Module):
|
|||||||
|
|
||||||
out = self.o_proj(attn_out)
|
out = self.o_proj(attn_out)
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
class DecoderBlock(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
dim: int,
|
|
||||||
n_heads: int,
|
|
||||||
dim_ffn: int,
|
|
||||||
n_kv_heads: int,
|
|
||||||
norm_eps: int,
|
|
||||||
use_qk_norm: bool,
|
|
||||||
use_gated_attention: bool,
|
|
||||||
layer_id: int,
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.attention = GQA(
|
|
||||||
dim,
|
|
||||||
n_heads,
|
|
||||||
n_kv_heads,
|
|
||||||
use_qk_norm,
|
|
||||||
norm_eps,
|
|
||||||
use_gated_attention,
|
|
||||||
layer_id,
|
|
||||||
)
|
|
||||||
self.input_norm = RMSNorm(dim, norm_eps)
|
|
||||||
self.mlp = MLP(dim, dim_ffn)
|
|
||||||
self.post_attention_norm = RMSNorm(dim, norm_eps)
|
|
||||||
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
x: Tensor,
|
|
||||||
rotary_emb: Tensor,
|
|
||||||
attention_mask: Optional[Tensor] = None,
|
|
||||||
paged_cache: Optional[KvcacheView] = None,
|
|
||||||
) -> Tensor:
|
|
||||||
attn_output = self.attention(
|
|
||||||
self.input_norm(x),
|
|
||||||
rotary_emb,
|
|
||||||
attention_mask,
|
|
||||||
paged_cache,
|
|
||||||
)
|
|
||||||
x = attn_output + x
|
|
||||||
x = self.mlp(self.post_attention_norm(x)) + x
|
|
||||||
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
class Embedding(nn.Module):
|
|
||||||
def __init__(self, vocab_size: int, embedding_dim: int):
|
|
||||||
super().__init__()
|
|
||||||
self.weight = nn.Parameter(torch.empty((vocab_size, embedding_dim)))
|
|
||||||
|
|
||||||
def forward(self, x: Tensor) -> Tensor:
|
|
||||||
return F.embedding(x, self.weight)
|
|
||||||
@@ -0,0 +1,58 @@
|
|||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch.nn as nn
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.inference.core.cache import KvcacheView
|
||||||
|
from astrai.model.components.attention import AttnFactory
|
||||||
|
from astrai.model.components.mlp import FFNFactory
|
||||||
|
from astrai.model.components.norm import RMSNorm
|
||||||
|
|
||||||
|
|
||||||
|
class DecoderBlock(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim: int,
|
||||||
|
n_heads: int,
|
||||||
|
dim_ffn: int,
|
||||||
|
n_kv_heads: int,
|
||||||
|
norm_eps: int,
|
||||||
|
use_qk_norm: bool,
|
||||||
|
use_gated_attention: bool,
|
||||||
|
layer_id: int,
|
||||||
|
attn_type: str = "gqa",
|
||||||
|
ffn_type: str = "mlp",
|
||||||
|
**moe_kwargs,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.attention = AttnFactory.create(
|
||||||
|
attn_type,
|
||||||
|
dim=dim,
|
||||||
|
n_heads=n_heads,
|
||||||
|
n_kv_heads=n_kv_heads,
|
||||||
|
use_qk_norm=use_qk_norm,
|
||||||
|
norm_eps=norm_eps,
|
||||||
|
use_gated_attention=use_gated_attention,
|
||||||
|
layer_id=layer_id,
|
||||||
|
)
|
||||||
|
self.input_norm = RMSNorm(dim, norm_eps)
|
||||||
|
self.post_attention_norm = RMSNorm(dim, norm_eps)
|
||||||
|
self.mlp = FFNFactory.create(ffn_type, dim, dim_ffn, **moe_kwargs)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
x: Tensor,
|
||||||
|
rotary_emb: Tensor,
|
||||||
|
attention_mask: Optional[Tensor] = None,
|
||||||
|
paged_cache: Optional[KvcacheView] = None,
|
||||||
|
) -> Tensor:
|
||||||
|
attn_output = self.attention(
|
||||||
|
self.input_norm(x),
|
||||||
|
rotary_emb,
|
||||||
|
attention_mask,
|
||||||
|
paged_cache,
|
||||||
|
)
|
||||||
|
x = attn_output + x
|
||||||
|
x = self.mlp(self.post_attention_norm(x)) + x
|
||||||
|
|
||||||
|
return x
|
||||||
@@ -0,0 +1,13 @@
|
|||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
class Embedding(nn.Module):
|
||||||
|
def __init__(self, vocab_size: int, embedding_dim: int):
|
||||||
|
super().__init__()
|
||||||
|
self.weight = nn.Parameter(torch.empty((vocab_size, embedding_dim)))
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
return F.embedding(x, self.weight)
|
||||||
@@ -0,0 +1,14 @@
|
|||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
class Linear(nn.Module):
|
||||||
|
def __init__(self, in_dim: int, out_dim: int, bias: bool = False):
|
||||||
|
super().__init__()
|
||||||
|
self.weight = nn.Parameter(torch.empty((out_dim, in_dim)))
|
||||||
|
self.bias = nn.Parameter(torch.zeros(out_dim)) if bias else None
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
return F.linear(x, self.weight, self.bias)
|
||||||
@@ -0,0 +1,94 @@
|
|||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.model.components.linear import Linear
|
||||||
|
|
||||||
|
|
||||||
|
class FFNFactory(BaseFactory[nn.Module]):
|
||||||
|
@classmethod
|
||||||
|
def create(cls, ffn_type: str, dim: int, dim_ffn: int, **kwargs) -> nn.Module:
|
||||||
|
return super().create(ffn_type, dim, dim_ffn, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
@FFNFactory.register("mlp")
|
||||||
|
class MLP(nn.Module):
|
||||||
|
def __init__(self, dim: int, dim_feed_forward: int, **kwargs):
|
||||||
|
super().__init__()
|
||||||
|
self.up = Linear(dim, dim_feed_forward)
|
||||||
|
self.gate = Linear(dim, dim_feed_forward)
|
||||||
|
self.down = Linear(dim_feed_forward, dim)
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
gated = self.up(x) * F.silu(self.gate(x))
|
||||||
|
out = self.down(gated)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
@FFNFactory.register("moe")
|
||||||
|
class DeepSeekMoE(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim: int,
|
||||||
|
dim_feed_forward: int,
|
||||||
|
n_routed_experts: int,
|
||||||
|
n_shared_experts: int = 1,
|
||||||
|
n_activated_experts: int = 2,
|
||||||
|
topk_method: str = "greedy",
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
self.n_routed_experts = n_routed_experts
|
||||||
|
self.n_shared_experts = n_shared_experts
|
||||||
|
self.n_activated_experts = n_activated_experts
|
||||||
|
self.topk_method = topk_method
|
||||||
|
|
||||||
|
self.router = Linear(dim, n_routed_experts, bias=False)
|
||||||
|
|
||||||
|
self.shared_experts = nn.ModuleList(
|
||||||
|
[MLP(dim, dim_feed_forward) for _ in range(n_shared_experts)]
|
||||||
|
)
|
||||||
|
self.routed_experts = nn.ModuleList(
|
||||||
|
[MLP(dim, dim_feed_forward) for _ in range(n_routed_experts)]
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
bsz, seq_len, dim = x.shape
|
||||||
|
x_flat = x.view(-1, dim)
|
||||||
|
|
||||||
|
shared_out = self._shared_forward(x_flat)
|
||||||
|
routed_out = self._routed_forward(x_flat)
|
||||||
|
|
||||||
|
out = (shared_out + routed_out).view(bsz, seq_len, dim)
|
||||||
|
return out
|
||||||
|
|
||||||
|
def _shared_forward(self, x: Tensor) -> Tensor:
|
||||||
|
if self.n_shared_experts == 0:
|
||||||
|
return torch.zeros_like(x)
|
||||||
|
return sum(e(x) for e in self.shared_experts) / self.n_shared_experts
|
||||||
|
|
||||||
|
def _routed_forward(self, x: Tensor) -> Tensor:
|
||||||
|
N, D = x.shape
|
||||||
|
K = self.n_activated_experts
|
||||||
|
|
||||||
|
router_logits = self.router(x)
|
||||||
|
router_probs = torch.softmax(router_logits.float(), dim=-1).to(x.dtype)
|
||||||
|
|
||||||
|
topk_weights, topk_indices = torch.topk(router_probs, K, dim=-1)
|
||||||
|
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
||||||
|
|
||||||
|
output = torch.zeros(N, D, device=x.device, dtype=x.dtype)
|
||||||
|
for expert_idx in range(self.n_routed_experts):
|
||||||
|
expert_mask = topk_indices == expert_idx
|
||||||
|
token_idx, k_idx = expert_mask.nonzero(as_tuple=True)
|
||||||
|
if token_idx.numel() == 0:
|
||||||
|
continue
|
||||||
|
expert_input = x[token_idx]
|
||||||
|
expert_output = self.routed_experts[expert_idx](expert_input)
|
||||||
|
weights = topk_weights[token_idx, k_idx].unsqueeze(-1)
|
||||||
|
output.index_add_(0, token_idx, expert_output * weights)
|
||||||
|
|
||||||
|
return output
|
||||||
@@ -0,0 +1,15 @@
|
|||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
class RMSNorm(nn.Module):
|
||||||
|
def __init__(self, dim, norm_eps):
|
||||||
|
super().__init__()
|
||||||
|
self.weight = nn.Parameter(torch.ones(dim))
|
||||||
|
self.normalized_shape = (dim,)
|
||||||
|
self.norm_eps = norm_eps
|
||||||
|
|
||||||
|
def forward(self, x: Tensor) -> Tensor:
|
||||||
|
return F.rms_norm(x, self.normalized_shape, self.weight, self.norm_eps)
|
||||||
@@ -0,0 +1,53 @@
|
|||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
|
||||||
|
def get_rotary_emb(
|
||||||
|
dim: int,
|
||||||
|
max_len: int,
|
||||||
|
base: float = 10000,
|
||||||
|
device: Optional[torch.device] = None,
|
||||||
|
) -> Tensor:
|
||||||
|
theta = base ** (-torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim)
|
||||||
|
t = torch.arange(0, max_len, dtype=torch.float64, device=device)
|
||||||
|
freqs = torch.outer(t, theta).float()
|
||||||
|
cos = torch.cos(freqs)
|
||||||
|
sin = torch.sin(freqs)
|
||||||
|
return torch.complex(cos, sin)
|
||||||
|
|
||||||
|
|
||||||
|
def apply_rotary_emb(x: torch.Tensor, freqs_cis: Tensor) -> Tensor:
|
||||||
|
dtype = x.dtype
|
||||||
|
x_ = x.float().reshape(*x.shape[:-1], -1, 2)
|
||||||
|
x_complex = torch.view_as_complex(x_)
|
||||||
|
freqs_cis = freqs_cis.unsqueeze(2)
|
||||||
|
x_rotated = x_complex * freqs_cis
|
||||||
|
x_out = torch.view_as_real(x_rotated).flatten(-2)
|
||||||
|
return x_out.to(dtype)
|
||||||
|
|
||||||
|
|
||||||
|
class RotaryEmbedding(nn.Module):
|
||||||
|
def __init__(self, dim: int, max_len: int, base: int = 10000):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
self.max_len = max_len
|
||||||
|
self.base = base
|
||||||
|
self._set_rotary_buffer(self.max_len)
|
||||||
|
|
||||||
|
def _set_rotary_buffer(self, max_len: int):
|
||||||
|
rotary_emb = get_rotary_emb(self.dim, max_len, self.base)
|
||||||
|
freqs_cis = torch.view_as_real(rotary_emb)
|
||||||
|
self.register_buffer("freqs_cis", freqs_cis, persistent=False)
|
||||||
|
|
||||||
|
def forward(self, x: Tensor, position_ids: Optional[Tensor] = None) -> Tensor:
|
||||||
|
if position_ids is None:
|
||||||
|
position_ids = (
|
||||||
|
torch.arange(x.size(1), device=x.device)
|
||||||
|
.unsqueeze(0)
|
||||||
|
.expand(x.size(0), -1)
|
||||||
|
)
|
||||||
|
position_freq_cis = self.freqs_cis[position_ids].float()
|
||||||
|
return torch.view_as_complex(position_freq_cis)
|
||||||
@@ -7,13 +7,11 @@ from torch import Tensor
|
|||||||
from astrai.config.model_config import ModelConfig
|
from astrai.config.model_config import ModelConfig
|
||||||
from astrai.inference.core.cache import KvcacheView
|
from astrai.inference.core.cache import KvcacheView
|
||||||
from astrai.model.automodel import AutoModel
|
from astrai.model.automodel import AutoModel
|
||||||
from astrai.model.module import (
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
DecoderBlock,
|
from astrai.model.components.embedding import Embedding
|
||||||
Embedding,
|
from astrai.model.components.linear import Linear
|
||||||
Linear,
|
from astrai.model.components.norm import RMSNorm
|
||||||
RMSNorm,
|
from astrai.model.components.rope import RotaryEmbedding
|
||||||
RotaryEmbedding,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def process_attention_mask(
|
def process_attention_mask(
|
||||||
@@ -71,6 +69,12 @@ class Transformer(AutoModel):
|
|||||||
config.use_qk_norm,
|
config.use_qk_norm,
|
||||||
config.use_gated_attention,
|
config.use_gated_attention,
|
||||||
layer_id,
|
layer_id,
|
||||||
|
attn_type=config.attn_type,
|
||||||
|
ffn_type=config.ffn_type,
|
||||||
|
n_routed_experts=config.n_routed_experts,
|
||||||
|
n_shared_experts=config.n_shared_experts,
|
||||||
|
n_activated_experts=config.n_activated_experts,
|
||||||
|
topk_method=config.moe_topk_method,
|
||||||
)
|
)
|
||||||
for layer_id in range(config.n_layers)
|
for layer_id in range(config.n_layers)
|
||||||
]
|
]
|
||||||
|
|||||||
Reference in New Issue
Block a user