docs: audit non-CUDA documentation

- Aligns CLI and strategy metric contracts
- Refreshes architecture, dataflow, preprocessing, distributed, and eval guides
- Corrects links, TOCs, defaults, and repository paths
This commit is contained in:
2026-08-02 07:39:24 +08:00
parent 020e2eff4e
commit 288ba20db1
13 changed files with 483 additions and 267 deletions
+110 -104
View File
@@ -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, BaseStrategyGRPOStrategy, StrategyFactory, BaseSchedulerWSDScheduler, SchedulerFactory, TrainCallback(Protocol)MetricCallback, CallbackFactory, RawRollout, RolloutResult, BaseRewardModel, RolloutGenerator, RolloutRunner | Training workflow |
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, PagePool, KVStorage, ReqToTokenPool, KVCache, Allocator, PrefixCache, Task, TaskManager, TaskStatus, StreamDecoder, GenerationRequest, GenerateResult, BaseSamplingStrategySamplingPipeline, FrequencyPenaltyStrategy, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuilder, StopChecker, GenContext, StopInfo, ChatMessage, FunctionDef, ToolDef, ChatCompletionRequest, AnthropicMessage, MessagesRequest, BaseToolParser, ToolParserFactory, SimpleJsonToolParser | Inference service |
@@ -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