@@ -22,7 +22,8 @@ classDiagram
+int n_layers
+float norm_eps
+int dim_ffn
+bool tie_weight
+Optional[ bool] tie_weight
+Optional[dict] rope_scaling
+int max_len
+float rope_theta
+str attn_type
@@ -52,6 +53,7 @@ classDiagram
+int n_kv_heads
+bool use_qk_norm
+bool use_gated_attention
+Optional[dict] rope_scaling
+Optional[str] pooling_type
+Optional[bool] normalize_embeddings
}
@@ -80,6 +82,7 @@ classDiagram
+str log_dir
+int log_interval
+List[str] metrics
+Optional[LoRAConfig] lora
+int random_seed
+int num_workers
+Optional[int] prefetch_factor
@@ -457,16 +460,15 @@ classDiagram
+on_train_end(context)
+on_epoch_begin(context)
+on_epoch_end(context)
+on_step_begin(context)
+on_step_end(context)
+on_batch_begin(context)
+on_batch_end(context)
+on_optimizer_step(context)
+on_error(context)
}
class GradientClippingCallback {
+float max_grad_norm
+on_step_begin (context)
+on_optimizer_step (context)
}
class GradientCheckpointingCallback {
@@ -512,7 +514,7 @@ classDiagram
class ValidationCallback {
+_run_validation(context)
+on_step_end (context)
+on_optimizer_ step(context)
}
class CallbackFactory {
@@ -747,56 +749,58 @@ classDiagram
+str model
+List[AnthropicMessage] messages
+Optional[str] system
+float temperature
+float top_p
+int top_k
+Optional[ float] temperature
+Optional[ float] top_p
+Optional[ int] top_k
+int max_tokens
+bool stream
+Optional[ bool] stream
+Optional[List[str]] stop_sequences
}
class ProtocolHandl er {
class ResponseBuild er {
<<abstract>>
+prepare(request, engine) Tuple[str, GenContext, List[str]]
+format_stream_start(ctx) List[str]
+format_chunk(token) str
+format_stream_end(ctx, stop) List[str]
+format_response(ctx, content, stop) Dict
}
class OpenAIResponseBuilder {
+prepare(request, engine) Tuple
+format_stream_start(ctx) List[str]
+format_chunk(token) str
+format_stream_end(ctx, stop) List[str]
+format_response(ctx, content, stop) Dict
}
class AnthropicResponseBuilder {
+prepare(request, engine) Tuple
+format_stream_start(ctx) List[str]
+format_chunk(token) str
+format_stream_end(ctx, stop) List[str]
+format_response(ctx, content, stop) Dict
}
class ProtocolHandler {
+request
+engine
+build_prompt() st r
+create_response_id() str
+get_stop_sequences() List[str]
+create_stop_checker() StopChecker
+on_token(ctx, token, stop_checker) Optional[str]
+format_stream_start(ctx) List[str]
+format_stream_token(ctx, token) str
+format_stream_end(ctx) List[str]
+format_non_stream_response(ctx, content) Dict
+builder: ResponseBuilde r
+handle() Union[StreamingResponse, Dict]
}
class OpenAIHandler {
+build_prompt() str
+create_response_id() str
}
class AnthropicHandler {
+build_prompt() str
+create_response_id() str
+on_token(ctx, token, stop_checker) Optional[str]
-_handle_stream(agen, ctx, stops) StreamingResponse
-_handle_non_stream(agen, ctx, stops) Dict
}
class StopChecker {
+has_sequences (property) bool
+check(text) Optional[str]
+trim(text, matched) str
}
class Stream Context {
class Gen Context {
+str resp_id
+int created
+str model
+int prompt_tokens
+int completion_tokens
+str accumulated
+Optional[str] stop_matched
+str last_yield_trimmed
}
class app {
@@ -876,6 +880,11 @@ classDiagram
+unwrap_model(model) nn.Module
}
class FSDPExecutor {
+_prepare_model(model) nn.Module
+unwrap_model(model) nn.Module
}
class ExecutorFactory {
+Registry _registry
+register(name) decorator
@@ -911,6 +920,7 @@ classDiagram
TrainCallback <|-- CheckpointCallback
TrainCallback <|-- ProgressBarCallback
TrainCallback <|-- MetricLoggerCallback
TrainCallback <|-- ValidationCallback
BaseDataset <|-- SEQDataset
BaseDataset <|-- SFTDataset
BaseDataset <|-- DPODataset
@@ -941,15 +951,14 @@ classDiagram
BaseFactory <|-- ConfigFactory
BaseExecutor <|-- NoneExecutor
BaseExecutor <|-- DDPExecutor
ProtocolHandler <|-- OpenAIHandle r
ProtocolHandler <|-- AnthropicHandl er
BaseExecutor <|-- FSDPExecuto r
ResponseBuilder <|-- OpenAIResponseBuild er
ResponseBuilder <|-- AnthropicResponseBuilder
%% --- Composition (strong ownership, part destroyed with whole) ---
KVCache *-- PagePool
KVCache *-- Storage
KVCache *-- TaskTable
PagePool *-- Allocator
PagePool *-- PrefixCache
InferenceEngine *-- InferenceScheduler
InferenceScheduler *-- KVCache
InferenceScheduler *-- Executor
@@ -963,7 +972,6 @@ classDiagram
DecoderBlock *-- RMSNorm
ChatCompletionRequest *-- ChatMessage
MessagesRequest *-- AnthropicMessage
AutoTokenizer *-- ChatTemplate
BaseFactory *-- Registry
BaseExecutor *-- GradientState
AccumOptimizer o-- GradientState
@@ -971,6 +979,9 @@ 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
@@ -998,6 +1009,7 @@ classDiagram
ConfigFactory ..> EncoderConfig : creates
ExecutorFactory ..> NoneExecutor : creates
ExecutorFactory ..> DDPExecutor : creates
ExecutorFactory ..> FSDPExecutor : creates
TrainContextBuilder ..> ExecutorFactory : creates
Trainer ..> TrainContextBuilder : uses
TrainContextBuilder ..> TrainContext : creates
@@ -1009,10 +1021,10 @@ classDiagram
KVCache ..> KvcacheView : binds
InferenceEngine ..> GenerationRequest : uses
InferenceEngine ..> GenerateResult : creates
OpenAIHandl er ..> ChatCompletionRequest : receives
AnthropicHandl er ..> MessagesRequest : receives
OpenAIResponseBuild er ..> ChatCompletionRequest : receives
AnthropicResponseBuild er ..> MessagesRequest : receives
ProtocolHandler ..> StopChecker : creates
ProtocolHandler ..> Stream Context : creates
ProtocolHandler ..> Gen Context : creates
%% --- Association (general usage) ---
Trainer --> TrainConfig
@@ -1026,7 +1038,6 @@ classDiagram
Executor --> AutoTokenizer
TaskManager --> AutoTokenizer
MultiSegmentFetcher --> BaseSegmentFetcher
ResumableDistributedSampler --> BaseDataset
```
@@ -1041,8 +1052,8 @@ classDiagram
| **astrai.model ** | AutoModel, AutoRegressiveLM, EmbeddingEncoder, DecoderBlock, GQA, MLA, MLP, DeepSeekMoE, AttnFactory, FFNFactory, RMSNorm, Linear, RotaryEmbedding, Embedding | Neural network model |
| **astrai.tokenize ** | AutoTokenizer, ChatTemplate | Tokenizer and chat template |
| **astrai.trainer ** | Trainer, TrainContext, TrainContextBuilder, BaseStrategy– GRPOStrategy, StrategyFactory, BaseScheduler– SGDRScheduler, SchedulerFactory, TrainCallback(Protocol)– ValidationCallback, CallbackFactory, Muon | Training workflow |
| **astrai.inference ** | InferenceEngine, InferenceScheduler, Executor, KVCache– KvcacheView, Allocator– Storage, Task, TaskManager, TaskStatus, GenerationRequest, GenerateResult, BaseSamplingStrategy– SamplingPipeline, ProtocolHandler– AnthropicHandl er, StopChecker, Stream Context, 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, GradientState, AccumOptimizer, AccumScheduler, ParallelModel, RowParallelLinear, ColumnParallelLinear | Distributed parallel & gradient accumulation |
| **astrai.inference ** | InferenceEngine, InferenceScheduler, Executor, KVCache– KvcacheView, Allocator– Storage, Task, TaskManager, TaskStatus, GenerationRequest, GenerateResult, BaseSamplingStrategy– SamplingPipeline, ProtocolHandler, ResponseBuilder, OpenAIResponseBuilder, AnthropicResponseBuild er, StopChecker, Gen Context, 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.factory ** | Registry, BaseFactory[T] | Component registration |
| **astrai.protocols ** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
@@ -1054,7 +1065,7 @@ classDiagram
| **Registry ** | `BaseFactory` , `Registry` | Component registration with category/priority |
| **Strategy ** | `SEQStrategy` , `SFTStrategy` , `DPOStrategy` , `GRPOStrategy` | Training strategy switching |
| **Strategy (Sampling) ** | `TemperatureStrategy` , `TopKStrategy` , `TopPStrategy` , `SamplingPipeline` | Composable logit transformations |
| **Template Method ** | `ProtocolHandl er` , `OpenAIHandl er` , `AnthropicHandl er` | HTTP API handler with format hooks |
| **Strategy (API) ** | `ResponseBuild er` , `OpenAIResponseBuild er` , `AnthropicResponseBuild er` | HTTP API handler with format hooks |
| **Builder ** | `TrainContextBuilder` | Chain-building training context |
| **Observer ** | `TrainCallback` , callback implementations | Training process monitoring |
| **Context ** | `TrainContext` | Unified training state bag |
@@ -1069,7 +1080,7 @@ classDiagram
1. **Config → Training ** : `TrainConfig` holds model, 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(parallel_mode, ** executor_kwargs)` → `NoneExecutor` (single) / `DDPExecutor` (distributed)
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 `KVCache` + `SamplingPipeline`
6. **Distributed ** : `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
7. **Dataset Loading ** : `DatasetFactory` creates datasets, `BaseStorage` (H5Storage/JSONStorage) loads via `BaseSegmentFetcher` + `MultiSegmentFetcher`