docs: audit non-CUDA documentation

- Aligns CLI and strategy metric contracts
- Refreshes architecture, dataflow, preprocessing, distributed, and eval guides
- Corrects links, TOCs, defaults, and repository paths
This commit is contained in:
2026-08-02 07:39:24 +08:00
parent 020e2eff4e
commit 288ba20db1
13 changed files with 483 additions and 267 deletions
+9 -7
View File
@@ -20,9 +20,6 @@ Run the following checks **in order** — CI will reject if any fail.
ruff format . ruff format .
``` ```
> **Note**: `ruff format` may rename parameters (e.g. `mask` → `attn_mask`).
> Always review the diff after formatting.
### 2. Import sorting ### 2. Import sorting
```bash ```bash
@@ -44,7 +41,7 @@ python -u -m pytest tests/ -v
> Failed tests may leave orphan tempdirs under `%TEMP%`. Clean them manually if needed. > Failed tests may leave orphan tempdirs under `%TEMP%`. Clean them manually if needed.
### 4. (Optional) Full pre-commit check ### 4. (Optional) Full pre-commit check script
If you have Git Bash available: If you have Git Bash available:
@@ -52,12 +49,17 @@ If you have Git Bash available:
bash scripts/pre_commit.sh bash scripts/pre_commit.sh
``` ```
This runs format check, import sort check, and tests in one go. The script installs development dependencies by default, then runs the format
check, import sort check, and tests. If dependencies are already installed, use:
```bash
bash scripts/pre_commit.sh --skip-deps
```
## Commit Style ## Commit Style
``` ```
fix/feat/chore/docs/refactor/perf/test/style/ci/build/revert : short description (~50 chars) type: short description (~50 chars)
- bullet point body (each ~60 chars) - bullet point body (each ~60 chars)
``` ```
@@ -73,7 +75,7 @@ fix/feat/chore/docs/refactor/perf/test/style/ci/build/revert : short description
|---------|-------|-----| |---------|-------|-----|
| `ruff check --select I` fails | Wrong import order | `ruff check . --select I --fix .` then `ruff format .` | | `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 | | `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` | | Pre-commit check script fails | Dependency install, tests, or lint failed | Fix the failing step; use `--skip-deps` only when dependencies are already installed |
| Tests fail with tempdir left | Test crash | Clean `%TEMP%` manually | | Tests fail with tempdir left | Test crash | Clean `%TEMP%` manually |
## Submitting Changes ## Submitting Changes
+4 -2
View File
@@ -56,6 +56,8 @@ End-to-end walkthrough in 5 steps:
**1. Install** **1. Install**
AstrAI requires Python 3.12+ and pins PyTorch exactly to `2.11.0`. Training, `scripts/tools/generate.py`, generation evaluations, and the generation demos require CUDA; CPU support is limited to components with an explicit CPU device path, such as the HTTP server and direct-scoring evaluations.
```bash ```bash
git clone https://github.com/ViperEkura/AstrAI.git git clone https://github.com/ViperEkura/AstrAI.git
cd AstrAI cd AstrAI
@@ -132,7 +134,7 @@ Check out the demos in the `scripts/demo/` folder:
# Download model weights (required before running demos) # Download model weights (required before running demos)
python scripts/demo/download.py # model → params/ python scripts/demo/download.py # model → params/
# Interactive streaming chat (multi-turn, maintains history) # Single-turn interactive streaming prompt loop (no conversation history)
python scripts/demo/stream_chat.py python scripts/demo/stream_chat.py
# Type your message after >>, type !exit to quit # Type your message after >>, type !exit to quit
@@ -183,7 +185,7 @@ docker run --gpus all -v /path/to/data:/data -it astrai:latest
# Docker Compose (GPU, default) # Docker Compose (GPU, default)
docker compose up -d docker compose up -d
# Docker Compose (CPU only) # Docker Compose CPU server profile (CUDA-only generation scripts/demos are unavailable)
docker compose --profile cpu up -d docker compose --profile cpu up -d
``` ```
+6 -4
View File
@@ -62,6 +62,8 @@
**1. 安装** **1. 安装**
AstrAI 需要 Python 3.12+,并精确固定 PyTorch 版本为 `2.11.0`。训练、`scripts/tools/generate.py`、生成式评估和生成演示需要 CUDA;CPU 支持仅适用于提供明确 CPU 设备路径的组件,例如 HTTP 服务和直接打分评估。
```bash ```bash
git clone https://github.com/ViperEkura/AstrAI.git git clone https://github.com/ViperEkura/AstrAI.git
cd AstrAI cd AstrAI
@@ -138,7 +140,7 @@ curl http://localhost:8000/v1/chat/completions \
# 下载模型权重(运行演示前必需) # 下载模型权重(运行演示前必需)
python scripts/demo/download.py # model → params/ python scripts/demo/download.py # model → params/
# 交互式流式聊天(多轮对话,保持历史记录 # 单轮交互式流式提示循环(不保留对话历史
python scripts/demo/stream_chat.py python scripts/demo/stream_chat.py
# 在 >> 后输入消息,输入 !exit 退出 # 在 >> 后输入消息,输入 !exit 退出
@@ -189,7 +191,7 @@ docker run --gpus all -v /path/to/data:/data -it astrai:latest
# Docker ComposeGPU,默认) # Docker ComposeGPU,默认)
docker compose up -d docker compose up -d
# Docker Compose(仅 CPU # Docker Compose CPU 服务配置(不支持仅限 CUDA 的生成脚本和演示
docker compose --profile cpu up -d docker compose --profile cpu up -d
``` ```
@@ -239,7 +241,7 @@ SSE 流式格式、错误码和统计端点详见[推理文档](guides/inference
### 贡献 ### 贡献
我们欢迎贡献!请参阅[贡献指南](../../CONTRIBUTING.md)了解详情。 我们欢迎贡献!请参阅[贡献指南](../CONTRIBUTING.md)了解详情。
1. Fork 本仓库。 1. Fork 本仓库。
2. 创建功能分支。 2. 创建功能分支。
@@ -256,7 +258,7 @@ SSE 流式格式、错误码和统计端点详见[推理文档](guides/inference
### 许可证 ### 许可证
本项目采用 [GPL-3.0 许可证](../../LICENSE)。 本项目采用 [GPL-3.0 许可证](../LICENSE)。
--- ---
+110 -104
View File
@@ -4,7 +4,7 @@
- [Class Diagram](#class-diagram) — Full Mermaid class diagram across 10+ namespaces - [Class Diagram](#class-diagram) — Full Mermaid class diagram across 10+ namespaces
- [Module Overview](#module-overview) — Component inventory per module - [Module Overview](#module-overview) — Component inventory per module
- [Design Patterns](#design-patterns) — 13 documented patterns with classes - [Design Patterns](#design-patterns) — 15 documented patterns with classes
- [Core Relationships](#core-relationships) — 11 key inter-component relationships - [Core Relationships](#core-relationships) — 11 key inter-component relationships
## Class Diagram ## Class Diagram
@@ -49,6 +49,11 @@ classDiagram
+Optional[int] n_shared_experts +Optional[int] n_shared_experts
+Optional[int] n_activated_experts +Optional[int] n_activated_experts
+Optional[str] topk_method +Optional[str] topk_method
+Optional[int] moe_intermediate_size
+Optional[int] shared_expert_intermediate_size
+bool norm_topk_prob
+int decoder_sparse_step
+Optional[List[int]] mlp_only_layers
} }
class EncoderConfig { class EncoderConfig {
@@ -63,6 +68,7 @@ classDiagram
+Optional[int] num_attention_heads +Optional[int] num_attention_heads
+Optional[int] num_key_value_heads +Optional[int] num_key_value_heads
+Optional[bool] use_qk_norm +Optional[bool] use_qk_norm
+Optional[bool] use_gated_attention
+str ffn_type +str ffn_type
+Optional[dict] rope_scaling +Optional[dict] rope_scaling
+Optional[str] pooling_type +Optional[str] pooling_type
@@ -114,22 +120,25 @@ classDiagram
+Dataset dataset +Dataset dataset
+Callable optimizer_fn +Callable optimizer_fn
+Callable scheduler_fn +Callable scheduler_fn
+Optional[str] optimizer_name
+Dict[str, Any] optimizer_hyperparameters
+int n_epoch +int n_epoch
+int batch_per_device +int batch_per_device
+int grad_accum_steps +int grad_accum_steps
+Optional[float] max_grad_norm +Optional[float] max_grad_norm
+list gradient_checkpointing_modules +list gradient_checkpointing_modules
+Optional[str] compile_mode
+int start_epoch +int start_epoch
+int start_samples +int start_samples
+str ckpt_dir +str ckpt_dir
+int ckpt_interval +int ckpt_interval
+str log_dir
+List[str] metrics +List[str] metrics
+Optional[LoRAConfig] lora +Optional[LoRAConfig] lora
+int random_seed +int random_seed
+int num_workers +int num_workers
+Optional[int] prefetch_factor +Optional[int] prefetch_factor
+bool pin_memory +bool pin_memory
+Optional[Callable] collate_fn
+int nprocs +int nprocs
+str backend +str backend
+str master_addr +str master_addr
@@ -140,6 +149,7 @@ classDiagram
+Optional[float] val_split +Optional[float] val_split
+int val_step +int val_step
+float neftune_alpha +float neftune_alpha
+float moe_aux_loss_coef
+str parallel_mode +str parallel_mode
+int rollout_interval +int rollout_interval
+float rollout_temperature +float rollout_temperature
@@ -149,7 +159,6 @@ classDiagram
+Optional[Callable] reward_model_fn +Optional[Callable] reward_model_fn
+dict executor_kwargs +dict executor_kwargs
+dict extra_kwargs +dict extra_kwargs
+validate()
} }
} }
@@ -205,10 +214,6 @@ classDiagram
-_fetch_record_key(key, index) Tensor -_fetch_record_key(key, index) Tensor
} }
class H5Store {
+load(path)
}
class MmapStore { class MmapStore {
+List _mmap_refs +List _mmap_refs
+load(path) +load(path)
@@ -260,11 +265,15 @@ classDiagram
} }
namespace model { namespace model {
class AutoModel { class ModelFactory {
+BaseModelConfig config
+Dict _entries +Dict _entries
+register(name) decorator +register(name) decorator
+get_component_class(name) Type +get_component_class(name) Type
}
class AutoModel {
<<nn.Module>>
+BaseModelConfig config
+from_pretrained(path, disable_random_init, strict) nn.Module +from_pretrained(path, disable_random_init, strict) nn.Module
+save_pretrained(save_directory) +save_pretrained(save_directory)
+to(*args, **kwargs) Self +to(*args, **kwargs) Self
@@ -299,7 +308,13 @@ classDiagram
+RMSNorm input_norm +RMSNorm input_norm
+nn.Module mlp # MLP or DeepSeekMoE via FFNFactory +nn.Module mlp # MLP or DeepSeekMoE via FFNFactory
+RMSNorm post_attention_norm +RMSNorm post_attention_norm
+forward(x, rotary_emb, attention_mask, kv_cache) Tensor +forward(x, rotary_emb, attention_mask, kv_cache, is_causal) DecoderOutput
}
class DecoderOutput {
<<TypedDict>>
+Tensor hidden_states
+Optional[Tensor] aux_loss
} }
class GQA { class GQA {
@@ -314,7 +329,7 @@ classDiagram
+Linear q_proj, k_proj, v_proj, o_proj +Linear q_proj, k_proj, v_proj, o_proj
+Linear gate # only if use_gated_attention +Linear gate # only if use_gated_attention
+RMSNorm q_norm, k_norm # only if use_qk_norm +RMSNorm q_norm, k_norm # only if use_qk_norm
+forward(x, rotary_emb, attn_mask, kv_cache) Tensor +forward(x, rotary_emb, attn_mask, kv_cache, is_causal) Tensor
} }
class MLA { class MLA {
@@ -334,12 +349,18 @@ classDiagram
+Linear gate # only if use_gated_attention +Linear gate # only if use_gated_attention
+RMSNorm kv_norm +RMSNorm kv_norm
+RMSNorm q_norm, k_norm # only if use_qk_norm +RMSNorm q_norm, k_norm # only if use_qk_norm
+forward(x, rotary_emb, attn_mask, kv_cache) Tensor +forward(x, rotary_emb, attn_mask, kv_cache, is_causal) Tensor
} }
class MLP { class MLP {
+Linear up, gate, down +Linear up, gate, down
+forward(x) Tensor +forward(x) FFNOutput
}
class FFNOutput {
<<TypedDict>>
+Tensor hidden_states
+Optional[Tensor] aux_loss
} }
class DeepSeekMoE { class DeepSeekMoE {
@@ -351,7 +372,7 @@ classDiagram
+Linear router +Linear router
+ModuleList shared_experts +ModuleList shared_experts
+ModuleList routed_experts +ModuleList routed_experts
+forward(x) Tensor +forward(x) FFNOutput
} }
class AttnFactory { class AttnFactory {
@@ -380,9 +401,8 @@ classDiagram
+int max_len +int max_len
+float base +float base
+Optional[Dict] rope_scaling +Optional[Dict] rope_scaling
+Tensor cos_table +Tensor freqs_cis
+Tensor sin_table +forward(x, position_ids=None) Tensor
+forward(x, position_ids=None) Tuple[Tensor, Tensor]
} }
class Embedding { class Embedding {
@@ -486,10 +506,6 @@ classDiagram
+save(output_dir, domain, shard_idx, tensors) +save(output_dir, domain, shard_idx, tensors)
} }
class H5Writer {
+save(output_dir, domain, shard_idx, tensors)
}
class Pipeline { class Pipeline {
+PipelineConfig config +PipelineConfig config
+List[str] paths +List[str] paths
@@ -559,7 +575,7 @@ classDiagram
class Trainer { class Trainer {
+TrainConfig train_config +TrainConfig train_config
+List[TrainCallback] callbacks +List[TrainCallback] callbacks
+train(resume_dir) +train(param_path=None, resume=False)
-_get_default_callbacks() List[TrainCallback] -_get_default_callbacks() List[TrainCallback]
} }
@@ -576,13 +592,17 @@ classDiagram
+int epoch +int epoch
+int consumed_samples +int consumed_samples
+float loss +float loss
+float grad_norm +Dict[str, float] metrics
+Optional[float] grad_norm
+GradSNRTracker grad_snr_tracker
+DataLoader val_dataloader +DataLoader val_dataloader
+float val_loss +Optional[float] val_loss
+int world_size +int world_size
+int rank +int rank
+dict kwargs +dict kwargs
+optimizer_step() int +stop_requested (property) bool
+optimizer_step (property) int
+request_stop()
} }
class TrainContextBuilder { class TrainContextBuilder {
@@ -594,11 +614,22 @@ classDiagram
class BaseStrategy { class BaseStrategy {
+Callable model +Callable model
+Optional[BaseExecutor] executor +Optional[BaseExecutor] executor
+Optional[Callable] model_fn +float moe_aux_loss_coef
+dict extra_kwargs +dict extra_kwargs
+str device +str device
+__call__(batch) Tensor +__call__(batch) LossOutput
+compute_loss(batch) Tensor +compute_loss(batch) Tensor
+compute_loss_output(batch) LossOutput
+supports_online() bool
+set_rollout_runner(runner)
+prepare_from_rollout(result) Dict
+on_optimizer_step()
}
class LossOutput {
<<TypedDict>>
+Tensor loss
+Dict[str, float] metrics
} }
class StrategyFactory { class StrategyFactory {
@@ -636,9 +667,12 @@ classDiagram
class RawRollout { class RawRollout {
+Tensor prompts +Tensor prompts
+Tensor prompt_mask
+Tensor responses +Tensor responses
+Tensor response_mask +Tensor response_mask
+Tensor logprobs_old +Tensor logprobs_old
+List[str] prompt_texts
+List[List[str]] response_texts
} }
class RolloutResult { class RolloutResult {
@@ -647,10 +681,18 @@ classDiagram
class BaseRewardModel { class BaseRewardModel {
<<abstract>> <<abstract>>
+score(prompts, responses) Tensor +score(List[str] prompts, List[List[str]] responses) Tensor
} }
class RolloutGenerator { class RolloutGenerator {
+InferenceScheduler scheduler
+int max_tokens
+int group_size
+float temperature
+int top_k
+float top_p
+float frequency_penalty
+int rep_window
+generate(batch) RawRollout +generate(batch) RawRollout
} }
@@ -740,7 +782,7 @@ classDiagram
} }
class MetricCallback { class MetricCallback {
+Path log_dir +Path ckpt_dir
+int save_interval +int save_interval
+List[str] metrics +List[str] metrics
+int val_step +int val_step
@@ -764,9 +806,9 @@ classDiagram
+nn.Module model +nn.Module model
+AutoTokenizer tokenizer +AutoTokenizer tokenizer
+InferenceScheduler scheduler +InferenceScheduler scheduler
+generate(prompt, stream, max_tokens, temperature, top_p, top_k) Union[Generator, str, List[str]] +generate(prompt, stream, max_tokens, temperature, top_p, top_k, frequency_penalty, rep_window) Union[Generator, str, List[str]]
+generate_with_request(request) Union[Generator, str, List[str]] +generate_with_request(request) Union[Generator, str, List[str]]
+generate_async(prompt, max_tokens, temperature, top_p, top_k) AsyncGenerator +generate_async(prompt, max_tokens, temperature, top_p, top_k, frequency_penalty, rep_window) AsyncGenerator
+get_stats() Dict +get_stats() Dict
+shutdown() +shutdown()
} }
@@ -774,18 +816,18 @@ classDiagram
class Executor { class Executor {
+AutoModel model +AutoModel model
+AutoTokenizer tokenizer +AutoTokenizer tokenizer
+KVCache page_cache +PagePool kv_cache
+Optional[str] device +Optional[str] device
+Optional[torch.dtype] dtype +Optional[torch.dtype] dtype
+execute_prefill(tasks, prompt_len, start_pos) +execute_prefill(tasks, prompt_len, start_pos)
+execute_decode(tasks) List[int] +execute_decode(tasks, return_logprobs=False) Union[List[int], List[Tuple[int, float]]]
} }
class InferenceScheduler { class InferenceScheduler {
+KVCache _page_cache +PagePool _cache
+Executor _executor +Executor _executor
+TaskManager _task_mgr +TaskManager _task_mgr
+bool _running +Event _stop_event
+Thread _loop_thread +Thread _loop_thread
+int max_seq_len +int max_seq_len
+str device +str device
@@ -795,6 +837,7 @@ classDiagram
+start() +start()
+stop() +stop()
+get_stats() Dict +get_stats() Dict
+run_batch(prompt_ids_list, max_tokens, temperature, top_p, top_k, frequency_penalty, rep_window, return_logprobs) Union[List[List[int]], List[Tuple[List[int], List[float]]]]
} }
class Allocator { class Allocator {
@@ -816,16 +859,6 @@ classDiagram
+record(page_idx, token_ids, logical_page_idx) +record(page_idx, token_ids, logical_page_idx)
} }
class PagePool {
-Allocator _alloc
-PrefixCache _prefix
+alloc() int
+free(idx)
+inc_ref(idx)
+lookup(token_ids) List[int]
+record(page_idx, token_ids, logical_page_idx)
}
class KVStorage { class KVStorage {
+int size +int size
+Tensor k_buffer +Tensor k_buffer
@@ -852,8 +885,7 @@ classDiagram
+Tensor seq_lens +Tensor seq_lens
+Tensor out_cache_loc +Tensor out_cache_loc
+int max_len +int max_len
+Optional[Tensor] page_table +Optional[Tensor] kv_indptr
+Optional[Tensor] decode_mask
} }
class PagePool { class PagePool {
@@ -878,6 +910,8 @@ classDiagram
+float temperature +float temperature
+float top_p +float top_p
+int top_k +int top_k
+float frequency_penalty
+int rep_window
+TaskStatus status +TaskStatus status
+List output_ids +List output_ids
+int input_tokens +int input_tokens
@@ -923,27 +957,29 @@ classDiagram
+float top_p +float top_p
+float temperature +float temperature
+Optional[int] max_tokens +Optional[int] max_tokens
+float frequency_penalty
+int rep_window
+bool stream +bool stream
} }
class BaseSamplingStrategy { class BaseSamplingStrategy {
<<abstract>> <<abstract>>
+apply(logits, filter_value) Tensor +apply(logits, filter_value, input_ids, input_mask) Tensor
} }
class TemperatureStrategy { class TemperatureStrategy {
+float temperature +float temperature
+apply(logits, filter_value) Tensor +apply(logits, filter_value, input_ids, input_mask) Tensor
} }
class TopKStrategy { class TopKStrategy {
+int top_k +int top_k
+apply(logits, filter_value) Tensor +apply(logits, filter_value, input_ids, input_mask) Tensor
} }
class TopPStrategy { class TopPStrategy {
+float top_p +float top_p
+apply(logits, filter_value) Tensor +apply(logits, filter_value, input_ids, input_mask) Tensor
} }
class FrequencyPenaltyStrategy { class FrequencyPenaltyStrategy {
@@ -953,8 +989,8 @@ classDiagram
class SamplingPipeline { class SamplingPipeline {
+List[BaseSamplingStrategy] strategies +List[BaseSamplingStrategy] strategies
+apply(logits, filter_value) Tensor +apply(logits, filter_value, input_ids, input_mask) Tensor
+sample(logits, filter_value) Tensor +sample(logits, filter_value, input_ids, input_mask, return_logprobs) Union[Tensor, Tuple[Tensor, Tensor]]
} }
class StreamDecoder { class StreamDecoder {
@@ -1029,7 +1065,7 @@ classDiagram
<<abstract>> <<abstract>>
+prepare(request, engine) Tuple[str, GenContext, List[str]] +prepare(request, engine) Tuple[str, GenContext, List[str]]
+format_stream_start(ctx) List[str] +format_stream_start(ctx) List[str]
+format_chunk(token) List[str] +format_chunk(token, **kwargs) List[str]
+format_stream_end(ctx, stop) List[str] +format_stream_end(ctx, stop) List[str]
+format_response(ctx, content, stop) Dict +format_response(ctx, content, stop) Dict
} }
@@ -1037,7 +1073,7 @@ classDiagram
class OpenAIResponseBuilder { class OpenAIResponseBuilder {
+prepare(request, engine) Tuple +prepare(request, engine) Tuple
+format_stream_start(ctx) List[str] +format_stream_start(ctx) List[str]
+format_chunk(token) List[str] +format_chunk(token, **kwargs) List[str]
+format_stream_end(ctx, stop) List[str] +format_stream_end(ctx, stop) List[str]
+format_response(ctx, content, stop) Dict +format_response(ctx, content, stop) Dict
} }
@@ -1045,7 +1081,7 @@ classDiagram
class AnthropicResponseBuilder { class AnthropicResponseBuilder {
+prepare(request, engine) Tuple +prepare(request, engine) Tuple
+format_stream_start(ctx) List[str] +format_stream_start(ctx) List[str]
+format_chunk(token) List[str] +format_chunk(token, **kwargs) List[str]
+format_stream_end(ctx, stop) List[str] +format_stream_end(ctx, stop) List[str]
+format_response(ctx, content, stop) Dict +format_response(ctx, content, stop) Dict
} }
@@ -1153,10 +1189,13 @@ classDiagram
class BaseExecutor { class BaseExecutor {
+GradientState gradient_state +GradientState gradient_state
+prepare(model_fn, optimizer_fn, scheduler_fn, before_wrap) tuple +prepare(model_fn, optimizer_fn, scheduler_fn, before_wrap, after_wrap) tuple
+accumulate(model) context manager +accumulate(model) context manager
+backward(loss) +backward(loss)
+unwrap_model(model) dict +unwrap_model(model) dict
+checkpoint_context(model) context manager
+clip_grad_norm(model, max_norm) float
+use_distributed (property) bool
+sync_gradients (property) bool +sync_gradients (property) bool
+grad_accum_steps (property) int +grad_accum_steps (property) int
} }
@@ -1173,7 +1212,8 @@ classDiagram
class FSDPExecutor { class FSDPExecutor {
-_prepare_model(model) nn.Module -_prepare_model(model) nn.Module
-_no_sync(model) context manager -_no_sync(model) context manager
+unwrap_model(model) dict +unwrap_model(model) Optional[dict]
+clip_grad_norm(model, max_norm) float
} }
class ExecutorFactory { class ExecutorFactory {
@@ -1182,33 +1222,6 @@ classDiagram
+create(parallel_mode, **kwargs) BaseExecutor +create(parallel_mode, **kwargs) BaseExecutor
} }
class ParallelModel {
+dist.ProcessGroup process_group
+int rank
+int world_size
}
class ColumnParallelLinear {
+int in_features
+int out_features
+int out_features_per_rank
+bool gather_results
+Parameter weight
+Optional[Parameter] bias
+forward(x) Tensor
+load_state_dict(state_dict)
}
class RowParallelLinear {
+int in_features
+int out_features
+int in_features_per_rank
+bool reduce_results
+Parameter weight
+Optional[Parameter] bias
+forward(x) Tensor
+load_state_dict(state_dict)
}
} }
%% Relationships — UML notation: <|-- generalization, *-- composition, o-- aggregation, --> association, ..> dependency %% Relationships — UML notation: <|-- generalization, *-- composition, o-- aggregation, --> association, ..> dependency
@@ -1230,11 +1243,8 @@ classDiagram
BaseDataset <|-- SFTDataset BaseDataset <|-- SFTDataset
BaseDataset <|-- DPODataset BaseDataset <|-- DPODataset
BaseDataset <|-- GRPODataset BaseDataset <|-- GRPODataset
Store <|-- H5Store
Store <|-- MmapStore Store <|-- MmapStore
Store <|-- JsonlStore Store <|-- JsonlStore
H5Store --|> Streamable
H5Store --|> Recordable
MmapStore --|> Streamable MmapStore --|> Streamable
MmapStore --|> Recordable MmapStore --|> Recordable
JsonlStore --|> Streamable JsonlStore --|> Streamable
@@ -1243,8 +1253,6 @@ classDiagram
BaseSamplingStrategy <|-- TopKStrategy BaseSamplingStrategy <|-- TopKStrategy
BaseSamplingStrategy <|-- TopPStrategy BaseSamplingStrategy <|-- TopPStrategy
BaseSamplingStrategy <|-- FrequencyPenaltyStrategy BaseSamplingStrategy <|-- FrequencyPenaltyStrategy
ParallelModel <|-- RowParallelLinear
ParallelModel <|-- ColumnParallelLinear
AutoModel <|-- AutoRegressiveLM AutoModel <|-- AutoRegressiveLM
AutoModel <|-- EmbeddingEncoder AutoModel <|-- EmbeddingEncoder
BaseConfig <|-- BaseModelConfig BaseConfig <|-- BaseModelConfig
@@ -1255,7 +1263,7 @@ classDiagram
BaseConfig <|-- PipelineConfig BaseConfig <|-- PipelineConfig
BaseModelConfig <|-- AutoRegressiveLMConfig BaseModelConfig <|-- AutoRegressiveLMConfig
BaseModelConfig <|-- EncoderConfig BaseModelConfig <|-- EncoderConfig
BaseFactory <|-- AutoModel BaseFactory <|-- ModelFactory
BaseFactory <|-- AttnFactory BaseFactory <|-- AttnFactory
BaseFactory <|-- FFNFactory BaseFactory <|-- FFNFactory
BaseFactory <|-- DatasetFactory BaseFactory <|-- DatasetFactory
@@ -1286,7 +1294,6 @@ classDiagram
PositionIdStrategy <|-- DocResetPositionId PositionIdStrategy <|-- DocResetPositionId
PositionIdStrategy <|-- ContinuousPositionId PositionIdStrategy <|-- ContinuousPositionId
StoreWriter <|-- BinWriter StoreWriter <|-- BinWriter
StoreWriter <|-- H5Writer
RawRollout <|-- RolloutResult RawRollout <|-- RolloutResult
LaunchStrategy <|-- TorchrunStrategy LaunchStrategy <|-- TorchrunStrategy
LaunchStrategy <|-- LocalStrategy LaunchStrategy <|-- LocalStrategy
@@ -1317,8 +1324,6 @@ classDiagram
%% --- Aggregation (weak ownership) --- %% --- Aggregation (weak ownership) ---
AutoModel o-- BaseModelConfig AutoModel o-- BaseModelConfig
AutoTokenizer o-- ChatTemplate AutoTokenizer o-- ChatTemplate
PagePool o-- Allocator
PagePool o-- PrefixCache
Trainer o-- TrainCallback Trainer o-- TrainCallback
TrainContext o-- BaseStrategy TrainContext o-- BaseStrategy
TrainContext o-- BaseScheduler TrainContext o-- BaseScheduler
@@ -1352,11 +1357,12 @@ classDiagram
FFNFactory ..> DeepSeekMoE : creates FFNFactory ..> DeepSeekMoE : creates
DecoderBlock ..> AttnFactory : uses DecoderBlock ..> AttnFactory : uses
DecoderBlock ..> FFNFactory : uses DecoderBlock ..> FFNFactory : uses
StoreFactory ..> H5Store : creates
StoreFactory ..> MmapStore : creates StoreFactory ..> MmapStore : creates
StoreFactory ..> JsonlStore : creates StoreFactory ..> JsonlStore : creates
ConfigFactory ..> AutoRegressiveLMConfig : creates ConfigFactory ..> AutoRegressiveLMConfig : creates
ConfigFactory ..> EncoderConfig : creates ConfigFactory ..> EncoderConfig : creates
ModelFactory ..> AutoRegressiveLM : creates
ModelFactory ..> EmbeddingEncoder : creates
ExecutorFactory ..> NoneExecutor : creates ExecutorFactory ..> NoneExecutor : creates
ExecutorFactory ..> DDPExecutor : creates ExecutorFactory ..> DDPExecutor : creates
ExecutorFactory ..> FSDPExecutor : creates ExecutorFactory ..> FSDPExecutor : creates
@@ -1399,10 +1405,10 @@ classDiagram
| Module | Components | Description | | Module | Components | Description |
|--------|------------|-------------| |--------|------------|-------------|
| **astrai.config** | BaseConfig, BaseModelConfig, AutoRegressiveLMConfig, EncoderConfig, ConfigFactory, TrainConfig, PipelineConfig, InputConfig, ProcessingConfig, OutputConfig | Configuration management (to_dict/from_dict, to_file/from_file) | | **astrai.config** | BaseConfig, BaseModelConfig, AutoRegressiveLMConfig, EncoderConfig, ConfigFactory, TrainConfig, PipelineConfig, InputConfig, ProcessingConfig, OutputConfig | Configuration management (to_dict/from_dict, to_file/from_file) |
| **astrai.preprocessing** | SectionRenderer, BaseMaskBuilder, MaskBuilderFactory, SectionedMaskBuilder, SingleOutputMaskBuilder, MultiOutputMaskBuilder, Pipeline, TokenizeTransform, PackingStrategy, PackingStrategyFactory, SimplePacking, BFDPacking, BFDSplitPacking, PositionIdStrategy, PositionIdStrategyFactory, NoPositionId, DocResetPositionId, ContinuousPositionId, StoreWriter, StoreWriterFactory, BinWriter, H5Writer | Declarative JSON-driven data preprocessing | | **astrai.preprocessing** | SectionRenderer, BaseMaskBuilder, MaskBuilderFactory, SectionedMaskBuilder, SingleOutputMaskBuilder, MultiOutputMaskBuilder, Pipeline, TokenizeTransform, PackingStrategy, PackingStrategyFactory, SimplePacking, BFDPacking, BFDSplitPacking, PositionIdStrategy, PositionIdStrategyFactory, NoPositionId, DocResetPositionId, ContinuousPositionId, StoreWriter, StoreWriterFactory, BinWriter | Declarative JSON-driven data preprocessing |
| **astrai.dataset** | BaseDataset, SEQDataset, SFTDataset, DPODataset, GRPODataset, Store, Streamable, Recordable, H5Store, MmapStore, JsonlSource, JsonlStore, StoreFactory, RDSampler, DatasetFactory | Dataset loading and management | | **astrai.dataset** | BaseDataset, SEQDataset, SFTDataset, DPODataset, GRPODataset, Store, Streamable, Recordable, MmapStore, JsonlSource, JsonlStore, StoreFactory, RDSampler, DatasetFactory | Dataset loading and management |
| **astrai.serialization** | Checkpoint | Model serialization | | **astrai.serialization** | Checkpoint | Model serialization |
| **astrai.model** | AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, LoRAConfig, LoRALinear, RotaryEmbedding, Embedding | Neural network model | | **astrai.model** | ModelFactory, AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, LoRAConfig, LoRALinear, 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, BaseStrategyGRPOStrategy, StrategyFactory, BaseSchedulerWSDScheduler, SchedulerFactory, TrainCallback(Protocol)MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow | | **astrai.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategyGRPOStrategy, StrategyFactory, BaseSchedulerWSDScheduler, SchedulerFactory, TrainCallback(Protocol)MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow |
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, PagePool, KVStorage, ReqToTokenPool, KVCache, Allocator, PrefixCache, Task, TaskManager, TaskStatus, StreamDecoder, GenerationRequest, GenerateResult, BaseSamplingStrategySamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service | | **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, PagePool, KVStorage, ReqToTokenPool, KVCache, Allocator, PrefixCache, Task, TaskManager, TaskStatus, StreamDecoder, GenerationRequest, GenerateResult, BaseSamplingStrategySamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service |
@@ -1415,7 +1421,7 @@ classDiagram
| Pattern | Classes | Purpose | | Pattern | Classes | Purpose |
|---------|---------|---------| |---------|---------|---------|
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory`, `MaskBuilderFactory`, `StoreWriterFactory`, `PackingStrategyFactory`, `PositionIdStrategyFactory`, `ToolParserFactory` | Decorator-based component creation | | **Factory** | `ModelFactory`, `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory`, `MaskBuilderFactory`, `StoreWriterFactory`, `PackingStrategyFactory`, `PositionIdStrategyFactory`, `ToolParserFactory` | Decorator-based component creation |
| **Registry** | `BaseFactory` | Component registration | | **Registry** | `BaseFactory` | Component registration |
| **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching | | **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching |
| **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations | | **Strategy (Sampling)** | `TemperatureStrategy`, `TopKStrategy`, `TopPStrategy`, `SamplingPipeline` | Composable logit transformations |
@@ -1427,9 +1433,9 @@ classDiagram
| **Strategy (Attention)** | `AttentionBackend`, `TorchNativeBackend`, `CudaBackend` | Attention computation backend switching via context manager | | **Strategy (Attention)** | `AttentionBackend`, `TorchNativeBackend`, `CudaBackend` | Attention computation backend switching via context manager |
| **Auto-dispatch (Rotary)** | `apply_rotary_emb`, `rotary_backend.py`, `rotary_ops.py` | Rotary embedding CUDA kernel auto-dispatch with torch fallback | | **Auto-dispatch (Rotary)** | `apply_rotary_emb`, `rotary_backend.py`, `rotary_ops.py` | Rotary embedding CUDA kernel auto-dispatch with torch fallback |
| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution | | **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution |
| **Storage** | `Store`, `H5Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support | | **Storage** | `Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support |
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching | | **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
| **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading | | **Model Registry** | `ModelFactory`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
## Core Relationships ## Core Relationships
@@ -1439,10 +1445,10 @@ classDiagram
4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)``NoneExecutor` / `DDPExecutor` / `FSDPExecutor` 4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)``NoneExecutor` / `DDPExecutor` / `FSDPExecutor`
5. **Inference Flow**: `InferenceEngine``InferenceScheduler``AutoRegressiveLM`, backed by `PagePool` + `KVCache` + `SamplingPipeline`. Attention backend selected via `attn_backend()` context manager (`TorchNativeBackend` default, `CudaBackend` for CUDA kernels). Rotary embedding auto-dispatches to CUDA kernel when available (inference mode), else torch complex multiply (training). 5. **Inference Flow**: `InferenceEngine``InferenceScheduler``AutoRegressiveLM`, backed by `PagePool` + `KVCache` + `SamplingPipeline`. Attention backend selected via `attn_backend()` context manager (`TorchNativeBackend` default, `CudaBackend` for CUDA kernels). Rotary embedding auto-dispatches to CUDA kernel when available (inference mode), else torch complex multiply (training).
6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP 6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (H5Store/MmapStore/JsonlStore) loads data with explicit `_length` and multi-segment `_data` 7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (`MmapStore`/`JsonlStore`) loads data with explicit `_length` and multi-segment `_data`
8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only), extra state saved as `{key}.pt` 8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata; `CheckpointCallback` performs rank-0 training saves, with extra state saved as `{key}.pt`
9. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`/`WSDScheduler` 9. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`/`WSDScheduler`
10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops 10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers 11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers
> Document Update Time: 2026-07-31 > Document Update Time: 2026-08-02
+69 -35
View File
@@ -14,26 +14,30 @@ This document describes the data pipeline: from raw text to model input tensors.
## Overview ## Overview
``` ```
JSONL Lines → Pipeline (mask builder) → Tokenized Tensors JSON / JSONL Records → Pipeline (mask builder) → Tokenized Tensors
.h5 or .bin storage .bin storage
Store.load() Store.load()
Store.fetch(begin, end, keys) Store.fetch(begin, end, keys)
BaseDataset.__getitem__(idx) Dataset.__getitem__(idx)
Sampler → DataLoader → Training / Inference RDSampler → DataLoader → Training
``` ```
## Data Preparation ## Data Preparation
Raw text is tokenized via `AutoTokenizer.encode()` and saved as HDF5 (`.h5`) or binary (`.bin` + `meta.json`) files with keyed tensor groups. The offline `Pipeline` accepts `.jsonl` records and `.json` files containing one
object or a list of objects. It tokenizes them and writes binary shards (`.bin`
plus `meta.json`) with keyed tensor groups. Binary is the only registered output
writer; the pipeline cannot emit JSONL.
### Tokenization ### Tokenization
The `Pipeline` reads JSONL lines, applies the mask builder (see [Preprocessing](../guides/preprocessing.md)), and produces flat token sequences: The `Pipeline` reads JSON/JSONL records, applies the mask builder (see
[Preprocessing](../guides/preprocessing.md)), and produces token sequences:
```python ```python
# Per JSONL line: messages → chat template → token IDs + loss mask # Per JSONL line: messages → chat template → token IDs + loss mask
@@ -42,84 +46,114 @@ loss_mask = [0, 0, 0, 1, 1, 1, 1, 1, 1] # 0=masked, 1=train
# Stored as flat tensors, packed with other lines by packing strategy # Stored as flat tensors, packed with other lines by packing strategy
``` ```
The output `meta.json` records the storage format, key names, dtype, total token count, and tensor shapes for each shard. For default single-output preprocessing, the stored keys are `sequence` and
`position_ids`, plus `loss_mask` when masking is required. Packing is supported
for single-output data with a `sequence` key. Shard flushing counts the primary
flat sequence for each record: `sequence` in single-output mode, otherwise the
first flat source output.
The exact shard `meta.json` schema is a top-level mapping from key to tensor
metadata. It does not contain a storage-format or total-token field:
```json
{
"sequence": {"shape": [123456], "dtype": "int32"},
"loss_mask": {"shape": [123456], "dtype": "bool"},
"position_ids": {"shape": [123456], "dtype": "int32"}
}
```
Record-aware binary data may also include `"offsets": [0, ...]` inside a key's
metadata, but the preprocessing `BinWriter` currently does not write offsets.
### Format Detection ### Format Detection
`detect_format(load_path)` inspects the path: `detect_format(load_path)` inspects the path:
- If `load_path` is a file: checks suffix — `.h5`/`.hdf5``"h5"`, `.jsonl` `"jsonl"`, unknown suffix raises `ValueError` - If `load_path` is a file: `.jsonl` selects `"jsonl"`; other suffixes raise `ValueError`.
- If `load_path` is a directory: recursively globs for `*.h5`/`*.hdf5` files → `"h5"`, `*.bin` + `**/meta.json` `"bin"`, or `*.jsonl` + `dataset_config.json` `"jsonl"` - If `load_path` is a directory: any recursive `*.bin` plus a `meta.json` selects `"bin"`; otherwise any recursive `*.jsonl` selects `"jsonl"`.
- Detection does not require `dataset_config.json`; configuration is selected later when `JsonlStore.load()` chooses a transform.
### Store Backends ### Store Backends
Storage format is auto-detected by `detect_format()`; backends are dispatched via registry: Storage format is auto-detected by `detect_format()`; backends are dispatched via registry:
``` ```
StoreFactory.create("h5") → H5Store
StoreFactory.create("bin") → MmapStore StoreFactory.create("bin") → MmapStore
StoreFactory.create("jsonl") → JsonlStore StoreFactory.create("jsonl") → JsonlStore
``` ```
All three inherit `Store` (base, owns `_data`/`_cum`/`_offsets`/`_normalize`) plus the `Streamable` and `Recordable` mixins, so every backend supports both `fetch(begin, end, keys)` (stream) and `fetch_record(index, keys)` (record) APIs. Both stores inherit `Store` and compose the `Streamable` and `Recordable`
access methods.
**H5Store**: Reads HDF5 files. Tensors are loaded into host memory and normalized into segmented storage. `segments_are_records=True` — each `data_i` dataset is one record.
**MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`. `segments_are_records=False` — bin segments are contiguous streams; record access is driven by `_offsets` (written when `save_bin(..., record_keys=...)` was used at preprocessing time). **MmapStore**: Memory-maps `.bin` files. OS page cache sharing is native — no explicit `share_memory_()` needed. Uses `torch.from_numpy(np.memmap(...))`. `segments_are_records=False` — bin segments are contiguous streams; record access is driven by `_offsets` (written when `save_bin(..., record_keys=...)` was used at preprocessing time).
**JsonlStore**: On-the-fly tokenization of raw JSONL files at load time. Requires a `dataset_config.json` alongside the `.jsonl` files following the same `PipelineConfig` schema with an additional `tokenizer_path` field. Two modes: eager (default, applies `TokenizeTransform` to all records at load) and lazy (`processor=fn` given, defers tokenisation to `fetch_record` — used by DPO/GRPO). **JsonlStore**: Reads a `.jsonl` file or the sorted top-level `*.jsonl` files in
a directory. Eager transform selection uses the first available route:
All backends normalise tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for bisect-based stream indexing) + `Store._offsets[Dict[str, List[int]]]` (per-record offsets for record-mode indexing). Nested keys (GRPO `responses`/`masks` as `List[List[Tensor]]`) are stored as-is and excluded from both bookkeepings — they are only accessed record-by-record. 1. An explicit `transform=` argument.
2. `dataset_config.json` in the JSONL directory. It follows `PipelineConfig` and may add `tokenizer_path`; when omitted, the config directory is used.
3. The built-in `messages` transform when `tokenizer_path=` is supplied. It masks system/user turns, trains assistant turns, and emits document-reset position IDs.
Only DPO gets an automatic lazy route from `DatasetFactory`: raw JSONL plus
`tokenizer_path` installs `dpo_processor` and tokenizes each record in
`fetch_record`. GRPO does not currently have an automatic lazy processor.
Eager-loaded stores normalize tensors into `Store._data[Dict[str, List[Tensor]]]` + `Store._cum[Dict[str, List[int]]]` (cumulative lengths for stream indexing) + `Store._offsets[Dict[str, List[int]]]` (per-record offsets for record indexing). Nested JSONL keys such as GRPO `responses`/`masks` are kept as record values and excluded from stream bookkeeping. Lazy DPO instead retains raw records and processes them in `fetch_record`.
## Data Keys by Training Type ## Data Keys by Training Type
| Type | Storage Keys | Access Mode | | Type | Storage Keys | Access Mode |
|------|-------------|-------------| |------|-------------|-------------|
| `seq` | `sequence` (→ input_ids, target_ids via offset-by-1) | stream (`fetch`) | | `seq` | `sequence`, `position_ids` by default (`SEQDataset` consumes only `sequence`) | stream (`fetch`) |
| `sft` | `sequence`, `loss_mask`, `position_ids` | stream (`fetch`) | | `sft` | `sequence`, `loss_mask`, `position_ids` | stream (`fetch`) |
| `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` | record (`fetch_record`) | | `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` | record (`fetch_record`) |
| `grpo` | `prompts`, `responses`, `masks`, `rewards` | record (`fetch_record`) | | `grpo` | `prompts`, `responses`, `masks`, `rewards` | record (`fetch_record`) |
Offline `.bin` output from DPO/GRPO preprocessing is not currently loadable for
training. DPO shards are written without record offsets, while GRPO response
groups are flattened without preserving record/group boundaries. Supported raw
routes are eager JSONL for SEQ/SFT and automatic lazy JSONL for DPO. GRPO
requires a caller-built, already-loaded record store.
## Dataset Architecture ## Dataset Architecture
``` ```
DatasetFactory.load( DatasetFactory.load(...)
train_type, load_path=None, window_size=0, stride=None,
storage_type=None, tokenizer_path=None,
max_len=2048, store=None
)
→ BaseDataset.load(load_path, storage_type=None)
→ detect_format(load_path) → detect_format(load_path)
StoreFactory.create(storage_type) optionally build dpo_processor for raw JSONL
→ Store.load(load_path) → StoreFactory.create(storage_type, window_size, stride)
→ _normalize(raw) # base Store, shared by both backends → Store.load(load_path, transform=... or processor=...)
→ Store._data[Dict[str, List[Tensor]]] → DatasetFactory.create(train_type, store=store)
+ _cum[Dict[str, List[int]]] (stream mode)
+ _offsets[Dict[str, List[int]]] (record mode)
Stream datasets (SEQ/SFT): Stream datasets (SEQ/SFT):
BaseDataset.__getitem__(idx) BaseDataset.__getitem__(idx)
get_index(idx) → [begin, end) Store.sample_window(idx) → [begin, end)
→ Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor] → Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor]
Record datasets (DPO/GRPO via RecordDataset): Record datasets (DPO/GRPO):
RecordDataset.__getitem__(idx) DPODataset/GRPODataset.__getitem__(idx)
→ Store.fetch_record(idx, keys) → Tensor / Dict[str, Tensor] → Store.fetch_record(idx, keys) → Tensor / Dict[str, Tensor]
``` ```
Class hierarchy: `BaseDataset` `SEQDataset` / `SFTDataset` (stream); `BaseDataset``RecordDataset``DPODataset` / `GRPODataset` (record). Class hierarchy: `BaseDataset` is the direct base of `SEQDataset`, `SFTDataset`,
`DPODataset`, and `GRPODataset`. There is no `RecordDataset` class.
`window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`, optional). Only meaningful for stream datasets — record datasets ignore both. `storage_type` defaults to `None` (auto-detect via `detect_format`). `window_size` = max input length, `stride` = step between consecutive samples (defaults to `window_size`, optional). Only meaningful for stream datasets — record datasets ignore both. `storage_type` defaults to `None` (auto-detect via `detect_format`).
`tokenizer_path` triggers lazy on-the-fly tokenisation for record datasets on raw JSONL (DPO builds a `dpo_processor`; SEQ/SFT/pre-tokenised backends ignore it). `store` (pre-built `Store`) bypasses `load_path`/`storage_type`/`tokenizer_path` entirely — the caller controls Store construction. For raw JSONL, `tokenizer_path` builds the lazy processor only for DPO. For
SEQ/SFT it is forwarded to `JsonlStore` so the built-in eager `messages`
transform can be selected when no `dataset_config.json` exists. GRPO receives no
automatic processor. A pre-built `store` bypasses path, format, tokenizer,
window, and stride setup entirely.
`Store.fetch(begin, end, keys)` (stream mode, on `Streamable`): accepts a single key (`str`) returning a `Tensor`, or a list of keys returning `Dict[str, Tensor]`. Internally uses `bisect` across multi-segment tensors. Raises `RuntimeError("Store not loaded")` if called before `load()`. `Store.fetch(begin, end, keys)` (stream mode, on `Streamable`): accepts a single key (`str`) returning a `Tensor`, or a list of keys returning `Dict[str, Tensor]`. Internally uses `bisect` across multi-segment tensors. Raises `RuntimeError("Store not loaded")` if called before `load()`.
`Store.fetch_record(index, keys)` (record mode, on `Recordable`): same key API. Uses `_offsets[key]` when present (bin layout with per-record offsets), otherwise indexes `_data[key]` directly (H5/JSONL where each segment is one record). `Store.fetch_record(index, keys)` (record mode, on `Recordable`): same key API. Uses `_offsets[key]` when present for binary record layouts; otherwise it indexes per-record JSONL tensors directly.
## Sampler ## Sampler
`ResumableDistributedSampler` supports checkpoint-aware distributed sampling: `RDSampler` supports checkpoint-aware distributed sampling:
- Tracks `start_epoch` / `start_iter` for resume - Tracks `start_epoch` / `start_iter` for resume
- Shuffle via `torch.Generator(seed + epoch)` - Shuffle via `torch.Generator(seed + epoch)`
+18 -10
View File
@@ -41,7 +41,12 @@ 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 $$ $$ 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 $$
`RotaryEmbedding` pre-computes `cos_table` and `sin_table` (f32, `[max_len, dim/2]`). `forward()` returns a `(cos, sin)` tuple indexed by `position_ids`. `apply_rotary_emb` applies the rotation: during training it uses torch complex multiply (autograd-compatible); during inference it auto-dispatches to a fused CUDA kernel when available. The key property is that the dot product $q_i^T k_j$ depends only on the relative position $i - j$, not the absolute positions. `RotaryEmbedding` pre-computes a complex `freqs_cis` buffer. `forward()` returns
a tensor indexed by `position_ids`. `apply_rotary_emb` applies the rotation:
during training it uses torch complex multiply (autograd-compatible); during
inference it auto-dispatches to a fused CUDA kernel when available. The key
property is that the dot product $q_i^T k_j$ depends only on the relative
position $i - j$, not the absolute positions.
**Critical for inference**: RoPE is applied **before** KV cache write, not after. If applied after caching, position encoding drift occurs because cached K/V would have stale rotation factors. **Critical for inference**: RoPE is applied **before** KV cache write, not after. If applied after caching, position encoding drift occurs because cached K/V would have stale rotation factors.
@@ -51,13 +56,13 @@ $$ 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_{
Next-token cross-entropy with optional label smoothing: Next-token cross-entropy with optional label smoothing:
$$ L_{\text{PT}} = -\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta) $$ $$ L_{\text{PT}} = -\frac{1}{T}\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta) $$
### SFT (Supervised Fine-Tuning) ### SFT (Supervised Fine-Tuning)
Masked cross-entropy (`ignore_index=-100`) over response tokens only: Masked cross-entropy (`ignore_index=-100`) over response tokens only:
$$ L_{\text{SFT}} = -\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta) $$ $$ L_{\text{SFT}} = -\frac{1}{L}\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta) $$
Prompt tokens are masked out via `loss_mask`; only response tokens contribute to the loss. Prompt tokens are masked out via `loss_mask`; only response tokens contribute to the loss.
@@ -98,8 +103,8 @@ on_train_begin
model.train() model.train()
on_epoch_begin on_epoch_begin
for batch in dataloader: for batch in dataloader:
on_batch_begin
with executor.accumulate(model): with executor.accumulate(model):
on_batch_begin
loss_output = strategy(batch) loss_output = strategy(batch)
context.loss = loss_output["loss"].item() context.loss = loss_output["loss"].item()
context.metrics = loss_output["metrics"] context.metrics = loss_output["metrics"]
@@ -113,6 +118,7 @@ on_train_begin
if executor.sync_gradients: if executor.sync_gradients:
on_optimizer_step on_optimizer_step
optimizer.step() optimizer.step()
strategy.on_optimizer_step()
optimizer.zero_grad() optimizer.zero_grad()
if scheduler: if scheduler:
scheduler.step() scheduler.step()
@@ -121,21 +127,23 @@ on_train_end
``` ```
The loss is divided by `grad_accum_steps` before `backward()`, so accumulated gradients sum to the correct mean. The loss is divided by `grad_accum_steps` before `backward()`, so accumulated gradients sum to the correct mean.
Strategy metrics are detached and converted to Python `float` values before the
`LossOutput` is returned; only `LossOutput.loss` remains a differentiable tensor.
## Callback Lifecycle ## Callback Lifecycle
| Hook | Fires | Default callback | | Hook | Fires | Default callback |
|------|-------|-----------------| |------|-------|-----------------|
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback` | | `on_train_begin` | Before training starts | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` | | `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
| `on_batch_begin` | Every batch | — | | `on_batch_begin` | Every batch | — |
| `on_optimizer_step` | Every accumulation window | `GradientClippingCallback`, `MetricCallback`, `ProgressBarCallback` | | `on_optimizer_step` | Every accumulation window | `MetricCallback`, `ProgressBarCallback`, `GradientClippingCallback` |
| `on_batch_end` | Every batch | `CheckpointCallback` | | `on_batch_end` | Every batch | `CheckpointCallback` |
| `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` | | `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` |
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` | | `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` | | `on_train_end` | Training exits after `on_train_begin` completes (via `finally`) | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping` (always registered; computes grad norm, clips only when `max_grad_norm` is not `None`). Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm, rank-0), `gradient_clipping`. The gradient-clipping callback is always registered and always calls `executor.clip_grad_norm()` with the numeric `max_grad_norm` value.
## KV Cache Mathematics ## KV Cache Mathematics
@@ -160,7 +168,7 @@ Three-layer separation (SGLang-inspired):
- **ReqToTokenPool**: Index table `[req_idx, pos] → physical token slot`, shared across all layers. - **ReqToTokenPool**: Index table `[req_idx, pos] → physical token slot`, shared across all layers.
- **Allocator + PrefixCache**: Paged-mode slot allocation with ref-counting, LRU eviction, and hash-based prefix sharing. - **Allocator + PrefixCache**: Paged-mode slot allocation with ref-counting, LRU eviction, and hash-based prefix sharing.
`PagePool` orchestrates all three. In contiguous mode (default), `req_to_token` is a trivial linear mapping. In paged mode, slots are allocated on demand with prefix caching support. `bind_tasks()` returns a `KVCache` dataclass with precomputed `page_table` and `decode_mask` fields (computed once per decode step, shared across all layers). Attention layers access buffers directly — no methods, no abstraction. `PagePool` orchestrates all three. In contiguous mode (default), `req_to_token` is a trivial linear mapping. In paged mode, slots are allocated on demand with prefix caching support. `bind_tasks()` returns a `KVCache` dataclass with `kv_indptr`, a prefix-sum index over sequence lengths computed once per step and shared across layers. Attention layers access buffers directly — no methods, no abstraction.
### Attention Backend ### Attention Backend
@@ -239,4 +247,4 @@ total_steps = (batches_per_replica // grad_accum_steps) * n_epoch
This accounts for data-parallel sharding — each rank processes `1/nprocs` of the dataset. This accounts for data-parallel sharding — each rank processes `1/nprocs` of the dataset.
> Document Update Time: 2026-07-31 > Document Update Time: 2026-08-02
+22 -3
View File
@@ -2,11 +2,23 @@
This guide walks you through installing AstrAI, downloading a model, running inference, preprocessing data, and launching your first training job. This guide walks you through installing AstrAI, downloading a model, running inference, preprocessing data, and launching your first training job.
## Contents
- [Prerequisites](#prerequisites)
- [1. Install](#1-install)
- [2. Download Model Weights](#2-download-model-weights)
- [3. Run Inference](#3-run-inference)
- [4. Preprocess Data](#4-preprocess-data)
- [5. Train](#5-train)
- [6. Evaluate](#6-evaluate)
- [7. Docker](#7-docker)
- [Next Steps](#next-steps)
## Prerequisites ## Prerequisites
- **Python 3.12+** - **Python 3.12+**
- **PyTorch 2.11+** (CUDA 12.8 recommended for GPU support) - **PyTorch 2.11.0** (the exact version pinned by AstrAI; CUDA 12.8 build recommended for GPU support)
- NVIDIA GPU with CUDA (optional but recommended; CPU works for inference) - NVIDIA GPU with CUDA for training, `scripts/tools/generate.py`, generation evaluations, and demos. The HTTP server and direct-scoring evaluations can run on CPU where their CLI exposes a CPU device.
## 1. Install ## 1. Install
@@ -55,7 +67,7 @@ python scripts/demo/stream_chat.py
# Type your message after >>, type !exit to quit # Type your message after >>, type !exit to quit
``` ```
This starts a multi-turn interactive chat session with streaming output. This starts a single-turn interactive prompt loop with streaming output. Each prompt is independent; conversation history is not retained.
### Start an HTTP Server ### Start an HTTP Server
@@ -192,6 +204,13 @@ See [Training Guide](guides/training.md) for loss formulas and strategies. See [
## 6. Evaluate ## 6. Evaluate
HumanEval and MMLU download their benchmark data through HuggingFace
`datasets`, which is not part of the base install:
```bash
pip install datasets
```
```bash ```bash
# HumanEval (code generation, auto-downloads dataset) # HumanEval (code generation, auto-downloads dataset)
python scripts/eval/evaluate_humaneval.py --param_path ./params --num_samples 20 python scripts/eval/evaluate_humaneval.py --param_path ./params --num_samples 20
+25 -17
View File
@@ -2,6 +2,18 @@
AstrAI supports three parallel modes: **single GPU** (`none`), **Data Parallel** (`ddp`), and **Fully Sharded Data Parallel** (`fsdp`). This guide covers when to use each, how to launch multi-GPU training, and how gradient accumulation works. AstrAI supports three parallel modes: **single GPU** (`none`), **Data Parallel** (`ddp`), and **Fully Sharded Data Parallel** (`fsdp`). This guide covers when to use each, how to launch multi-GPU training, and how gradient accumulation works.
## Contents
- [Quick Start](#quick-start)
- [Parallel Modes](#parallel-modes)
- [Gradient Accumulation](#gradient-accumulation)
- [Process Launching](#process-launching)
- [NCCL Troubleshooting](#nccl-troubleshooting)
- [Checkpoint Saving](#checkpoint-saving)
- [Total Steps Calculation](#total-steps-calculation)
- [Real Examples](#real-examples)
- [CLI Parameters](#cli-parameters)
## Quick Start ## Quick Start
### Single GPU ### Single GPU
@@ -21,9 +33,6 @@ python scripts/tools/train.py \
```bash ```bash
export CUDA_VISIBLE_DEVICES=0,1,2,3 export CUDA_VISIBLE_DEVICES=0,1,2,3
export NCCL_P2P_DISABLE=1
export NCCL_NET_GDR_LEVEL=0
python scripts/tools/train.py \ python scripts/tools/train.py \
--train_type=sft \ --train_type=sft \
--param_path ./params \ --param_path ./params \
@@ -38,9 +47,6 @@ python scripts/tools/train.py \
```bash ```bash
export CUDA_VISIBLE_DEVICES=0,1,2,3 export CUDA_VISIBLE_DEVICES=0,1,2,3
export NCCL_P2P_DISABLE=1
export NCCL_NET_GDR_LEVEL=0
python scripts/tools/train.py \ python scripts/tools/train.py \
--train_type=sft \ --train_type=sft \
--param_path ./params \ --param_path ./params \
@@ -110,7 +116,7 @@ AstrAI auto-detects the launch method:
| Detection | Strategy | Use Case | | Detection | Strategy | Use Case |
|-----------|----------|----------| |-----------|----------|----------|
| `torchelastic` / `torchrun` env vars | `TorchrunStrategy` | External orchestrator (torchrun, SLURM, K8s) | | `torchelastic` / `torchrun` env vars | `TorchrunStrategy` | External orchestrator (`torchrun`, K8s) |
| `RANK` + `WORLD_SIZE` env vars | `TorchrunStrategy` | External launch | | `RANK` + `WORLD_SIZE` env vars | `TorchrunStrategy` | External launch |
| Neither | `LocalStrategy` | `python scripts/tools/train.py` (in-process spawn) | | Neither | `LocalStrategy` | `python scripts/tools/train.py` (in-process spawn) |
@@ -126,23 +132,28 @@ For multi-node or SLURM environments:
torchrun --nproc_per_node=4 scripts/tools/train.py \ torchrun --nproc_per_node=4 scripts/tools/train.py \
--train_type=sft \ --train_type=sft \
--parallel_mode=ddp \ --parallel_mode=ddp \
--nprocs=4 \
--param_path ./params \ --param_path ./params \
--data_root_path ./dataset \ --data_root_path ./dataset \
--batch_per_device=4 --batch_per_device=4
``` ```
When launched via torchrun, AstrAI reads `RANK`, `WORLD_SIZE`, `LOCAL_RANK` from the environment and uses `TorchrunStrategy`. The `--nprocs` flag is ignored (the orchestrator controls process count). When launched via `torchrun`, the launcher creates the worker processes. AstrAI reads `RANK`, `WORLD_SIZE`, and `LOCAL_RANK` from the environment and uses `TorchrunStrategy`; `--nprocs` does not control process creation in this mode.
## NCCL Environment Variables The current training CLI still uses `--nprocs` when calculating scheduler `total_steps`. Set it to the global `WORLD_SIZE` so the step count reflects data-parallel sharding, including multi-node runs.
For multi-GPU training, you **must** set these environment variables: Raw Slurm variables such as `SLURM_PROCID`, `SLURM_NTASKS`, and `SLURM_LOCALID` are not recognized automatically. Launch through `torchrun`, or map the scheduler's variables to `RANK`, `WORLD_SIZE`, `LOCAL_RANK`, `MASTER_ADDR`, and `MASTER_PORT` before starting AstrAI. The same requirement applies to launchers that expose only OpenMPI-specific variables.
## NCCL Troubleshooting
The following variables are troubleshooting options for hardware or network configurations where NCCL hangs or fails. They are not general requirements and can reduce performance by disabling peer-to-peer or GPUDirect RDMA paths:
```bash ```bash
export NCCL_P2P_DISABLE=1 export NCCL_P2P_DISABLE=1
export NCCL_NET_GDR_LEVEL=0 export NCCL_NET_GDR_LEVEL=0
``` ```
These are required on certain hardware configurations (see `AGENTS.md`). Without them, NCCL may hang or crash during collective operations. These are set in the training shell scripts (`train-seq.sh`, `train-sft.sh`, `train-dpo.sh`) but not in Python code — you must export them before launching. Apply them only after confirming the relevant NCCL transport is the source of the failure. AstrAI does not set them in Python.
## Checkpoint Saving ## Checkpoint Saving
@@ -176,9 +187,6 @@ This ensures the LR schedule is correctly scaled regardless of the number of GPU
```bash ```bash
export CUDA_VISIBLE_DEVICES=0,1,2,3 export CUDA_VISIBLE_DEVICES=0,1,2,3
export NCCL_P2P_DISABLE=1
export NCCL_NET_GDR_LEVEL=0
python scripts/tools/train.py \ python scripts/tools/train.py \
--train_type=seq \ --train_type=seq \
--param_path ./params \ --param_path ./params \
@@ -240,7 +248,7 @@ python scripts/tools/train.py \
| Parameter | Default | Description | | Parameter | Default | Description |
|-----------|---------|-------------| |-----------|---------|-------------|
| `--nprocs` | 1 | Number of GPUs / processes | | `--nprocs` | 1 | Local process count for AstrAI's launcher; under `torchrun`, set it to global `WORLD_SIZE` for total-step calculation |
| `--parallel_mode` | `fsdp` | `none`, `ddp`, or `fsdp` | | `--parallel_mode` | `fsdp` | `none`, `ddp`, or `fsdp` |
| `--start_method` | `spawn` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) | | `--start_method` | `spawn` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) |
| `--backend` | `nccl` | Distributed backend (`nccl`, `gloo`) | | `--backend` | `nccl` | Distributed backend (`nccl`, `gloo`) |
@@ -248,8 +256,8 @@ python scripts/tools/train.py \
| `--master_port` | `29500` | Master node port | | `--master_port` | `29500` | Master node port |
| `--device_type` | `cuda` | Device type | | `--device_type` | `cuda` | Device type |
> `--tp_size` is parsed but **not yet wired** — tensor parallelism is future work. `ColumnParallelLinear` / `RowParallelLinear` exist in `astrai/parallel/module.py` but are not used by the model. > `--tp_size` is accepted by the CLI but discarded before configuration. Tensor parallelism is not implemented, and there is no tensor-parallel module or model integration.
Full parameter reference: [CLI Reference](params.md). Training loop and strategies: [Training Guide](training.md). Full parameter reference: [CLI Reference](params.md). Training loop and strategies: [Training Guide](training.md).
> Document Update Time: 2026-07-30 > Document Update Time: 2026-08-02
+46 -12
View File
@@ -2,6 +2,29 @@
AstrAI provides 7 evaluation scripts in `scripts/eval/` covering code generation, knowledge QA, perplexity, summarization, data quality, instruction following, and weight analysis. AstrAI provides 7 evaluation scripts in `scripts/eval/` covering code generation, knowledge QA, perplexity, summarization, data quality, instruction following, and weight analysis.
## Contents
- [Prerequisites](#prerequisites)
- [Overview](#overview)
- [HumanEval](#humaneval-code-generation)
- [MMLU](#mmlu-knowledge-qa)
- [Perplexity](#perplexity-ppl)
- [ROUGE](#rouge)
- [IFD](#ifd-instruction-following-difficulty)
- [IFEval](#ifeval-instruction-following)
- [Weight Analysis](#weight-analysis)
- [Tips](#tips)
## Prerequisites
HumanEval, MMLU, and IFEval import HuggingFace `datasets` to download their benchmark data. This package is not installed by AstrAI's base dependencies, so install it before running those scripts:
```bash
pip install datasets
```
The generation-based scripts require CUDA because they load the model on `cuda` with `bfloat16`. Direct-scoring and metric scripts support the devices shown below.
## Overview ## Overview
| Script | Metric | Model Invocation | External Dataset | | Script | Metric | Model Invocation | External Dataset |
@@ -18,7 +41,15 @@ Two invocation patterns exist:
- **Generation benchmarks** (HumanEval, IFEval): use `InferenceEngine` to generate responses, then score them. - **Generation benchmarks** (HumanEval, IFEval): use `InferenceEngine` to generate responses, then score them.
- **Scoring benchmarks** (MMLU, PPL, IFD): call `model()` directly under `torch.inference_mode()` for log-likelihood computation. - **Scoring benchmarks** (MMLU, PPL, IFD): call `model()` directly under `torch.inference_mode()` for log-likelihood computation.
Common defaults: `--param_path` defaults to `./params`; dtype defaults to `bfloat16` on CUDA, `float32` on CPU. | Script | Device support |
|--------|----------------|
| HumanEval | CUDA for generation; `--test_only` can score existing completions without loading a model |
| IFEval | CUDA only |
| MMLU | CUDA or CPU via `--device`; auto-selects CUDA when available |
| PPL | CUDA or CPU via `--device`; auto-selects CUDA when available |
| IFD | CUDA or CPU via `--device`; auto-selects CUDA when available |
| ROUGE | CPU-only metric computation; no model is loaded |
| Weight analysis | CUDA by default; CPU supported via `--device cpu` |
--- ---
@@ -30,7 +61,7 @@ Generates completions for 164 programming problems, executes them against hidden
python scripts/eval/evaluate_humaneval.py \ python scripts/eval/evaluate_humaneval.py \
--param_path ./params \ --param_path ./params \
--num_samples 20 \ --num_samples 20 \
--batch_size 32 \ --batch_size 64 \
--max_tokens 512 \ --max_tokens 512 \
--output results/humaneval.json --output results/humaneval.json
``` ```
@@ -47,7 +78,8 @@ python scripts/eval/evaluate_humaneval.py \
| `--temperature` | 0.8 | Sampling temperature | | `--temperature` | 0.8 | Sampling temperature |
| `--top_p` | 0.95 | Nucleus sampling threshold | | `--top_p` | 0.95 | Nucleus sampling threshold |
| `--top_k` | 50 | Top-k sampling | | `--top_k` | 50 | Top-k sampling |
| `--batch_size` | 32 | Generation batch size | | `--batch_size` | 64 | Generation batch size |
| `--max_seq_len` | 4096 | KV cache sequence length |
| `--test_workers` | 8 | ProcessPoolExecutor workers for test execution | | `--test_workers` | 8 | ProcessPoolExecutor workers for test execution |
| `--test_timeout` | 3.0 | Per-subprocess timeout (seconds) | | `--test_timeout` | 3.0 | Per-subprocess timeout (seconds) |
| `--problems` | None | Restrict to specific problem indices | | `--problems` | None | Restrict to specific problem indices |
@@ -66,7 +98,7 @@ python scripts/eval/evaluate_humaneval.py \
python scripts/eval/evaluate_mmlu.py \ python scripts/eval/evaluate_mmlu.py \
--param_path ./params \ --param_path ./params \
--n_shot 5 \ --n_shot 5 \
--subjects math_algebra history_us \ --subjects abstract_algebra high_school_us_history \
--output results/mmlu.json --output results/mmlu.json
``` ```
@@ -82,12 +114,13 @@ python scripts/eval/evaluate_mmlu.py \
| `--device` | auto | Device (`cuda` / `cpu`) | | `--device` | auto | Device (`cuda` / `cpu`) |
| `--dtype` | auto | `bfloat16` on CUDA, `float32` on CPU | | `--dtype` | auto | `bfloat16` on CUDA, `float32` on CPU |
| `--seed` | 0 | Seed for option permutation (0 = enabled, -1 = disabled) | | `--seed` | 0 | Seed for option permutation (0 = enabled, -1 = disabled) |
| `--batch_size` | 4 | Questions per batch; each question produces four choice rows |
**How it works**: For each question, builds a prompt with n-shot examples, then scores each choice (A/B/C/D) by computing the summed log-likelihood of the choice token given the context. The choice with the highest log-prob is the prediction. **How it works**: For each question, builds a prompt with n-shot examples, then scores each choice (A/B/C/D) by computing the summed log-likelihood of the choice token given the context. The choice with the highest log-prob is the prediction.
**Output**: stdout prints per-subject accuracy and overall. With `--output`, writes per-subject `{accuracy, correct, total}` + `_overall` aggregate. **Output**: stdout prints per-subject accuracy and overall. With `--output`, writes per-subject `{accuracy, correct, total}` + `_overall` aggregate.
**Data**: Auto-downloads `cais/mmlu` from HuggingFace. Stored as per-subject CSVs in `<data_dir>/<split>/` and `<data_dir>/dev/` (for few-shot). **Data**: Auto-downloads `cais/mmlu` from HuggingFace. Stored as per-subject CSVs in `<data_dir>/<split>/` and `<data_dir>/dev/` (for few-shot). `--subjects` accepts canonical MMLU names such as `abstract_algebra`, `college_computer_science`, `high_school_us_history`, and `world_religions`.
--- ---
@@ -100,7 +133,7 @@ python scripts/eval/evaluate_ppl.py \
--param_path ./params \ --param_path ./params \
--input_path data.jsonl \ --input_path data.jsonl \
--output_dir ppl_results/ \ --output_dir ppl_results/ \
--batch_size 4 \ --batch_size 64 \
--max_length 2048 --max_length 2048
``` ```
@@ -110,7 +143,7 @@ python scripts/eval/evaluate_ppl.py \
| `--input_path` | required | Input file, glob, or directory | | `--input_path` | required | Input file, glob, or directory |
| `--output_dir` | required | Output directory for `summary.json` + token JSONL | | `--output_dir` | required | Output directory for `summary.json` + token JSONL |
| `--text_key` | `text` | Key for the text field in input data | | `--text_key` | `text` | Key for the text field in input data |
| `--batch_size` | 4 | Batch size | | `--batch_size` | 64 | Batch size |
| `--max_length` | 2048 | Max sequence length (tokens) | | `--max_length` | 2048 | Max sequence length (tokens) |
| `--token_level` | False | Store per-token log_probs + token-type analysis | | `--token_level` | False | Store per-token log_probs + token-type analysis |
| `--max_samples` | None | Random subsample per file | | `--max_samples` | None | Random subsample per file |
@@ -119,7 +152,7 @@ python scripts/eval/evaluate_ppl.py \
**Input**: JSONL or JSON files. Each item must have a field named by `--text_key` (default `text`). If `--input_path` is a directory, recursively collects `*.jsonl` and `*.json`. **Input**: JSONL or JSON files. Each item must have a field named by `--text_key` (default `text`). If `--input_path` is a directory, recursively collects `*.jsonl` and `*.json`.
**Output**: `summary.json` with per-file stats (tokens, mean/median loss, perplexity, p50/p90/p95/p99). With `--token_level`, also writes per-token JSONL with token IDs and log-probs. **Output**: `summary.json` with per-file token count, mean loss, perplexity, and p50/p90/p95/p99 loss. Median loss is included only with `--token_level`; that mode also writes per-token JSONL with token IDs and log-probs.
--- ---
@@ -210,7 +243,8 @@ python scripts/eval/evaluate_ifeval.py \
| `--top_p` | 0.95 | Top-p sampling | | `--top_p` | 0.95 | Top-p sampling |
| `--top_k` | 50 | Top-k sampling | | `--top_k` | 50 | Top-k sampling |
| `--num_samples` | 1 | Samples per problem (best-of-n scoring) | | `--num_samples` | 1 | Samples per problem (best-of-n scoring) |
| `--batch_size` | 1 | Inference batch size | | `--batch_size` | 64 | Inference batch size |
| `--max_seq_len` | 4096 | KV cache sequence length |
| `--limit` | None | Limit to first N problems (quick testing) | | `--limit` | None | Limit to first N problems (quick testing) |
| `--dump_responses` | None | Path to dump raw responses as JSONL | | `--dump_responses` | None | Path to dump raw responses as JSONL |
@@ -232,7 +266,7 @@ python scripts/eval/analyze_weights.py \
| Parameter | Default | Description | | Parameter | Default | Description |
|-----------|---------|-------------| |-----------|---------|-------------|
| `--ckpt_dir` | required | Checkpoint dir with `model.safetensors` + `config.json` | | `--ckpt_dir` | required | Checkpoint directory containing `model.safetensors` |
| `--compare` | None | Additional checkpoint dirs to compare | | `--compare` | None | Additional checkpoint dirs to compare |
| `--no_svd` | False | Skip SVD; show only weight stats (faster) | | `--no_svd` | False | Skip SVD; show only weight stats (faster) |
| `--output` | None | Save results as JSON | | `--output` | None | Save results as JSON |
@@ -245,8 +279,8 @@ python scripts/eval/analyze_weights.py \
## Tips ## Tips
- **Quick test**: Use `--limit` (IFEval) or `--problems` (HumanEval) to run on a small subset first. - **Quick test**: Use `--limit` (IFEval) or `--problems` (HumanEval) to run on a small subset first.
- **Auto-download**: HumanEval, MMLU, and IFEval auto-download their datasets on first run. The other scripts expect user-provided data. - **Auto-download**: After installing `datasets`, HumanEval, MMLU, and IFEval auto-download their datasets on first run. The other scripts expect user-provided data.
- **Output formats**: `--output` writes a single JSON for most scripts. PPL and IFD write an `--output_dir` containing `summary.json` plus per-file artifacts. - **Output formats**: `--output` writes a single JSON for most scripts. PPL and IFD write an `--output_dir` containing `summary.json` plus per-file artifacts.
- **CPU mode**: All scripts auto-detect CUDA. To force CPU, use `--device cpu --dtype float32`. - **CPU mode**: MMLU, PPL, and IFD support `--device cpu --dtype float32`; weight analysis supports `--device cpu`. HumanEval generation and IFEval are CUDA-only.
> Document Update Time: 2026-07-30 > Document Update Time: 2026-07-30
+52 -13
View File
@@ -49,8 +49,7 @@ KVCache
├── seq_lens [batch_size] ├── seq_lens [batch_size]
├── out_cache_loc [batch, seq_len] — write indices for this forward ├── out_cache_loc [batch, seq_len] — write indices for this forward
├── max_len int — max(seq_lens), avoids GPU sync in decode ├── max_len int — max(seq_lens), avoids GPU sync in decode
── page_table [batch, max_len] — precomputed gather indices for decode (None for prefill) ── kv_indptr [batch + 1] int32 — prefix sum of seq_lens, precomputed once per step
└── decode_mask [batch, max_len] bool — precomputed position validity mask (None for single-batch decode)
``` ```
Attention layers do raw buffer indexing: `k_buffer[layer_id, out_cache_loc] = k` to write, `k_buffer[layer_id, indices]` to gather. Attention layers do raw buffer indexing: `k_buffer[layer_id, out_cache_loc] = k` to write, `k_buffer[layer_id, indices]` to gather.
@@ -87,7 +86,9 @@ Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_b
- **CUDA kernel** (`rotary_emb.cu`): fused cos/sin lookup + rotation in a single kernel, used when the kernel is available, input is on CUDA, and `torch.is_grad_enabled()` is `False` (inference mode) - **CUDA kernel** (`rotary_emb.cu`): fused cos/sin lookup + rotation in a single kernel, used when the kernel is available, input is on CUDA, and `torch.is_grad_enabled()` is `False` (inference mode)
- **Torch fallback**: complex multiply path (`torch.view_as_complex``torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd backward) or when the CUDA kernel is not available - **Torch fallback**: complex multiply path (`torch.view_as_complex``torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd backward) or when the CUDA kernel is not available
`RotaryEmbedding` stores `cos_table`/`sin_table` as f32 buffers and returns a `(cos, sin)` tuple from `forward()`. Both attention backends share the same rotary dispatch — it is backend-agnostic. `RotaryEmbedding` stores a complex `freqs_cis` buffer and returns a tensor
from `forward()`. Both attention backends share the same rotary dispatch — it
is backend-agnostic.
## Continuous Batching ## Continuous Batching
@@ -183,22 +184,58 @@ curl -X POST http://localhost:8000/v1/messages \
-d '{"model":"astrai","system":"You are helpful.","messages":[{"role":"user","content":"Hello"}],"max_tokens":512}' -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`. Supports `stop_sequences` and streaming via `event: content_block_delta`. Anthropic streams also end with the shared `data: [DONE]` sentinel after `event: message_stop`.
### GenerationRequest Parameters ### Request Parameters
The HTTP protocols and direct engine API have distinct request models and defaults.
**OpenAI** (`ChatCompletionRequest`):
| Param | Type | Default | Description | | Param | Type | Default | Description |
|-------|------|---------|-------------| |-------|------|---------|-------------|
| `model` | str | `"astrai"` | Model name returned in responses |
| `messages` | List[dict] | required | Chat messages (role, content) | | `messages` | List[dict] | required | Chat messages (role, content) |
| `top_k` | int | 50 | Top-k count | | `temperature` | Optional[float] | 1.0 | Sampling temperature (0.0-2.0) |
| `top_p` | float | 1.0 | Nucleus threshold | | `top_p` | Optional[float] | 1.0 | Nucleus threshold (0.0-1.0) |
| `temperature` | float | 1.0 | Sampling temperature (> 0.0) | | `top_k` | Optional[int] | 50 | Top-k count |
| `max_tokens` | Optional[int] | None | Max generation length | | `max_tokens` | Optional[int] | 2048 | Max generation length |
| `stream` | bool | False | Stream output | | `stream` | Optional[bool] | False | Stream output |
| `stop` | Optional[Union[str, List[str]]] | None | Stop sequences | | `stop` | Optional[Union[str, List[str]]] | None | Stop sequences |
| `frequency_penalty` | float | 0.0 | Frequency penalty | | `n` | Optional[int] | 1 | Number of choices requested |
| `tools` | Optional[List[dict]] | None | Tool definitions for function calling | | `presence_penalty` | Optional[float] | 0.0 | Presence penalty (-2.0 to 2.0) |
| `tool_choice` | Optional[str] | None | Tool selection mode | | `frequency_penalty` | Optional[float] | 0.0 | Frequency penalty (-2.0 to 2.0) |
| `logit_bias` | Optional[Dict[int, float]] | None | Per-token logit bias |
| `user` | Optional[str] | None | End-user identifier |
| `tools` | Optional[List[ToolDef]] | None | Tool definitions for function calling |
| `tool_choice` | Optional[Union[str, Dict[str, Any]]] | `"auto"` | Tool selection mode or explicit tool choice |
**Anthropic** (`MessagesRequest`):
| Param | Type | Default | Description |
|-------|------|---------|-------------|
| `model` | str | `"astrai"` | Model name returned in responses |
| `messages` | List[AnthropicMessage] | required | User/assistant messages |
| `system` | Optional[str] | None | System prompt |
| `max_tokens` | int | 1024 | Max generation length |
| `temperature` | Optional[float] | 1.0 | Sampling temperature (0.0-2.0) |
| `top_p` | Optional[float] | 1.0 | Nucleus threshold (0.0-1.0) |
| `top_k` | Optional[int] | 50 | Top-k count |
| `stream` | Optional[bool] | False | Stream output |
| `stop_sequences` | Optional[List[str]] | None | Stop sequences |
**Engine** (`GenerationRequest`):
| Param | Type | Default | Description |
|-------|------|---------|-------------|
| `messages` | List[Dict[str, str]] | required | Messages to format before generation |
| `top_k` | int | 50 | Top-k count; 0 disables filtering |
| `top_p` | float | 1.0 | Nucleus threshold |
| `temperature` | float | 1.0 | Sampling temperature; 0 enables greedy decoding |
| `max_tokens` | Optional[int] | None | Max generation length |
| `frequency_penalty` | float | 0.0 | Frequency penalty (-2.0 to 2.0) |
| `rep_window` | int | 64 | Recent-token window used by the frequency penalty |
| `stream` | bool | False | Stream output |
### SSE Streaming Format ### SSE Streaming Format
@@ -240,6 +277,8 @@ data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":
event: message_stop event: message_stop
data: {"type":"message_stop"} data: {"type":"message_stop"}
data: [DONE]
``` ```
### Error Responses ### Error Responses
+37 -30
View File
@@ -13,9 +13,11 @@
| Parameter | Description | Default | | Parameter | Description | Default |
|-----------|-------------|---------| |-----------|-------------|---------|
| `--config`, `-c` | YAML config file; explicit CLI options override YAML values | None |
| `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`, `online_grpo`, `online_dpo`) | required | | `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`, `online_grpo`, `online_dpo`) | required |
| `--data_root_path` | Dataset root directory | required | | `--data_root_path` | Dataset root directory | required |
| `--param_path` | Model parameters or checkpoint path | required | | `--param_path` | Model parameters or checkpoint path | required |
| `--resume` | Resume training from `--param_path` | False |
| `--n_epoch` | Total training epochs | 1 | | `--n_epoch` | Total training epochs | 1 |
| `--batch_per_device` | Batch size per device | 1 | | `--batch_per_device` | Batch size per device | 1 |
| `--grad_accum_steps` | Gradient accumulation steps between optimizer steps | 1 | | `--grad_accum_steps` | Gradient accumulation steps between optimizer steps | 1 |
@@ -26,7 +28,7 @@
|-----------|-------------|---------| |-----------|-------------|---------|
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 | | `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 | | `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
| `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | 1.0 | | `--max_grad_norm` | Maximum gradient norm for clipping; the current CLI requires a positive number | 1.0 |
### Optimizer ### Optimizer
@@ -36,9 +38,9 @@ non-matrix parameters through **AdamW** (`fused=True`).
| Parameter | Description | Default | | Parameter | Description | Default |
|-----------|-------------|---------| |-----------|-------------|---------|
| `--optimizer` | Built-in optimizer (`muon_adamw`, `nora_nadamw`, `mano_adamw`) | `muon_adamw` | | `--optimizer` | Built-in optimizer (`muon_adamw`, `nora_nadamw`, `mano_adamw`) | `muon_adamw` |
| `--weight_decay` | Weight decay (applied to Muon matrix params; non-matrix use 0) | 0.1 | | `--weight_decay` | Weight decay for optimizer parameter groups that are eligible for decay | 0.1 |
| `--muon_momentum` | Muon momentum factor | 0.95 | | `--muon_momentum` | Muon momentum factor | 0.95 |
| `--muon_nesterov` | Enable Nesterov momentum for Muon | True | | `--muon_nesterov`, `--no-muon_nesterov` | Enable or disable Nesterov momentum for Muon | enabled |
| `--muon_ns_steps` | Newton-Schulz iteration steps for Muon | 5 | | `--muon_ns_steps` | Newton-Schulz iteration steps for Muon | 5 |
| `--muon_adjust_lr` | Muon LR adjustment strategy (`original`, `match_rms_adamw`) | `match_rms_adamw` | | `--muon_adjust_lr` | Muon LR adjustment strategy (`original`, `match_rms_adamw`) | `match_rms_adamw` |
@@ -56,15 +58,18 @@ under DTensor sharding and rejects layouts sharded along the last dimension.
| `--nora_weight_decay` | Nora matrix weight decay | 0.0 | | `--nora_weight_decay` | Nora matrix weight decay | 0.0 |
`mano_adamw` routes internal `Linear.weight` matrices to **Mano** (manifold `mano_adamw` routes internal `Linear.weight` matrices to **Mano** (manifold
normalized optimizer) and the remaining parameters to **NAdamW**. Mano projects normalized optimizer) and the remaining parameters to **AdamW**. Mano projects
the momentum onto the tangent space of the Oblique manifold and normalizes it, the momentum onto the tangent space of the Oblique manifold and normalizes it,
alternating the projection axis (row/column) each step — replacing Muon's alternating the projection axis (row/column) each step — replacing Muon's
Newton-Schulz iteration with a cheaper normalization. Newton-Schulz iteration with a cheaper normalization.
| Parameter | Description | Default | | Parameter | Description | Default |
|-----------|-------------|---------| |-----------|-------------|---------|
| `--mano_momentum` | Mano momentum factor | 0.95 | | `--mano_momentum` | Accepted by the CLI but currently ignored by optimizer construction | 0.95 |
| `--mano_nesterov` | Enable Nesterov momentum for Mano | True | | `--mano_nesterov`, `--no-mano_nesterov` | Accepted by the CLI but currently ignored by optimizer construction | enabled |
The two Mano-specific flags are reserved for future wiring; do not rely on them
to change optimizer behavior in the current release.
Optimizer identity and hyperparameters are saved in checkpoint metadata. Optimizer Optimizer identity and hyperparameters are saved in checkpoint metadata. Optimizer
states are intentionally not interchangeable: resume older MuonAdamW checkpoints states are intentionally not interchangeable: resume older MuonAdamW checkpoints
@@ -78,7 +83,7 @@ with `--optimizer=muon_adamw`.
| `--stride` | Stride for sliding window over sequences | None | | `--stride` | Stride for sliding window over sequences | None |
| `--random_seed` | Random seed for reproducibility | 3407 | | `--random_seed` | Random seed for reproducibility | 3407 |
| `--num_workers` | DataLoader worker processes | 4 | | `--num_workers` | DataLoader worker processes | 4 |
| `--no_pin_memory` | Disable pin_memory (enabled by default) | (flag) | | `--pin_memory`, `--no-pin_memory` | Enable or disable DataLoader pinned memory | enabled |
### Checkpoint & Resume ### Checkpoint & Resume
@@ -100,14 +105,20 @@ with `--optimizer=muon_adamw`.
| Parameter | Description | Default | | Parameter | Description | Default |
|-----------|-------------|---------| |-----------|-------------|---------|
| `--log_dir` | Directory for metric logs | checkpoint/logs | | `--metrics` | Repeatable metric option (for example, `--metrics loss --metrics lr --metrics val_loss`) | `loss`, `lr`, `grad_norm`, `grad_snr` |
| `--metrics` | Metrics to log (e.g. --metrics loss lr val_loss) | ["loss", "lr", "grad_norm"] |
### Gradient Checkpointing ### Gradient Checkpointing
| Parameter | Description | Default | | Parameter | Description | Default |
|-----------|-------------|---------| |-----------|-------------|---------|
| `--gradient_checkpointing` | Enable activation checkpointing for DecoderBlock modules | False | | `--gradient_checkpointing`, `--no-gradient_checkpointing` | Enable or disable activation checkpointing for DecoderBlock modules | disabled |
### Miscellaneous
| Parameter | Description | Default |
|-----------|-------------|---------|
| `--compile` | Enable `torch.compile` with mode `default`, `reduce-overhead`, or `max-autotune`; omit to disable | None |
| `--dry-run` | Validate the merged configuration and print the training plan without training | False |
### Distributed Training ### Distributed Training
@@ -120,21 +131,25 @@ with `--optimizer=muon_adamw`.
| `--backend` | Distributed training backend | nccl | | `--backend` | Distributed training backend | nccl |
| `--master_addr` | Master node address | localhost | | `--master_addr` | Master node address | localhost |
| `--master_port` | Master node port | 29500 | | `--master_port` | Master node port | 29500 |
| `--tp_size` | Reserved tensor-parallel size; accepted but currently ignored | None |
### Strategy-specific ### Strategy-specific
| 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`, `online_dpo` |
| `--label_smoothing` | Label smoothing for cross-entropy loss | 0.0 | `seq`, `sft` | | `--label_smoothing` | Label smoothing for cross-entropy loss | 0.0 | `seq`, `sft` |
| `--group_size` | GRPO group size | 4 | `grpo` | | `--group_size` | GRPO/rollout group size | 4 | `grpo`, `online_grpo`, `online_dpo` |
| `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo` | | `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo`, `online_grpo` |
| `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo` | | `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo`, `online_grpo` |
| `--neftune_alpha` | NEFTune noise alpha (0=disabled, typical: 5.0) | 0.0 | `sft` | | `--neftune_alpha` | NEFTune noise alpha (0=disabled, typical: 5.0) | 0.0 | `sft` |
### Online Rollout ### Online Rollout
These options apply to `online_grpo` and `online_dpo`. Online strategies require `online_grpo` and `online_dpo` are factory aliases for the existing `grpo` and
`dpo` strategy classes; online behavior is enabled by rollout components rather
than separate strategy subclasses. These options apply to the online aliases.
Online strategies require
a `BaseRewardModel` factory in `TrainConfig`; `train.py` does not currently a `BaseRewardModel` factory in `TrainConfig`; `train.py` does not currently
provide a command-line option for configuring one. provide a command-line option for configuring one.
@@ -151,7 +166,7 @@ provide a command-line option for configuring one.
| Parameter | Description | Default | | Parameter | Description | Default |
|-----------|-------------|---------| |-----------|-------------|---------|
| `--schedule_type` | LR scheduler type (`cosine`, `sgdr`, `wsd`) | cosine | | `--schedule_type` | LR scheduler type (`cosine`, `sgdr`, `wsd`) | cosine |
| `--min_rate` | Minimum LR as fraction of base LR | None (scheduler default: 0.05 for cosine/SGDR, 0.0 for WSD) | | `--min_rate` | Minimum LR as fraction of base LR | None (all current schedulers use their effective default of 0.01) |
| `--cycle_length` | SGDR first cycle length in steps | None (total_steps - warmup_steps) | | `--cycle_length` | SGDR first cycle length in steps | None (total_steps - warmup_steps) |
| `--t_mult` | SGDR cycle length multiplier per restart | 2 | | `--t_mult` | SGDR cycle length multiplier per restart | 2 |
| `--stable_steps` | WSD stable plateau steps | None (80% of post-warmup steps) | | `--stable_steps` | WSD stable plateau steps | None (80% of post-warmup steps) |
@@ -204,14 +219,6 @@ python scripts/tools/server.py --param_path ./params --device cuda --dtype bfloa
See [Inference Guide](inference.md) for HTTP API documentation. See [Inference Guide](inference.md) for HTTP API documentation.
# Preprocess
```bash
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c config.json
```
See [Preprocessing Guide](preprocessing.md) for config file format and examples.
## Generate (`generate.py`) ## Generate (`generate.py`)
| Parameter | Type | Default | Description | | Parameter | Type | Default | Description |
@@ -221,13 +228,12 @@ See [Preprocessing Guide](preprocessing.md) for config file format and examples.
| `--output_json_file` | str | required | Path to the output JSONL file | | `--output_json_file` | str | required | Path to the output JSONL file |
| `--question_key` | str | `question` | Key for the question in input JSON | | `--question_key` | str | `question` | Key for the question in input JSON |
| `--response_key` | str | `response` | Key for the response in output JSON | | `--response_key` | str | `response` | Key for the response in output JSON |
| `--temperature` | float | `0.60` | Sampling temperature | | `--temperature` | float | `0.8` | Sampling temperature |
| `--top_k` | int | `30` | Top-k filtering | | `--top_k` | int | `50` | Top-k filtering |
| `--top_p` | float | `0.95` | Nucleus sampling threshold | | `--top_p` | float | `0.95` | Nucleus sampling threshold |
| `--batch_size` | int | `1` | Batch size for generation | | `--batch_size` | int | `1` | Batch size for generation |
| `--num_samples` | int | `1` | Responses per prompt | | `--num_samples` | int | `1` | Responses per prompt |
| `--max_tokens` | int | model config `max_position_embeddings` | Maximum tokens to generate | | `--max_seq_len` | int | `2048` | KV cache sequence length |
| `--cache_len` | int | `2048` | KV cache length |
| `--frequency_penalty` | float | `0.0` | Frequency penalty | | `--frequency_penalty` | float | `0.0` | Frequency penalty |
| `--rep_window` | int | `64` | Window size for frequency penalty | | `--rep_window` | int | `64` | Window size for frequency penalty |
@@ -243,14 +249,15 @@ python scripts/tools/generate.py \
| Parameter | Type | Default | Description | | Parameter | Type | Default | Description |
|-----------|------|---------|-------------| |-----------|------|---------|-------------|
| `input_files` | path(s) | required | Input JSONL file(s), supports glob (`data/*.jsonl`) | | `input_files` | path(s) | required | One or more existing `.jsonl` or `.json` paths. Wildcards work only when expanded by the invoking shell; the CLI does not expand globs itself. |
| `--output_dir`, `-o` | path | required | Output directory for processed data | | `--output_dir`, `-o` | path | required | Output directory for processed data |
| `--config`, `-c` | path | required | Preprocessing pipeline config (JSON) | | `--config`, `-c` | path | required | Preprocessing pipeline config (JSON) |
| `--tokenizer_path` | str | `params` | Path to tokenizer directory | | `--tokenizer_path` | str | `params` | Path to tokenizer directory |
| `--batch_size` | int | config value (`256` by default) | Override records processed per batch; must be at least 1 |
Usage: Usage:
```bash ```bash
python scripts/tools/preprocess.py data/*.jsonl -o output/ -c sft.json python scripts/tools/preprocess.py data/part-000.jsonl data/part-001.jsonl -o output/ -c sft.json
``` ```
See [Preprocessing Guide](preprocessing.md) for config file format and examples. See [Preprocessing Guide](preprocessing.md) for config file format and examples.
+60 -11
View File
@@ -10,6 +10,7 @@ Declarative JSON-driven data preprocessing. `MaskBuilderFactory` supports three
- [Configuration Reference](#configuration-reference) — all fields - [Configuration Reference](#configuration-reference) — all fields
- [Mask Algorithm](#mask-algorithm) - [Mask Algorithm](#mask-algorithm)
- [Output Layout](#output-layout) - [Output Layout](#output-layout)
- [Training Compatibility](#training-compatibility)
- [CLI](#cli) - [CLI](#cli)
- [Python API](#python-api) - [Python API](#python-api)
@@ -40,7 +41,7 @@ A single config file captures the entire pipeline, reusable and version-controll
| Field | Type | Default | Description | | Field | Type | Default | Description |
|-------|------|---------|-------------| |-------|------|---------|-------------|
| `field` | str | -- | JSONL key to read | | `field` | str | -- | JSONL key to read |
| `action` | str | -- | `"train"` / `"mask"` / `"$role"` | | `action` | str | -- | `"train"` / `"mask"` / `"$role"` / `"value"`; `"value"` copies raw values without tokenization |
| `template` | bool | `false` | Apply `chat_template` per message | | `template` | bool | `false` | Apply `chat_template` per message |
| `add_special_tokens` | bool | `true` for first non-template section | Add special tokens during encode | | `add_special_tokens` | bool | `true` for first non-template section | Add special tokens during encode |
@@ -89,7 +90,7 @@ Config:
} }
``` ```
Output keys: `sequence` (int32), `loss_mask` (bool) Output keys: `sequence` (int32), `loss_mask` (bool), `position_ids` (int32)
### SFT Instruction ### SFT Instruction
@@ -116,7 +117,7 @@ Config:
} }
``` ```
Output keys: `sequence`, `loss_mask` Output keys: `sequence`, `loss_mask`, `position_ids`
### Pretrain ### Pretrain
@@ -142,7 +143,7 @@ Config:
} }
``` ```
Output keys: `sequence` (no `loss_mask` — all tokens trained) Output keys: `sequence`, `position_ids` (no `loss_mask` — all tokens trained)
### DPO ### DPO
@@ -180,6 +181,11 @@ Config:
Output keys: `chosen`, `chosen_mask`, `rejected`, `rejected_mask` Output keys: `chosen`, `chosen_mask`, `rejected`, `rejected_mask`
The offline `Pipeline` can construct these keys, but its `.bin` output is not
currently loadable for DPO training because the writer does not preserve
per-record offsets. Train DPO directly from raw JSONL instead; see
[Training Compatibility](#training-compatibility).
### GRPO ### GRPO
Input JSONL: Input JSONL:
@@ -228,6 +234,11 @@ Output keys: `prompts`, `prompts_mask`, `responses`, `masks`, `rewards` (float32
- `mask_key: "masks"` — rename the auto-generated mask key (default: `responses_mask`) - `mask_key: "masks"` — rename the auto-generated mask key (default: `responses_mask`)
- `prompts_mask` is auto-generated (all masked) and unused by GRPOStrategy - `prompts_mask` is auto-generated (all masked) and unused by GRPOStrategy
The offline `Pipeline` flattens GRPO response groups for `.bin` output without
preserving their boundaries, and there is no automatic raw-JSONL GRPO processor
in `DatasetFactory`. See
[Training Compatibility](#training-compatibility) for the supported routes.
--- ---
## Configuration Reference ## Configuration Reference
@@ -257,7 +268,7 @@ When `sources` is set, `sections` is ignored.
| `max_chars` | int | `2000000` | Skip text-mode items longer than this | | `max_chars` | int | `2000000` | Skip text-mode items longer than this |
| `max_items` | int or null | `null` | Stop after N documents | | `max_items` | int or null | `null` | Stop after N documents |
| `batch_size` | int | `256` | Records per tokenization batch | | `batch_size` | int | `256` | Records per tokenization batch |
| `packing_strategy` | str | `"simple"` | Packing strategy: `"simple"`, `"bfd"`, `"bfd_split"` | | `packing_strategy` | str | `"simple"` | Packing is supported for single-output data with a `sequence` key: `"simple"`, `"bfd"`, or `"bfd_split"`. Multi-output DPO/GRPO data is not record-preserving packed output. |
| `max_packed_len` | int | `8192` | Maximum length of a packed bin | | `max_packed_len` | int | `8192` | Maximum length of a packed bin |
| `truncation_mode` | str | `"keep_start"` | How to truncate sequences: `"keep_start"` or `"keep_end"` | | `truncation_mode` | str | `"keep_start"` | How to truncate sequences: `"keep_start"` or `"keep_end"` |
@@ -266,8 +277,8 @@ When `sources` is set, `sections` is ignored.
| Field | Type | Default | Description | | Field | Type | Default | Description |
|-------|------|---------|-------------| |-------|------|---------|-------------|
| `domain_key` | str or null | `null` | JSONL key for domain grouping | | `domain_key` | str or null | `null` | JSONL key for domain grouping |
| `storage_format` | str | `"bin"` | `"bin"` (mmap). Reading also supports `"jsonl"` for on-the-fly tokenization | | `storage_format` | str | `"bin"` | Pipeline output format. Only `"bin"` has a registered writer; `"jsonl"` is accepted by config validation but cannot be emitted by `Pipeline`. |
| `max_tokens_per_shard` | int | `100000000` | Flush threshold in cumulative tokens | | `max_tokens_per_shard` | int | `100000000` | Flush threshold counted from each record's primary flat sequence: `sequence` for single-output data, otherwise the first flat source output |
| `dtype` | dict[str, str] | `{}` | Per-key tensor dtype override (e.g. `{"loss_mask": "bool"}`) | | `dtype` | dict[str, str] | `{}` | Per-key tensor dtype override (e.g. `{"loss_mask": "bool"}`) |
| `position_ids_mode` | str | `"doc_reset"` | How to compute position_ids: `"none"`, `"doc_reset"`, `"continuous"` | | `position_ids_mode` | str | `"doc_reset"` | How to compute position_ids: `"none"`, `"doc_reset"`, `"continuous"` |
@@ -304,11 +315,13 @@ output/
meta.json meta.json
sequence.bin sequence.bin
loss_mask.bin loss_mask.bin
position_ids.bin
wiki/ wiki/
shard_0000/ shard_0000/
meta.json meta.json
sequence.bin sequence.bin
loss_mask.bin loss_mask.bin
position_ids.bin
``` ```
### Multi-Shard (`bin`) ### Multi-Shard (`bin`)
@@ -322,13 +335,44 @@ output/
meta.json meta.json
sequence.bin sequence.bin
loss_mask.bin loss_mask.bin
position_ids.bin
shard_0001/ shard_0001/
meta.json meta.json
sequence.bin sequence.bin
loss_mask.bin loss_mask.bin
position_ids.bin
``` ```
For `bin` format, `MmapStore` discovers all shards under the domain directory via `rglob("meta.json")`. For `h5` format, `H5Store` discovers `.h5`/`.hdf5` files via recursive glob. `MmapStore` discovers binary shards recursively through their `meta.json` files.
Each shard's metadata is a top-level object keyed by tensor name:
```json
{
"sequence": {"shape": [123456], "dtype": "int32"},
"loss_mask": {"shape": [123456], "dtype": "bool"},
"position_ids": {"shape": [123456], "dtype": "int32"}
}
```
An optional `offsets` array may appear for record-oriented binary data written
through `save_bin(..., record_keys=...)`; the preprocessing `BinWriter` does not
currently request those offsets.
---
## Training Compatibility
| Training type | Supported input route |
|---------------|-----------------------|
| `seq` | Offline preprocessed `.bin`, or raw `.jsonl` eagerly transformed by `JsonlStore` using `dataset_config.json` or the built-in `messages` config |
| `sft` | Offline preprocessed `.bin`, or raw `.jsonl` through the same eager transform routes |
| `dpo` | Raw `.jsonl` through the automatic lazy DPO processor selected by `DatasetFactory` when `tokenizer_path` is supplied, or a caller-provided record store |
| `grpo` | A caller-provided, already-loaded `Store` with record-shaped `prompts`, `responses`, `masks`, and `rewards`; no automatic raw-JSONL processor is currently wired |
Offline DPO and GRPO preprocessing configs describe the intended token fields,
but their `.bin` output is not currently loadable for training. DPO binary
shards lack per-record offsets. GRPO response groups are flattened before the
binary writer and their record/group boundaries are not preserved.
--- ---
@@ -336,15 +380,20 @@ For `bin` format, `MmapStore` discovers all shards under the domain directory vi
```bash ```bash
# SFT # SFT
python scripts/tools/preprocess.py data/sft/*.jsonl -o output/sft/ -c configs/sft_chat.json python scripts/tools/preprocess.py data/sft/part-000.jsonl -o output/sft/ -c configs/sft_chat.json --batch_size 128
# DPO # DPO
python scripts/tools/preprocess.py data/dpo/*.jsonl -o output/dpo/ -c configs/dpo.json --tokenizer_path params python scripts/tools/preprocess.py data/dpo/part-000.jsonl -o output/dpo/ -c configs/dpo.json --tokenizer_path params
# GRPO # GRPO
python scripts/tools/preprocess.py data/grpo/*.jsonl -o output/grpo/ -c configs/grpo.json python scripts/tools/preprocess.py data/grpo/part-000.jsonl -o output/grpo/ -c configs/grpo.json
``` ```
Inputs may be `.jsonl` files or `.json` files containing one object or a list of
objects. Each positional path must exist. A wildcard such as `data/*.jsonl`
works only when the invoking shell expands it before Click receives the
arguments; otherwise pass the files explicitly.
--- ---
## Python API ## Python API
+18 -12
View File
@@ -41,7 +41,10 @@ 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 $$ $$ 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 $$
`RotaryEmbedding` pre-computes `cos_table` and `sin_table` (f32, `[max_len, dim/2]`). `forward()` returns a `(cos, sin)` tuple indexed by `position_ids`. `apply_rotary_emb` applies the rotation: during training it uses torch complex multiply (autograd-compatible); during inference it auto-dispatches to a fused CUDA kernel when available. `RotaryEmbedding` pre-computes a complex `freqs_cis` buffer. `forward()` returns
a tensor indexed by `position_ids`. `apply_rotary_emb` applies the rotation:
during training it uses torch complex multiply (autograd-compatible); during
inference it auto-dispatches to a fused CUDA kernel when available.
## Training Loop ## Training Loop
@@ -52,8 +55,8 @@ on_train_begin
model.train() model.train()
on_epoch_begin on_epoch_begin
for batch in dataloader: for batch in dataloader:
on_batch_begin
with executor.accumulate(model): with executor.accumulate(model):
on_batch_begin
loss_output = strategy(batch) loss_output = strategy(batch)
context.loss = loss_output["loss"].item() context.loss = loss_output["loss"].item()
context.metrics = loss_output["metrics"] context.metrics = loss_output["metrics"]
@@ -67,6 +70,7 @@ on_train_begin
if executor.sync_gradients: if executor.sync_gradients:
on_optimizer_step on_optimizer_step
optimizer.step() optimizer.step()
strategy.on_optimizer_step()
optimizer.zero_grad() optimizer.zero_grad()
if scheduler: if scheduler:
scheduler.step() scheduler.step()
@@ -78,16 +82,16 @@ on_train_end
| Hook | Fires | Default callback | | Hook | Fires | Default callback |
|------|-------|-----------------| |------|-------|-----------------|
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback` | | `on_train_begin` | Before training starts | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |
| `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` | | `on_epoch_begin` | Start of each epoch | `ProgressBarCallback` |
| `on_batch_begin` | Every batch | — | | `on_batch_begin` | Every batch | — |
| `on_optimizer_step` | Every accumulation window | `GradientClippingCallback`, `MetricCallback`, `ProgressBarCallback` | | `on_optimizer_step` | Every accumulation window | `MetricCallback`, `ProgressBarCallback`, `GradientClippingCallback` |
| `on_batch_end` | Every batch | `CheckpointCallback` | | `on_batch_end` | Every batch | `CheckpointCallback` |
| `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` | | `on_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` |
| `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` | | `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` |
| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` | | `on_train_end` | Training exits after `on_train_begin` completes (via `finally`) | `GradientCheckpointingCallback`, `CheckpointCallback`, `MetricCallback` |
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm), `gradient_clipping` (always registered; computes grad norm, clips only when `max_grad_norm` is not `None`). Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric` (JSONL + validation, rank-0), `progress_bar` (tqdm, rank-0), `gradient_clipping`. The gradient-clipping callback is always registered and always calls `executor.clip_grad_norm()` with the numeric `max_grad_norm` value.
Strategies return `{"loss": Tensor, "metrics": Dict[str, float]}` when called by the trainer. Built-in metrics include the task-specific loss and, for MoE models, `moe_aux_loss` plus `moe_aux_loss_weighted`. Direct `compute_loss(batch)` calls continue to return a single loss tensor. Strategies return `{"loss": Tensor, "metrics": Dict[str, float]}` when called by the trainer. Built-in metrics include the task-specific loss and, for MoE models, `moe_aux_loss` plus `moe_aux_loss_weighted`. Direct `compute_loss(batch)` calls continue to return a single loss tensor.
@@ -98,7 +102,7 @@ Strategies return `{"loss": Tensor, "metrics": Dict[str, float]}` when called by
Next-token cross-entropy with optional label smoothing: Next-token cross-entropy with optional label smoothing:
$$ $$
L_{\text{PT}} = -\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta) L_{\text{PT}} = -\frac{1}{T}\sum_{t=1}^{T} \log P(x_t \mid x_{\lt t}; \theta)
$$ $$
Keys: `input_ids`, `target_ids`. Optional: `label_smoothing`. Keys: `input_ids`, `target_ids`. Optional: `label_smoothing`.
@@ -108,7 +112,7 @@ Keys: `input_ids`, `target_ids`. Optional: `label_smoothing`.
Masked cross-entropy (`ignore_index=-100`) over response tokens: 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) L_{\text{SFT}} = -\frac{1}{L}\sum_{t=P+1}^{P+L} \log P(s_t \mid s_{\lt t}; \theta)
$$ $$
Keys: `input_ids`, `target_ids`, `loss_mask`, `position_ids`. Optional: `label_smoothing`. Keys: `input_ids`, `target_ids`, `loss_mask`, `position_ids`. Optional: `label_smoothing`.
@@ -168,9 +172,9 @@ model factory.
|------|-------|-------------| |------|-------|-------------|
| Cosine | `CosineScheduler` | Linear warmup → cosine decay to `min_rate` | | Cosine | `CosineScheduler` | Linear warmup → cosine decay to `min_rate` |
| SGDR | `SGDRScheduler` | Cosine annealing with warm restarts (`t_mult=2`) | | SGDR | `SGDRScheduler` | Cosine annealing with warm restarts (`t_mult=2`) |
| WSD | `WSDScheduler` | Warmup-Stable-Decay with sqrt cooldown | | WSD | `WSDScheduler` | Warmup-Stable-Decay with quadratic decay |
Created by `SchedulerFactory.create(schedule_type, optimizer, **kwargs)`. Valid types: `"cosine"`, `"sgdr"`, `"wsd"`. Omit to use no scheduler. Created by `SchedulerFactory.create(schedule_type, optimizer, **kwargs)`. Valid types: `"cosine"`, `"sgdr"`, `"wsd"`. The training CLI always creates a scheduler and defaults `--schedule_type` to `"cosine"`.
## Gradient Checkpointing ## Gradient Checkpointing
@@ -188,10 +192,12 @@ Callback wraps each `DecoderBlock.forward` with `torch.utils.checkpoint.checkpoi
``` ```
Checkpoint(state_dict, epoch, consumed_samples, extra, meta, config) Checkpoint(state_dict, epoch, consumed_samples, extra, meta, config)
├── save(save_dir) rank-0 only: meta.json (epoch/consumed_samples/timestamp) + config.json (model config) + model.safetensors + optional {key}.pt (optimizer.pt, scheduler.pt) ├── save(save_dir) meta.json (epoch/consumed_samples/timestamp) + config.json (model config) + model.safetensors + optional {key}.pt (optimizer.pt, scheduler.pt)
└── load(save_dir, broadcast=False) loads from local disk; set broadcast=True to broadcast metadata from rank-0 └── load(save_dir, broadcast=False) loads from local disk; set broadcast=True to broadcast metadata from rank-0
``` ```
`Checkpoint.save()` writes whenever it is called. During training, `CheckpointCallback` uses the executor checkpoint context so only rank 0 receives a state dict and calls `save()`.
Optimizer/scheduler state persisted by default via `Checkpoint.extra`. Optimizer/scheduler state persisted by default via `Checkpoint.extra`.
Model config (`context.model_config`) saved into `config.json` during training via `CheckpointCallback`. Model config (`context.model_config`) saved into `config.json` during training via `CheckpointCallback`.
@@ -235,4 +241,4 @@ nohup python scripts/tools/train.py \
Full parameter reference at [params.md](params.md). Full parameter reference at [params.md](params.md).
> Document Update Time: 2026-07-31 > Document Update Time: 2026-08-02