From eee7f547894a7dee67f4a92e444829efa3f85790 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Mon, 20 Jul 2026 15:23:30 +0800 Subject: [PATCH] docs: sync training and architecture guides --- assets/docs/architecture.md | 263 ++++++++++++++++++++++++++++++------ assets/docs/params.md | 23 +++- assets/docs/training.md | 30 ++-- 3 files changed, 262 insertions(+), 54 deletions(-) diff --git a/assets/docs/architecture.md b/assets/docs/architecture.md index 10b326b..ced6f88 100644 --- a/assets/docs/architecture.md +++ b/assets/docs/architecture.md @@ -141,6 +141,12 @@ classDiagram +int val_step +float neftune_alpha +str parallel_mode + +int rollout_interval + +float rollout_temperature + +int rollout_top_k + +float rollout_top_p + +int rollout_max_tokens + +Optional[Callable] reward_model_fn +dict executor_kwargs +dict extra_kwargs +validate() @@ -166,13 +172,6 @@ classDiagram +__getitem__(index) Dict } - class RecordDataset { - +Optional[Callable] processor - +load(load_path, storage_type) - +__getitem__(index) - +__len__() - } - class DPODataset { +__getitem__(index) Dict } @@ -222,7 +221,12 @@ classDiagram +fetch_record(index, keys) } - class ResumableDistributedSampler { + class JsonlSource { + +Path path + +load() List[dict] + } + + class RDSampler { +int epoch +int iter } @@ -385,19 +389,103 @@ classDiagram +forward(x) Tensor +set_neftune_alpha(alpha) } + + class LoRAConfig { + +int r + +int alpha + +tuple target_modules + } + + class LoRALinear { + +Linear weight + +Parameter lora_A, lora_B + +forward(x) Tensor + +merge() + } } namespace preprocessing { + class SectionRenderer { + +process_sections(item, sections, config, tokenizer) Tuple + +process_list_field(item, sections, config, tokenizer) Tuple + } + class BaseMaskBuilder { <> +build(item, config, tokenizer) Optional[dict] } - class SectionedMaskBuilder { + class SingleOutputMaskBuilder { +SectionRenderer renderer +build(item, config, tokenizer) Optional[dict] - +_build_single(item, config, tokenizer) Optional[dict] - +_build_multi(item, sources_spec, config, tokenizer) Optional[dict] + } + + class MultiOutputMaskBuilder { + +SectionRenderer renderer + +build(item, config, tokenizer) Optional[dict] + } + + class SectionedMaskBuilder { + +build(item, config, tokenizer) Optional[dict] + } + + class PackingStrategy { + <> + +apply(keys, max_packed_len, truncation_mode) Dict + } + + class PackingStrategyFactory { + +create(name, *args, **kwargs) PackingStrategy + } + + class SimplePacking { + +apply(keys, max_packed_len, truncation_mode) Dict + } + + class BFDPacking { + +apply(keys, max_packed_len, truncation_mode) Dict + } + + class BFDSplitPacking { + +apply(keys, max_packed_len, truncation_mode) Dict + } + + class PositionIdStrategy { + <> + +generate(sequences) List[int] + } + + class PositionIdStrategyFactory { + +create(name, *args, **kwargs) PositionIdStrategy + } + + class NoPositionId { + +generate(sequences) List[int] + } + + class DocResetPositionId { + +generate(sequences) List[int] + } + + class ContinuousPositionId { + +generate(sequences) List[int] + } + + class StoreWriter { + <> + +save(output_dir, domain, shard_idx, tensors) + } + + class StoreWriterFactory { + +create(name, *args, **kwargs) StoreWriter + } + + class BinWriter { + +save(output_dir, domain, shard_idx, tensors) + } + + class H5Writer { + +save(output_dir, domain, shard_idx, tensors) } class Pipeline { @@ -497,7 +585,7 @@ classDiagram class TrainContextBuilder { +TrainConfig config - +with_resume_dir(resume_dir) TrainContextBuilder + +with_param_path(param_path, resume) TrainContextBuilder +build() TrainContext } @@ -544,6 +632,32 @@ classDiagram +sync_old_model() } + class RawRollout { + +Tensor prompts + +Tensor responses + +Tensor response_mask + +Tensor logprobs_old + } + + class RolloutResult { + +Tensor rewards + } + + class BaseRewardModel { + <> + +score(prompts, responses) Tensor + } + + class RolloutGenerator { + +generate(batch) RawRollout + } + + class RolloutRunner { + +step() + +clear_cache() + +__call__(batch) Tuple[RolloutResult, bool] + } + class BaseScheduler { +get_lr() List[float] +step() @@ -857,12 +971,21 @@ classDiagram +apply(logits, filter_value) Tensor } + class FrequencyPenaltyStrategy { + +float penalty + +apply(logits, filter_value, input_ids, input_mask) Tensor + } + class SamplingPipeline { +List[BaseSamplingStrategy] strategies +apply(logits, filter_value) Tensor +sample(logits, filter_value) Tensor } + class StreamDecoder { + +push(token_id) str + } + class GenerateResult { +List[Tuple[int, str]] tokens +List[str] results @@ -881,6 +1004,17 @@ classDiagram +Optional[str] tool_call_id } + class FunctionDef { + +str name + +Optional[str] description + +Optional[Dict] parameters + } + + class ToolDef { + +str type + +FunctionDef function + } + class ChatCompletionRequest { +str model +List[ChatMessage] messages @@ -969,9 +1103,20 @@ classDiagram +str yielded } - class get_app { - <> - +get_app() FastAPI + class BaseToolParser { + <> + +feed(body, current_token_ids, delta_token_ids) List[Dict] + +parse_complete(body) Optional[Dict] + +has_tool_calls (property) bool + } + + class ToolParserFactory { + +create(name, *args, **kwargs) BaseToolParser + } + + class SimpleJsonToolParser { + +feed(body, current_token_ids, delta_token_ids) List[Dict] + +parse_complete(body) Optional[Dict] } } @@ -994,14 +1139,17 @@ classDiagram } namespace parallel { - class setup { - <> - +spawn_parallel_fn(func, world_size, backend, master_addr, master_port, device_type, start_method, **kwargs) - +setup_parallel(rank, world_size, backend, master_addr, master_port, device_type) contextmanager - +get_current_device() str - +get_world_size() int - +get_rank() int - +only_on_rank(rank, sync=False) decorator + class LaunchStrategy { + <> + +launch(func, **kwargs) + } + + class TorchrunStrategy { + +launch(func, **kwargs) + } + + class LocalStrategy { + +launch(func, **kwargs) } class GradientState { @@ -1030,7 +1178,7 @@ classDiagram class BaseExecutor { +GradientState gradient_state - +prepare(model, optimizer, dataloader, scheduler) tuple + +prepare(model_fn, optimizer_fn, scheduler_fn, before_wrap) tuple +accumulate(model) context manager +backward(loss) +unwrap_model(model) dict @@ -1052,6 +1200,12 @@ classDiagram +unwrap_model(model) dict } + class FSDP2Executor { + -_prepare_model(model) nn.Module + -_no_sync(model) context manager + +unwrap_model(model) dict + } + class ExecutorFactory { +Dict _entries +register(name) decorator @@ -1104,9 +1258,8 @@ classDiagram TrainCallback <|-- MetricCallback BaseDataset <|-- SEQDataset BaseDataset <|-- SFTDataset - BaseDataset <|-- RecordDataset - RecordDataset <|-- DPODataset - RecordDataset <|-- GRPODataset + BaseDataset <|-- DPODataset + BaseDataset <|-- GRPODataset Store <|-- H5Store Store <|-- MmapStore Store <|-- JsonlStore @@ -1119,6 +1272,7 @@ classDiagram BaseSamplingStrategy <|-- TemperatureStrategy BaseSamplingStrategy <|-- TopKStrategy BaseSamplingStrategy <|-- TopPStrategy + BaseSamplingStrategy <|-- FrequencyPenaltyStrategy ParallelModel <|-- RowParallelLinear ParallelModel <|-- ColumnParallelLinear AutoModel <|-- AutoRegressiveLM @@ -1142,12 +1296,31 @@ classDiagram BaseFactory <|-- ExecutorFactory BaseFactory <|-- ConfigFactory BaseFactory <|-- MaskBuilderFactory + BaseFactory <|-- PackingStrategyFactory + BaseFactory <|-- PositionIdStrategyFactory + BaseFactory <|-- StoreWriterFactory + BaseFactory <|-- ToolParserFactory BaseExecutor <|-- NoneExecutor BaseExecutor <|-- DDPExecutor BaseExecutor <|-- FSDPExecutor + BaseExecutor <|-- FSDP2Executor ResponseBuilder <|-- OpenAIResponseBuilder ResponseBuilder <|-- AnthropicResponseBuilder + BaseToolParser <|-- SimpleJsonToolParser BaseMaskBuilder <|-- SectionedMaskBuilder + BaseMaskBuilder <|-- SingleOutputMaskBuilder + BaseMaskBuilder <|-- MultiOutputMaskBuilder + PackingStrategy <|-- SimplePacking + PackingStrategy <|-- BFDPacking + BFDPacking <|-- BFDSplitPacking + PositionIdStrategy <|-- NoPositionId + PositionIdStrategy <|-- DocResetPositionId + PositionIdStrategy <|-- ContinuousPositionId + StoreWriter <|-- BinWriter + StoreWriter <|-- H5Writer + RawRollout <|-- RolloutResult + LaunchStrategy <|-- TorchrunStrategy + LaunchStrategy <|-- LocalStrategy KVCache <|-- PageCache KVCache <|-- ContiguousCache CacheView <|-- PageCacheView @@ -1169,6 +1342,8 @@ classDiagram EmbeddingEncoder *-- Embedding DecoderBlock *-- RMSNorm ChatCompletionRequest *-- ChatMessage + ChatCompletionRequest *-- ToolDef + ToolDef *-- FunctionDef MessagesRequest *-- AnthropicMessage BaseExecutor *-- GradientState AccumOptimizer o-- GradientState @@ -1191,6 +1366,9 @@ classDiagram Pipeline o-- PipelineConfig Pipeline o-- BaseMaskBuilder Pipeline o-- AutoTokenizer + Pipeline o-- PackingStrategy + Pipeline o-- PositionIdStrategy + Pipeline o-- StoreWriter TokenizeTransform o-- AutoTokenizer TokenizeTransform o-- BaseMaskBuilder @@ -1198,6 +1376,9 @@ classDiagram TrainConfig ..> BaseStrategy : selects PipelineConfig ..> MaskBuilderFactory : selects MaskBuilderFactory ..> BaseMaskBuilder : creates + PackingStrategyFactory ..> PackingStrategy : creates + PositionIdStrategyFactory ..> PositionIdStrategy : creates + StoreWriterFactory ..> StoreWriter : creates StrategyFactory ..> BaseStrategy : creates SchedulerFactory ..> BaseScheduler : creates DatasetFactory ..> BaseDataset : creates @@ -1216,12 +1397,13 @@ classDiagram ExecutorFactory ..> NoneExecutor : creates ExecutorFactory ..> DDPExecutor : creates ExecutorFactory ..> FSDPExecutor : creates + ExecutorFactory ..> FSDP2Executor : creates + ToolParserFactory ..> BaseToolParser : creates TrainContextBuilder ..> ExecutorFactory : creates Trainer ..> TrainContextBuilder : uses TrainContextBuilder ..> TrainContext : creates - Trainer ..> Functions : spawns TrainContextBuilder ..> StrategyFactory : uses - TrainContextBuilder ..> ResumableDistributedSampler : creates + TrainContextBuilder ..> RDSampler : creates Checkpoint ..> Checkpoint : serializes CheckpointCallback ..> Checkpoint : creates PageCache ..> PageCacheView : binds @@ -1232,6 +1414,9 @@ classDiagram AnthropicResponseBuilder ..> MessagesRequest : receives ProtocolHandler ..> StopChecker : creates ProtocolHandler ..> GenContext : creates + RolloutGenerator ..> InferenceScheduler : uses + RolloutRunner ..> RolloutGenerator : uses + RolloutRunner ..> BaseRewardModel : uses %% --- Association (general usage) --- Trainer --> TrainConfig @@ -1253,14 +1438,14 @@ 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** | BaseMaskBuilder, MaskBuilderFactory, SectionedMaskBuilder, SingleOutputMaskBuilder, MultiOutputMaskBuilder, Pipeline, TokenizeTransform, filter_by_length, PackingStrategy, PackingStrategyFactory, plan_bfd, PositionIdStrategy, PositionIdStrategyFactory, StoreWriter, StoreWriterFactory, core (shared helpers) | Declarative JSON-driven data preprocessing | -| **astrai.dataset** | BaseDataset–RecordDataset–DPO/GRPODataset, SEQDataset, SFTDataset, Store, Streamable, Recordable, H5Store, MmapStore, JsonlStore, StoreFactory, ResumableDistributedSampler, DatasetFactory | Dataset loading and management | +| **astrai.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.serialization** | Checkpoint | Model serialization | -| **astrai.model** | AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, RotaryEmbedding, Embedding | Neural network model | +| **astrai.model** | 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 | Training workflow | -| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, KVCache–ContiguousCache/PageCache, CacheView–ContiguousCacheView/PageCacheView, Allocator–Storage, Task, TaskManager, TaskStatus, GenerationRequest, GenerateResult, BaseSamplingStrategy–SamplingPipeline, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, ChatMessage–MessagesRequest, app | Inference service | -| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, GradientState, AccumOptimizer, AccumScheduler, ParallelModel, RowParallelLinear, ColumnParallelLinear | Distributed parallel & gradient accumulation | +| **astrai.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, KVCache–ContiguousCache/PageCache, CacheView–ContiguousCacheView/PageCacheView, Allocator–Storage, 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 | +| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, LaunchStrategy, TorchrunStrategy, LocalStrategy, BaseExecutor, ExecutorFactory, NoneExecutor, DDPExecutor, FSDPExecutor, FSDP2Executor, GradientState, AccumOptimizer, AccumScheduler, ParallelModel, RowParallelLinear, ColumnParallelLinear | Distributed parallel & gradient accumulation | | **astrai.factory** | BaseFactory | Component registration | | **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers | @@ -1268,7 +1453,7 @@ classDiagram | Pattern | Classes | Purpose | |---------|---------|---------| -| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StoreFactory`, `ConfigFactory`, `ExecutorFactory`, `MaskBuilderFactory`, `StoreWriterFactory`, `PackingStrategyFactory`, `PositionIdStrategyFactory` | Decorator-based component creation | +| **Factory** | `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 | @@ -1277,7 +1462,7 @@ classDiagram | **Observer** | `TrainCallback`, callback implementations | Training process monitoring | | **Context** | `TrainContext` | Unified training state bag | | **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction | -| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution | +| **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor`, `FSDP2Executor` | Gradient accumulation & model distribution | | **Storage** | `Store`, `H5Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support | | **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching | | **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading | @@ -1287,7 +1472,7 @@ classDiagram 1. **Config → Training**: `TrainConfig` holds `model_fn`, `dataset`, `optimizer_fn`, `scheduler_fn`, `parallel_mode`, `executor_kwargs` 2. **Training Flow**: `Trainer` → `TrainContextBuilder` → `TrainContext`, uses `BaseStrategy` for loss, `BaseExecutor` for gradient accumulation + model distribution 3. **Strategy Selection**: `StrategyFactory` creates strategy by `train_type` -4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)` → `NoneExecutor` / `DDPExecutor` / `FSDPExecutor` +4. **Executor Selection**: `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)` → `NoneExecutor` / `DDPExecutor` / `FSDPExecutor` / `FSDP2Executor` 5. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `AutoRegressiveLM`, backed by `KVCache` + `SamplingPipeline` 6. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP 7. **Dataset Loading**: `DatasetFactory` creates datasets, `Store` (H5Store/MmapStore/JsonlStore) loads data with explicit `_length` and multi-segment `_data` @@ -1296,4 +1481,4 @@ classDiagram 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-19 +> Document Update Time: 2026-07-20 diff --git a/assets/docs/params.md b/assets/docs/params.md index 7f10942..65b75ef 100644 --- a/assets/docs/params.md +++ b/assets/docs/params.md @@ -13,7 +13,7 @@ | Parameter | Description | Default | |-----------|-------------|---------| -| `--train_type` | Training type (`seq`, `sft`, `dpo`, `grpo`) | required | +| `--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 | | `--n_epoch` | Total training epochs | 1 | @@ -100,18 +100,31 @@ Combined optimizer: matrix parameters via **Muon**, non-matrix via **AdamW** (`f | `--group_size` | GRPO group size | 4 | `grpo` | | `--grpo_clip_eps` | GRPO clipping epsilon | 0.2 | `grpo` | | `--grpo_kl_coef` | GRPO KL penalty coefficient | 0.01 | `grpo` | -| `--grpo_sync_interval` | GRPO ref_model sync interval (steps) | 200 | `grpo` | | `--neftune_alpha` | NEFTune noise alpha (0=disabled, typical: 5.0) | 0.0 | `sft` | +### Online Rollout + +These options apply to `online_grpo` and `online_dpo`. Online strategies require +a `BaseRewardModel` factory in `TrainConfig`; `train.py` does not currently +provide a command-line option for configuring one. + +| Parameter | Description | Default | +|-----------|-------------|---------| +| `--rollout_interval` | Optimizer steps between rollout refreshes | 512 | +| `--rollout_temperature` | Rollout sampling temperature | 0.7 | +| `--rollout_top_k` | Rollout top-k filtering (`0` disables) | 0 | +| `--rollout_top_p` | Rollout nucleus sampling threshold | 0.9 | +| `--rollout_max_tokens` | Maximum generated tokens per response | 1024 | + ### Scheduler | Parameter | Description | Default | |-----------|-------------|---------| | `--schedule_type` | LR scheduler type (`cosine`, `sgdr`, `wsd`) | cosine | -| `--min_rate` | Minimum LR as fraction of base LR | None (scheduler default: 0.01) | +| `--min_rate` | Minimum LR as fraction of base LR | None (scheduler default: 0.05 for cosine/SGDR, 0.0 for WSD) | | `--cycle_length` | SGDR first cycle length in steps | None (total_steps - warmup_steps) | | `--t_mult` | SGDR cycle length multiplier per restart | 2 | -| `--stable_steps` | WSD stable plateau steps | None (required for wsd) | +| `--stable_steps` | WSD stable plateau steps | None (80% of post-warmup steps) | | `--decay_steps` | WSD decay steps | None (total_steps - warmup_steps - stable_steps) | ### Usage Example @@ -201,4 +214,4 @@ See [Preprocessing Guide](preprocessing.md) for config file format and examples. --- -> Document Update Time: 2026-07-19 \ No newline at end of file +> Document Update Time: 2026-07-20 diff --git a/assets/docs/training.md b/assets/docs/training.md index c81b704..22bc829 100644 --- a/assets/docs/training.md +++ b/assets/docs/training.md @@ -6,7 +6,7 @@ - [Causal Mask](#causal-mask) - [Rotary Position Embedding (RoPE)](#rotary-position-embedding-rope) - [Training Loop](#training-loop) -- [Strategies](#strategies) — SEQ, SFT, DPO, GRPO +- [Strategies](#strategies) — SEQ, SFT, DPO, GRPO, online rollout - [LR Schedulers](#lr-schedulers) - [Gradient Checkpointing](#gradient-checkpointing) - [Checkpoint](#checkpoint) @@ -146,6 +146,19 @@ Parameters: `group_size=4`, `clip_eps=0.2`, `kl_coef=0.01`. External sync of `ol Keys: `prompts`, `responses`, `masks`, `rewards`. +### Online Rollout + +`online_grpo` and `online_dpo` use the respective GRPO and DPO strategies with +a `RolloutRunner`. The runner renders prompts through the tokenizer chat +template, generates grouped responses through `InferenceScheduler`, then scores +them with a `BaseRewardModel`. It refreshes cached rollouts every +`rollout_interval` optimizer steps. `online_grpo` synchronizes `old_model` when +a fresh rollout is produced. + +Online strategies require `TrainConfig.reward_model_fn`. `train.py` exposes the +rollout sampling parameters but does not yet offer a CLI argument for the reward +model factory. + ## LR Schedulers | Type | Class | Description | @@ -162,6 +175,7 @@ Trades compute for memory by recomputing activations during backward pass. Speci ```python from astrai.model.components.decoder_block import DecoderBlock + config = TrainConfig(..., gradient_checkpointing_modules=[DecoderBlock]) ``` @@ -181,18 +195,14 @@ Model config (`context.model_config`) saved into `config.json` during training v ## TrainContextBuilder (Builder Pattern) ```python -context = ( - TrainContextBuilder(config) - .with_resume_dir(resume_dir) - .build() -) +context = TrainContextBuilder(config).with_param_path(param_path, resume=True).build() # Returns TrainContext with model, strategy, optimizer, scheduler, dataloader, checkpoint ``` -- Loads checkpoint weights if provided +- Loads checkpoint weights before the model is wrapped - Creates executor via `ExecutorFactory.create(cfg.parallel_mode, grad_accum_steps=cfg.grad_accum_steps, **cfg.executor_kwargs)` -- Calls `executor.prepare(model, optimizer, dataloader, scheduler)` for model distribution (e.g. DDP) + gradient accumulation wrappers -- Creates `ResumableDistributedSampler` for shuffle+resume +- Calls `executor.prepare(model_fn, optimizer_fn, scheduler_fn, before_wrap=...)`; the executor creates, wraps, then builds the optimizer and scheduler for the wrapped model +- Creates `RDSampler` for shuffle+resume - Builds strategy via `StrategyFactory.create(train_type, model, device, **kwargs)` ## Training CLI @@ -222,4 +232,4 @@ nohup python scripts/tools/train.py \ Full parameter reference at [params.md](params.md). -> Document Update Time: 2026-07-19 +> Document Update Time: 2026-07-20