From 288ba20db1f1d4b642f526b74db96d771f496fd6 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Sun, 2 Aug 2026 07:39:24 +0800 Subject: [PATCH] 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 --- CONTRIBUTING.md | 16 +-- README.md | 8 +- docs/README-zh-CN.md | 12 +- docs/developer/architecture.md | 214 +++++++++++++++++---------------- docs/developer/dataflow.md | 114 ++++++++++++------ docs/developer/internals.md | 28 +++-- docs/get-started.md | 25 +++- docs/guides/distributed.md | 42 ++++--- docs/guides/evaluation.md | 58 +++++++-- docs/guides/inference.md | 65 ++++++++-- docs/guides/params.md | 67 ++++++----- docs/guides/preprocessing.md | 71 +++++++++-- docs/guides/training.md | 30 +++-- 13 files changed, 483 insertions(+), 267 deletions(-) diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 683c508..6a3b957 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -20,9 +20,6 @@ Run the following checks **in order** — CI will reject if any fail. ruff format . ``` -> **Note**: `ruff format` may rename parameters (e.g. `mask` → `attn_mask`). -> Always review the diff after formatting. - ### 2. Import sorting ```bash @@ -44,7 +41,7 @@ python -u -m pytest tests/ -v > Failed tests may leave orphan tempdirs under `%TEMP%`. Clean them manually if needed. -### 4. (Optional) Full pre-commit check +### 4. (Optional) Full pre-commit check script If you have Git Bash available: @@ -52,12 +49,17 @@ If you have Git Bash available: 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 ``` -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) ``` @@ -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 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 | ## Submitting Changes diff --git a/README.md b/README.md index e391e6e..506b363 100644 --- a/README.md +++ b/README.md @@ -56,6 +56,8 @@ End-to-end walkthrough in 5 steps: **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 git clone https://github.com/ViperEkura/AstrAI.git cd AstrAI @@ -132,7 +134,7 @@ Check out the demos in the `scripts/demo/` folder: # Download model weights (required before running demos) 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 # 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 up -d -# Docker Compose (CPU only) +# Docker Compose CPU server profile (CUDA-only generation scripts/demos are unavailable) docker compose --profile cpu up -d ``` @@ -256,4 +258,4 @@ This project is licensed under the [GPL-3.0 License](LICENSE).
A lightweight Transformer framework designed for both high performance and ease of use. -
\ No newline at end of file + diff --git a/docs/README-zh-CN.md b/docs/README-zh-CN.md index df638dd..09b8879 100644 --- a/docs/README-zh-CN.md +++ b/docs/README-zh-CN.md @@ -62,6 +62,8 @@ **1. 安装** +AstrAI 需要 Python 3.12+,并精确固定 PyTorch 版本为 `2.11.0`。训练、`scripts/tools/generate.py`、生成式评估和生成演示需要 CUDA;CPU 支持仅适用于提供明确 CPU 设备路径的组件,例如 HTTP 服务和直接打分评估。 + ```bash git clone https://github.com/ViperEkura/AstrAI.git cd AstrAI @@ -138,7 +140,7 @@ curl http://localhost:8000/v1/chat/completions \ # 下载模型权重(运行演示前必需) python scripts/demo/download.py # model → params/ -# 交互式流式聊天(多轮对话,保持历史记录) +# 单轮交互式流式提示循环(不保留对话历史) python scripts/demo/stream_chat.py # 在 >> 后输入消息,输入 !exit 退出 @@ -189,7 +191,7 @@ docker run --gpus all -v /path/to/data:/data -it astrai:latest # Docker Compose(GPU,默认) docker compose up -d -# Docker Compose(仅 CPU) +# Docker Compose CPU 服务配置(不支持仅限 CUDA 的生成脚本和演示) docker compose --profile cpu up -d ``` @@ -239,7 +241,7 @@ SSE 流式格式、错误码和统计端点详见[推理文档](guides/inference ### 贡献 -我们欢迎贡献!请参阅[贡献指南](../../CONTRIBUTING.md)了解详情。 +我们欢迎贡献!请参阅[贡献指南](../CONTRIBUTING.md)了解详情。 1. Fork 本仓库。 2. 创建功能分支。 @@ -256,10 +258,10 @@ SSE 流式格式、错误码和统计端点详见[推理文档](guides/inference ### 许可证 -本项目采用 [GPL-3.0 许可证](../../LICENSE)。 +本项目采用 [GPL-3.0 许可证](../LICENSE)。 ---
专为高性能与易用性设计的轻量级 Transformer 框架。 -
\ No newline at end of file + diff --git a/docs/developer/architecture.md b/docs/developer/architecture.md index be76761..cae26f1 100644 --- a/docs/developer/architecture.md +++ b/docs/developer/architecture.md @@ -4,7 +4,7 @@ - [Class Diagram](#class-diagram) — Full Mermaid class diagram across 10+ namespaces - [Module Overview](#module-overview) — Component inventory per module -- [Design Patterns](#design-patterns) — 13 documented patterns with classes +- [Design Patterns](#design-patterns) — 15 documented patterns with classes - [Core Relationships](#core-relationships) — 11 key inter-component relationships ## Class Diagram @@ -49,6 +49,11 @@ classDiagram +Optional[int] n_shared_experts +Optional[int] n_activated_experts +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 { @@ -63,6 +68,7 @@ classDiagram +Optional[int] num_attention_heads +Optional[int] num_key_value_heads +Optional[bool] use_qk_norm + +Optional[bool] use_gated_attention +str ffn_type +Optional[dict] rope_scaling +Optional[str] pooling_type @@ -114,22 +120,25 @@ classDiagram +Dataset dataset +Callable optimizer_fn +Callable scheduler_fn + +Optional[str] optimizer_name + +Dict[str, Any] optimizer_hyperparameters +int n_epoch +int batch_per_device +int grad_accum_steps +Optional[float] max_grad_norm +list gradient_checkpointing_modules + +Optional[str] compile_mode +int start_epoch +int start_samples +str ckpt_dir +int ckpt_interval - +str log_dir +List[str] metrics +Optional[LoRAConfig] lora +int random_seed +int num_workers +Optional[int] prefetch_factor +bool pin_memory + +Optional[Callable] collate_fn +int nprocs +str backend +str master_addr @@ -140,6 +149,7 @@ classDiagram +Optional[float] val_split +int val_step +float neftune_alpha + +float moe_aux_loss_coef +str parallel_mode +int rollout_interval +float rollout_temperature @@ -149,7 +159,6 @@ classDiagram +Optional[Callable] reward_model_fn +dict executor_kwargs +dict extra_kwargs - +validate() } } @@ -205,10 +214,6 @@ classDiagram -_fetch_record_key(key, index) Tensor } - class H5Store { - +load(path) - } - class MmapStore { +List _mmap_refs +load(path) @@ -260,11 +265,15 @@ classDiagram } namespace model { - class AutoModel { - +BaseModelConfig config + class ModelFactory { +Dict _entries +register(name) decorator +get_component_class(name) Type + } + + class AutoModel { + <> + +BaseModelConfig config +from_pretrained(path, disable_random_init, strict) nn.Module +save_pretrained(save_directory) +to(*args, **kwargs) Self @@ -299,7 +308,13 @@ classDiagram +RMSNorm input_norm +nn.Module mlp # MLP or DeepSeekMoE via FFNFactory +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 { + <> + +Tensor hidden_states + +Optional[Tensor] aux_loss } class GQA { @@ -314,7 +329,7 @@ classDiagram +Linear q_proj, k_proj, v_proj, o_proj +Linear gate # only if use_gated_attention +RMSNorm q_norm, k_norm # only if use_qk_norm - +forward(x, rotary_emb, attn_mask, kv_cache) Tensor + +forward(x, rotary_emb, attn_mask, kv_cache, is_causal) Tensor } class MLA { @@ -334,12 +349,18 @@ classDiagram +Linear gate # only if use_gated_attention +RMSNorm kv_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 { +Linear up, gate, down - +forward(x) Tensor + +forward(x) FFNOutput + } + + class FFNOutput { + <> + +Tensor hidden_states + +Optional[Tensor] aux_loss } class DeepSeekMoE { @@ -351,7 +372,7 @@ classDiagram +Linear router +ModuleList shared_experts +ModuleList routed_experts - +forward(x) Tensor + +forward(x) FFNOutput } class AttnFactory { @@ -380,9 +401,8 @@ classDiagram +int max_len +float base +Optional[Dict] rope_scaling - +Tensor cos_table - +Tensor sin_table - +forward(x, position_ids=None) Tuple[Tensor, Tensor] + +Tensor freqs_cis + +forward(x, position_ids=None) Tensor } class Embedding { @@ -486,10 +506,6 @@ classDiagram +save(output_dir, domain, shard_idx, tensors) } - class H5Writer { - +save(output_dir, domain, shard_idx, tensors) - } - class Pipeline { +PipelineConfig config +List[str] paths @@ -559,7 +575,7 @@ classDiagram class Trainer { +TrainConfig train_config +List[TrainCallback] callbacks - +train(resume_dir) + +train(param_path=None, resume=False) -_get_default_callbacks() List[TrainCallback] } @@ -576,13 +592,17 @@ classDiagram +int epoch +int consumed_samples +float loss - +float grad_norm + +Dict[str, float] metrics + +Optional[float] grad_norm + +GradSNRTracker grad_snr_tracker +DataLoader val_dataloader - +float val_loss + +Optional[float] val_loss +int world_size +int rank +dict kwargs - +optimizer_step() int + +stop_requested (property) bool + +optimizer_step (property) int + +request_stop() } class TrainContextBuilder { @@ -594,11 +614,22 @@ classDiagram class BaseStrategy { +Callable model +Optional[BaseExecutor] executor - +Optional[Callable] model_fn + +float moe_aux_loss_coef +dict extra_kwargs +str device - +__call__(batch) Tensor + +__call__(batch) LossOutput +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 { + <> + +Tensor loss + +Dict[str, float] metrics } class StrategyFactory { @@ -636,9 +667,12 @@ classDiagram class RawRollout { +Tensor prompts + +Tensor prompt_mask +Tensor responses +Tensor response_mask +Tensor logprobs_old + +List[str] prompt_texts + +List[List[str]] response_texts } class RolloutResult { @@ -647,10 +681,18 @@ classDiagram class BaseRewardModel { <> - +score(prompts, responses) Tensor + +score(List[str] prompts, List[List[str]] responses) Tensor } 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 } @@ -740,7 +782,7 @@ classDiagram } class MetricCallback { - +Path log_dir + +Path ckpt_dir +int save_interval +List[str] metrics +int val_step @@ -764,9 +806,9 @@ classDiagram +nn.Module model +AutoTokenizer tokenizer +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_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 +shutdown() } @@ -774,18 +816,18 @@ classDiagram class Executor { +AutoModel model +AutoTokenizer tokenizer - +KVCache page_cache + +PagePool kv_cache +Optional[str] device +Optional[torch.dtype] dtype +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 { - +KVCache _page_cache + +PagePool _cache +Executor _executor +TaskManager _task_mgr - +bool _running + +Event _stop_event +Thread _loop_thread +int max_seq_len +str device @@ -795,6 +837,7 @@ classDiagram +start() +stop() +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 { @@ -816,16 +859,6 @@ classDiagram +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 { +int size +Tensor k_buffer @@ -852,8 +885,7 @@ classDiagram +Tensor seq_lens +Tensor out_cache_loc +int max_len - +Optional[Tensor] page_table - +Optional[Tensor] decode_mask + +Optional[Tensor] kv_indptr } class PagePool { @@ -878,6 +910,8 @@ classDiagram +float temperature +float top_p +int top_k + +float frequency_penalty + +int rep_window +TaskStatus status +List output_ids +int input_tokens @@ -923,27 +957,29 @@ classDiagram +float top_p +float temperature +Optional[int] max_tokens + +float frequency_penalty + +int rep_window +bool stream } class BaseSamplingStrategy { <> - +apply(logits, filter_value) Tensor + +apply(logits, filter_value, input_ids, input_mask) Tensor } class TemperatureStrategy { +float temperature - +apply(logits, filter_value) Tensor + +apply(logits, filter_value, input_ids, input_mask) Tensor } class TopKStrategy { +int top_k - +apply(logits, filter_value) Tensor + +apply(logits, filter_value, input_ids, input_mask) Tensor } class TopPStrategy { +float top_p - +apply(logits, filter_value) Tensor + +apply(logits, filter_value, input_ids, input_mask) Tensor } class FrequencyPenaltyStrategy { @@ -953,8 +989,8 @@ classDiagram class SamplingPipeline { +List[BaseSamplingStrategy] strategies - +apply(logits, filter_value) Tensor - +sample(logits, filter_value) Tensor + +apply(logits, filter_value, input_ids, input_mask) Tensor + +sample(logits, filter_value, input_ids, input_mask, return_logprobs) Union[Tensor, Tuple[Tensor, Tensor]] } class StreamDecoder { @@ -1029,7 +1065,7 @@ classDiagram <> +prepare(request, engine) Tuple[str, GenContext, 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_response(ctx, content, stop) Dict } @@ -1037,7 +1073,7 @@ classDiagram class OpenAIResponseBuilder { +prepare(request, engine) Tuple +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_response(ctx, content, stop) Dict } @@ -1045,7 +1081,7 @@ classDiagram class AnthropicResponseBuilder { +prepare(request, engine) Tuple +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_response(ctx, content, stop) Dict } @@ -1153,10 +1189,13 @@ classDiagram class BaseExecutor { +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 +backward(loss) +unwrap_model(model) dict + +checkpoint_context(model) context manager + +clip_grad_norm(model, max_norm) float + +use_distributed (property) bool +sync_gradients (property) bool +grad_accum_steps (property) int } @@ -1173,7 +1212,8 @@ classDiagram class FSDPExecutor { -_prepare_model(model) nn.Module -_no_sync(model) context manager - +unwrap_model(model) dict + +unwrap_model(model) Optional[dict] + +clip_grad_norm(model, max_norm) float } class ExecutorFactory { @@ -1182,33 +1222,6 @@ classDiagram +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 @@ -1230,11 +1243,8 @@ classDiagram BaseDataset <|-- SFTDataset BaseDataset <|-- DPODataset BaseDataset <|-- GRPODataset - Store <|-- H5Store Store <|-- MmapStore Store <|-- JsonlStore - H5Store --|> Streamable - H5Store --|> Recordable MmapStore --|> Streamable MmapStore --|> Recordable JsonlStore --|> Streamable @@ -1243,8 +1253,6 @@ classDiagram BaseSamplingStrategy <|-- TopKStrategy BaseSamplingStrategy <|-- TopPStrategy BaseSamplingStrategy <|-- FrequencyPenaltyStrategy - ParallelModel <|-- RowParallelLinear - ParallelModel <|-- ColumnParallelLinear AutoModel <|-- AutoRegressiveLM AutoModel <|-- EmbeddingEncoder BaseConfig <|-- BaseModelConfig @@ -1255,7 +1263,7 @@ classDiagram BaseConfig <|-- PipelineConfig BaseModelConfig <|-- AutoRegressiveLMConfig BaseModelConfig <|-- EncoderConfig - BaseFactory <|-- AutoModel + BaseFactory <|-- ModelFactory BaseFactory <|-- AttnFactory BaseFactory <|-- FFNFactory BaseFactory <|-- DatasetFactory @@ -1286,7 +1294,6 @@ classDiagram PositionIdStrategy <|-- DocResetPositionId PositionIdStrategy <|-- ContinuousPositionId StoreWriter <|-- BinWriter - StoreWriter <|-- H5Writer RawRollout <|-- RolloutResult LaunchStrategy <|-- TorchrunStrategy LaunchStrategy <|-- LocalStrategy @@ -1317,8 +1324,6 @@ classDiagram %% --- Aggregation (weak ownership) --- AutoModel o-- BaseModelConfig AutoTokenizer o-- ChatTemplate - PagePool o-- Allocator - PagePool o-- PrefixCache Trainer o-- TrainCallback TrainContext o-- BaseStrategy TrainContext o-- BaseScheduler @@ -1352,11 +1357,12 @@ classDiagram FFNFactory ..> DeepSeekMoE : creates DecoderBlock ..> AttnFactory : uses DecoderBlock ..> FFNFactory : uses - StoreFactory ..> H5Store : creates StoreFactory ..> MmapStore : creates StoreFactory ..> JsonlStore : creates ConfigFactory ..> AutoRegressiveLMConfig : creates ConfigFactory ..> EncoderConfig : creates + ModelFactory ..> AutoRegressiveLM : creates + ModelFactory ..> EmbeddingEncoder : creates ExecutorFactory ..> NoneExecutor : creates ExecutorFactory ..> DDPExecutor : creates ExecutorFactory ..> FSDPExecutor : creates @@ -1399,10 +1405,10 @@ classDiagram | 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.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.dataset** | BaseDataset, SEQDataset, SFTDataset, DPODataset, GRPODataset, Store, Streamable, Recordable, H5Store, MmapStore, JsonlSource, JsonlStore, StoreFactory, RDSampler, DatasetFactory | Dataset loading and management | +| **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, MmapStore, JsonlSource, JsonlStore, StoreFactory, RDSampler, DatasetFactory | Dataset loading and management | | **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.trainer** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy–GRPOStrategy, StrategyFactory, BaseScheduler–WSDScheduler, 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, BaseSamplingStrategy–SamplingPipeline, 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 | |---------|---------|---------| -| **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 | | **Strategy** | `SEQStrategy`, `SFTStrategy`, `DPOStrategy`, `GRPOStrategy` | Training strategy switching | | **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 | | **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 | -| **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 | -| **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading | +| **Model Registry** | `ModelFactory`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading | ## 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` 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 -7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (H5Store/MmapStore/JsonlStore) loads data with explicit `_length` and multi-segment `_data` -8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only), extra state saved as `{key}.pt` +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; `CheckpointCallback` performs rank-0 training saves, with extra state saved as `{key}.pt` 9. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`/`WSDScheduler` 10. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops 11. **Protocols**: `OptimizerProtocol` / `SchedulerProtocol` — structural subtyping for `AccumOptimizer` / `AccumScheduler` wrappers -> Document Update Time: 2026-07-31 +> Document Update Time: 2026-08-02 diff --git a/docs/developer/dataflow.md b/docs/developer/dataflow.md index 5fff96f..4ba066d 100644 --- a/docs/developer/dataflow.md +++ b/docs/developer/dataflow.md @@ -14,26 +14,30 @@ This document describes the data pipeline: from raw text to model input tensors. ## Overview ``` -JSONL Lines → Pipeline (mask builder) → Tokenized Tensors - ↓ - .h5 or .bin storage - ↓ - Store.load() +JSON / JSONL Records → Pipeline (mask builder) → Tokenized Tensors + ↓ + .bin storage + ↓ + Store.load() ↓ Store.fetch(begin, end, keys) ↓ - BaseDataset.__getitem__(idx) - ↓ - Sampler → DataLoader → Training / Inference + Dataset.__getitem__(idx) + ↓ + RDSampler → DataLoader → Training ``` ## 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 -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 # 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 ``` -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 `detect_format(load_path)` inspects the path: -- If `load_path` is a file: checks suffix — `.h5`/`.hdf5` → `"h5"`, `.jsonl` → `"jsonl"`, unknown suffix raises `ValueError` -- If `load_path` is a directory: recursively globs for `*.h5`/`*.hdf5` files → `"h5"`, `*.bin` + `**/meta.json` → `"bin"`, or `*.jsonl` + `dataset_config.json` → `"jsonl"` +- If `load_path` is a file: `.jsonl` selects `"jsonl"`; other suffixes raise `ValueError`. +- 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 Storage format is auto-detected by `detect_format()`; backends are dispatched via registry: ``` -StoreFactory.create("h5") → H5Store StoreFactory.create("bin") → MmapStore 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. - -**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. +Both stores inherit `Store` and compose the `Streamable` and `Recordable` +access methods. **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 | 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`) | | `dpo` | `chosen`, `rejected`, `chosen_mask`, `rejected_mask` | 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 ``` -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) - → StoreFactory.create(storage_type) - → Store.load(load_path) - → _normalize(raw) # base Store, shared by both backends - → Store._data[Dict[str, List[Tensor]]] - + _cum[Dict[str, List[int]]] (stream mode) - + _offsets[Dict[str, List[int]]] (record mode) +DatasetFactory.load(...) + → detect_format(load_path) + → optionally build dpo_processor for raw JSONL + → StoreFactory.create(storage_type, window_size, stride) + → Store.load(load_path, transform=... or processor=...) + → DatasetFactory.create(train_type, store=store) Stream datasets (SEQ/SFT): BaseDataset.__getitem__(idx) - → get_index(idx) → [begin, end) + → Store.sample_window(idx) → [begin, end) → Store.fetch(begin, end, keys) → Tensor / Dict[str, Tensor] -Record datasets (DPO/GRPO via RecordDataset): - RecordDataset.__getitem__(idx) +Record datasets (DPO/GRPO): + DPODataset/GRPODataset.__getitem__(idx) → 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`). -`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_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 -`ResumableDistributedSampler` supports checkpoint-aware distributed sampling: +`RDSampler` supports checkpoint-aware distributed sampling: - Tracks `start_epoch` / `start_iter` for resume - Shuffle via `torch.Generator(seed + epoch)` diff --git a/docs/developer/internals.md b/docs/developer/internals.md index 479e818..4a0608b 100644 --- a/docs/developer/internals.md +++ b/docs/developer/internals.md @@ -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 $$ -`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. @@ -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: -$$ 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) 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. @@ -98,8 +103,8 @@ on_train_begin model.train() on_epoch_begin for batch in dataloader: - on_batch_begin with executor.accumulate(model): + on_batch_begin loss_output = strategy(batch) context.loss = loss_output["loss"].item() context.metrics = loss_output["metrics"] @@ -113,6 +118,7 @@ on_train_begin if executor.sync_gradients: on_optimizer_step optimizer.step() + strategy.on_optimizer_step() optimizer.zero_grad() if scheduler: 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. +Strategy metrics are detached and converted to Python `float` values before the +`LossOutput` is returned; only `LossOutput.loss` remains a differentiable tensor. ## Callback Lifecycle | 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_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_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` | | `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` | -| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` | +| `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 @@ -160,7 +168,7 @@ Three-layer separation (SGLang-inspired): - **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. -`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 @@ -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. -> Document Update Time: 2026-07-31 +> Document Update Time: 2026-08-02 diff --git a/docs/get-started.md b/docs/get-started.md index 24f70ff..c60e386 100644 --- a/docs/get-started.md +++ b/docs/get-started.md @@ -2,11 +2,23 @@ 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 - **Python 3.12+** -- **PyTorch 2.11+** (CUDA 12.8 recommended for GPU support) -- NVIDIA GPU with CUDA (optional but recommended; CPU works for inference) +- **PyTorch 2.11.0** (the exact version pinned by AstrAI; CUDA 12.8 build recommended for GPU support) +- 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 @@ -55,7 +67,7 @@ python scripts/demo/stream_chat.py # 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 @@ -192,6 +204,13 @@ See [Training Guide](guides/training.md) for loss formulas and strategies. See [ ## 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 # HumanEval (code generation, auto-downloads dataset) python scripts/eval/evaluate_humaneval.py --param_path ./params --num_samples 20 diff --git a/docs/guides/distributed.md b/docs/guides/distributed.md index 8462fdc..95adb10 100644 --- a/docs/guides/distributed.md +++ b/docs/guides/distributed.md @@ -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. +## 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 ### Single GPU @@ -21,9 +33,6 @@ python scripts/tools/train.py \ ```bash export CUDA_VISIBLE_DEVICES=0,1,2,3 -export NCCL_P2P_DISABLE=1 -export NCCL_NET_GDR_LEVEL=0 - python scripts/tools/train.py \ --train_type=sft \ --param_path ./params \ @@ -38,9 +47,6 @@ python scripts/tools/train.py \ ```bash export CUDA_VISIBLE_DEVICES=0,1,2,3 -export NCCL_P2P_DISABLE=1 -export NCCL_NET_GDR_LEVEL=0 - python scripts/tools/train.py \ --train_type=sft \ --param_path ./params \ @@ -110,7 +116,7 @@ AstrAI auto-detects the launch method: | 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 | | 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 \ --train_type=sft \ --parallel_mode=ddp \ + --nprocs=4 \ --param_path ./params \ --data_root_path ./dataset \ --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 export NCCL_P2P_DISABLE=1 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 @@ -176,9 +187,6 @@ This ensures the LR schedule is correctly scaled regardless of the number of GPU ```bash export CUDA_VISIBLE_DEVICES=0,1,2,3 -export NCCL_P2P_DISABLE=1 -export NCCL_NET_GDR_LEVEL=0 - python scripts/tools/train.py \ --train_type=seq \ --param_path ./params \ @@ -240,7 +248,7 @@ python scripts/tools/train.py \ | 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` | | `--start_method` | `spawn` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) | | `--backend` | `nccl` | Distributed backend (`nccl`, `gloo`) | @@ -248,8 +256,8 @@ python scripts/tools/train.py \ | `--master_port` | `29500` | Master node port | | `--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). -> Document Update Time: 2026-07-30 +> Document Update Time: 2026-08-02 diff --git a/docs/guides/evaluation.md b/docs/guides/evaluation.md index 80481f5..2e13491 100644 --- a/docs/guides/evaluation.md +++ b/docs/guides/evaluation.md @@ -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. +## 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 | 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. - **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 \ --param_path ./params \ --num_samples 20 \ - --batch_size 32 \ + --batch_size 64 \ --max_tokens 512 \ --output results/humaneval.json ``` @@ -47,7 +78,8 @@ python scripts/eval/evaluate_humaneval.py \ | `--temperature` | 0.8 | Sampling temperature | | `--top_p` | 0.95 | Nucleus sampling threshold | | `--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_timeout` | 3.0 | Per-subprocess timeout (seconds) | | `--problems` | None | Restrict to specific problem indices | @@ -66,7 +98,7 @@ python scripts/eval/evaluate_humaneval.py \ python scripts/eval/evaluate_mmlu.py \ --param_path ./params \ --n_shot 5 \ - --subjects math_algebra history_us \ + --subjects abstract_algebra high_school_us_history \ --output results/mmlu.json ``` @@ -82,12 +114,13 @@ python scripts/eval/evaluate_mmlu.py \ | `--device` | auto | Device (`cuda` / `cpu`) | | `--dtype` | auto | `bfloat16` on CUDA, `float32` on CPU | | `--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. **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 `//` and `/dev/` (for few-shot). +**Data**: Auto-downloads `cais/mmlu` from HuggingFace. Stored as per-subject CSVs in `//` and `/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 \ --input_path data.jsonl \ --output_dir ppl_results/ \ - --batch_size 4 \ + --batch_size 64 \ --max_length 2048 ``` @@ -110,7 +143,7 @@ python scripts/eval/evaluate_ppl.py \ | `--input_path` | required | Input file, glob, or directory | | `--output_dir` | required | Output directory for `summary.json` + token JSONL | | `--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) | | `--token_level` | False | Store per-token log_probs + token-type analysis | | `--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`. -**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_k` | 50 | Top-k sampling | | `--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) | | `--dump_responses` | None | Path to dump raw responses as JSONL | @@ -232,7 +266,7 @@ python scripts/eval/analyze_weights.py \ | 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 | | `--no_svd` | False | Skip SVD; show only weight stats (faster) | | `--output` | None | Save results as JSON | @@ -245,8 +279,8 @@ python scripts/eval/analyze_weights.py \ ## Tips - **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. -- **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 diff --git a/docs/guides/inference.md b/docs/guides/inference.md index c03e3b9..19785b8 100644 --- a/docs/guides/inference.md +++ b/docs/guides/inference.md @@ -49,8 +49,7 @@ KVCache ├── seq_lens [batch_size] ├── out_cache_loc [batch, seq_len] — write indices for this forward ├── max_len int — max(seq_lens), avoids GPU sync in decode - ├── page_table [batch, max_len] — precomputed gather indices for decode (None for prefill) - └── decode_mask [batch, max_len] bool — precomputed position validity mask (None for single-batch decode) + └── kv_indptr [batch + 1] int32 — prefix sum of seq_lens, precomputed once per step ``` 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) - **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 @@ -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}' ``` -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 | |-------|------|---------|-------------| +| `model` | str | `"astrai"` | Model name returned in responses | | `messages` | List[dict] | required | Chat messages (role, content) | -| `top_k` | int | 50 | Top-k count | -| `top_p` | float | 1.0 | Nucleus threshold | -| `temperature` | float | 1.0 | Sampling temperature (> 0.0) | -| `max_tokens` | Optional[int] | None | Max generation length | -| `stream` | bool | False | Stream output | +| `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 | +| `max_tokens` | Optional[int] | 2048 | Max generation length | +| `stream` | Optional[bool] | False | Stream output | | `stop` | Optional[Union[str, List[str]]] | None | Stop sequences | -| `frequency_penalty` | float | 0.0 | Frequency penalty | -| `tools` | Optional[List[dict]] | None | Tool definitions for function calling | -| `tool_choice` | Optional[str] | None | Tool selection mode | +| `n` | Optional[int] | 1 | Number of choices requested | +| `presence_penalty` | Optional[float] | 0.0 | Presence penalty (-2.0 to 2.0) | +| `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 @@ -240,6 +277,8 @@ data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence": event: message_stop data: {"type":"message_stop"} + +data: [DONE] ``` ### Error Responses diff --git a/docs/guides/params.md b/docs/guides/params.md index 86f0160..578be8c 100644 --- a/docs/guides/params.md +++ b/docs/guides/params.md @@ -13,9 +13,11 @@ | 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 | | `--data_root_path` | Dataset root directory | required | | `--param_path` | Model parameters or checkpoint path | required | +| `--resume` | Resume training from `--param_path` | False | | `--n_epoch` | Total training epochs | 1 | | `--batch_per_device` | Batch size per device | 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 | | `--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 @@ -36,9 +38,9 @@ non-matrix parameters through **AdamW** (`fused=True`). | Parameter | Description | Default | |-----------|-------------|---------| | `--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_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_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 | `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, alternating the projection axis (row/column) each step — replacing Muon's Newton-Schulz iteration with a cheaper normalization. | Parameter | Description | Default | |-----------|-------------|---------| -| `--mano_momentum` | Mano momentum factor | 0.95 | -| `--mano_nesterov` | Enable Nesterov momentum for Mano | True | +| `--mano_momentum` | Accepted by the CLI but currently ignored by optimizer construction | 0.95 | +| `--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 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 | | `--random_seed` | Random seed for reproducibility | 3407 | | `--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 @@ -100,14 +105,20 @@ with `--optimizer=muon_adamw`. | Parameter | Description | Default | |-----------|-------------|---------| -| `--log_dir` | Directory for metric logs | checkpoint/logs | -| `--metrics` | Metrics to log (e.g. --metrics loss lr val_loss) | ["loss", "lr", "grad_norm"] | +| `--metrics` | Repeatable metric option (for example, `--metrics loss --metrics lr --metrics val_loss`) | `loss`, `lr`, `grad_norm`, `grad_snr` | ### Gradient Checkpointing | 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 @@ -120,21 +131,25 @@ with `--optimizer=muon_adamw`. | `--backend` | Distributed training backend | nccl | | `--master_addr` | Master node address | localhost | | `--master_port` | Master node port | 29500 | +| `--tp_size` | Reserved tensor-parallel size; accepted but currently ignored | None | ### Strategy-specific | 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` | -| `--group_size` | GRPO group size | 4 | `grpo` | -| `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo` | -| `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo` | +| `--group_size` | GRPO/rollout group size | 4 | `grpo`, `online_grpo`, `online_dpo` | +| `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo`, `online_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` | ### 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 provide a command-line option for configuring one. @@ -151,7 +166,7 @@ provide a command-line option for configuring one. | Parameter | Description | Default | |-----------|-------------|---------| | `--schedule_type` | LR scheduler type (`cosine`, `sgdr`, `wsd`) | cosine | -| `--min_rate` | Minimum LR as fraction of base LR | None (scheduler default: 0.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) | | `--t_mult` | SGDR cycle length multiplier per restart | 2 | | `--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. -# 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`) | 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 | | `--question_key` | str | `question` | Key for the question in input JSON | | `--response_key` | str | `response` | Key for the response in output JSON | -| `--temperature` | float | `0.60` | Sampling temperature | -| `--top_k` | int | `30` | Top-k filtering | +| `--temperature` | float | `0.8` | Sampling temperature | +| `--top_k` | int | `50` | Top-k filtering | | `--top_p` | float | `0.95` | Nucleus sampling threshold | | `--batch_size` | int | `1` | Batch size for generation | | `--num_samples` | int | `1` | Responses per prompt | -| `--max_tokens` | int | model config `max_position_embeddings` | Maximum tokens to generate | -| `--cache_len` | int | `2048` | KV cache length | +| `--max_seq_len` | int | `2048` | KV cache sequence length | | `--frequency_penalty` | float | `0.0` | Frequency penalty | | `--rep_window` | int | `64` | Window size for frequency penalty | @@ -243,14 +249,15 @@ python scripts/tools/generate.py \ | 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 | | `--config`, `-c` | path | required | Preprocessing pipeline config (JSON) | | `--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: ```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. diff --git a/docs/guides/preprocessing.md b/docs/guides/preprocessing.md index 25b8961..9df0f95 100644 --- a/docs/guides/preprocessing.md +++ b/docs/guides/preprocessing.md @@ -10,6 +10,7 @@ Declarative JSON-driven data preprocessing. `MaskBuilderFactory` supports three - [Configuration Reference](#configuration-reference) — all fields - [Mask Algorithm](#mask-algorithm) - [Output Layout](#output-layout) +- [Training Compatibility](#training-compatibility) - [CLI](#cli) - [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` | 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 | | `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 @@ -116,7 +117,7 @@ Config: } ``` -Output keys: `sequence`, `loss_mask` +Output keys: `sequence`, `loss_mask`, `position_ids` ### 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 @@ -180,6 +181,11 @@ Config: 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 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`) - `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 @@ -257,7 +268,7 @@ When `sources` is set, `sections` is ignored. | `max_chars` | int | `2000000` | Skip text-mode items longer than this | | `max_items` | int or null | `null` | Stop after N documents | | `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 | | `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 | |-------|------|---------|-------------| | `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 | -| `max_tokens_per_shard` | int | `100000000` | Flush threshold in cumulative tokens | +| `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 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"}`) | | `position_ids_mode` | str | `"doc_reset"` | How to compute position_ids: `"none"`, `"doc_reset"`, `"continuous"` | @@ -304,11 +315,13 @@ output/ meta.json sequence.bin loss_mask.bin + position_ids.bin wiki/ shard_0000/ meta.json sequence.bin loss_mask.bin + position_ids.bin ``` ### Multi-Shard (`bin`) @@ -322,13 +335,44 @@ output/ meta.json sequence.bin loss_mask.bin + position_ids.bin shard_0001/ meta.json sequence.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 # 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 -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 -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 diff --git a/docs/guides/training.md b/docs/guides/training.md index fcc3117..e4ea14a 100644 --- a/docs/guides/training.md +++ b/docs/guides/training.md @@ -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 $$ -`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 @@ -52,8 +55,8 @@ on_train_begin model.train() on_epoch_begin for batch in dataloader: - on_batch_begin with executor.accumulate(model): + on_batch_begin loss_output = strategy(batch) context.loss = loss_output["loss"].item() context.metrics = loss_output["metrics"] @@ -67,6 +70,7 @@ on_train_begin if executor.sync_gradients: on_optimizer_step optimizer.step() + strategy.on_optimizer_step() optimizer.zero_grad() if scheduler: scheduler.step() @@ -78,16 +82,16 @@ on_train_end | 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_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_epoch_end` | End of each epoch | `MetricCallback`, `ProgressBarCallback` | | `on_error` | On exception during training | `CheckpointCallback`, `MetricCallback` | -| `on_train_end` | Training ends (always via finally) | `CheckpointCallback`, `MetricCallback`, `GradientCheckpointingCallback` | +| `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. @@ -98,7 +102,7 @@ Strategies return `{"loss": Tensor, "metrics": Dict[str, float]}` when called by 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`. @@ -108,7 +112,7 @@ Keys: `input_ids`, `target_ids`. Optional: `label_smoothing`. 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`. @@ -168,9 +172,9 @@ model factory. |------|-------|-------------| | Cosine | `CosineScheduler` | Linear warmup → cosine decay to `min_rate` | | SGDR | `SGDRScheduler` | Cosine annealing with warm restarts (`t_mult=2`) | -| WSD | `WSDScheduler` | Warmup-Stable-Decay with sqrt cooldown | +| 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 @@ -188,10 +192,12 @@ Callback wraps each `DecoderBlock.forward` with `torch.utils.checkpoint.checkpoi ``` 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 ``` +`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`. 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). -> Document Update Time: 2026-07-31 +> Document Update Time: 2026-08-02