docs: audit non-CUDA documentation
- Aligns CLI and strategy metric contracts - Refreshes architecture, dataflow, preprocessing, distributed, and eval guides - Corrects links, TOCs, defaults, and repository paths
This commit is contained in:
+110
-104
@@ -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 {
|
||||
<<nn.Module>>
|
||||
+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 {
|
||||
<<TypedDict>>
|
||||
+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 {
|
||||
<<TypedDict>>
|
||||
+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 {
|
||||
<<TypedDict>>
|
||||
+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 {
|
||||
<<abstract>>
|
||||
+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 {
|
||||
<<abstract>>
|
||||
+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
|
||||
<<abstract>>
|
||||
+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
|
||||
|
||||
+74
-40
@@ -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)`
|
||||
|
||||
+18
-10
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user