Compare commits
15
Commits
v1.3.6
..
65ab69543b
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
65ab69543b | ||
|
|
1d26aa2e93 | ||
|
|
a548d4553e | ||
|
|
dd1b39f435 | ||
|
|
94d6e713e9 | ||
|
|
47c37e4876 | ||
|
|
737585a32a | ||
|
|
a4688021bf | ||
|
|
7df6eb9211 | ||
|
|
82a3f2626f | ||
|
|
7fa69572c0 | ||
|
|
3ab4f237e5 | ||
|
|
8cbf3f36e2 | ||
|
|
0594ce1017 | ||
|
|
ff509ff39f |
+99
-17
@@ -88,12 +88,12 @@ classDiagram
|
||||
+str backend
|
||||
+str master_addr
|
||||
+str master_port
|
||||
+Callable parallel_wrapper
|
||||
+Callable state_dict_fn
|
||||
+str start_method
|
||||
+str device_type
|
||||
+Optional[Dataset] val_dataset
|
||||
+int val_step
|
||||
+str parallel_mode
|
||||
+dict executor_kwargs
|
||||
+dict extra_kwargs
|
||||
+validate()
|
||||
}
|
||||
@@ -257,11 +257,13 @@ classDiagram
|
||||
+int qk_rope_head_dim
|
||||
+int n_rep
|
||||
+int layer_id
|
||||
+bool use_qk_norm
|
||||
+bool use_gated_attention
|
||||
+Linear q_proj, kv_a_proj, kv_b_proj
|
||||
+Linear o_proj
|
||||
+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, paged_cache) Tensor
|
||||
}
|
||||
|
||||
@@ -364,10 +366,11 @@ classDiagram
|
||||
+nn.Module model
|
||||
+BaseStrategy strategy
|
||||
+DataLoader dataloader
|
||||
+Optimizer optimizer
|
||||
+LRScheduler scheduler
|
||||
+OptimizerProtocol optimizer
|
||||
+SchedulerProtocol scheduler
|
||||
+Checkpoint checkpoint
|
||||
+TrainConfig config
|
||||
+BaseExecutor executor
|
||||
+int epoch
|
||||
+int iteration
|
||||
+float loss
|
||||
@@ -802,6 +805,24 @@ classDiagram
|
||||
}
|
||||
}
|
||||
|
||||
namespace protocols {
|
||||
class OptimizerProtocol {
|
||||
<<protocol>>
|
||||
+step(closure)
|
||||
+zero_grad()
|
||||
+state_dict() dict
|
||||
+load_state_dict(d)
|
||||
}
|
||||
|
||||
class SchedulerProtocol {
|
||||
<<protocol>>
|
||||
+step()
|
||||
+state_dict() dict
|
||||
+load_state_dict(d)
|
||||
+get_last_lr()
|
||||
}
|
||||
}
|
||||
|
||||
namespace parallel {
|
||||
class Functions {
|
||||
<<module>>
|
||||
@@ -813,6 +834,54 @@ classDiagram
|
||||
+only_on_rank(rank, sync) decorator
|
||||
}
|
||||
|
||||
class GradientState {
|
||||
+int num_steps
|
||||
+sync_gradients (property) bool
|
||||
}
|
||||
|
||||
class AccumOptimizer {
|
||||
+Optimizer optimizer
|
||||
+GradientState gradient_state
|
||||
+step(closure)
|
||||
+zero_grad()
|
||||
+state_dict() dict
|
||||
+load_state_dict(d)
|
||||
}
|
||||
|
||||
class AccumScheduler {
|
||||
+LRScheduler scheduler
|
||||
+GradientState gradient_state
|
||||
+step()
|
||||
+state_dict() dict
|
||||
+load_state_dict(d)
|
||||
+get_last_lr()
|
||||
}
|
||||
|
||||
class BaseExecutor {
|
||||
+GradientState gradient_state
|
||||
+prepare(model, optimizer, dataloader, scheduler) tuple
|
||||
+accumulate(model) context manager
|
||||
+backward(loss)
|
||||
+unwrap_model(model) nn.Module
|
||||
+sync_gradients (property) bool
|
||||
+grad_accum_steps (property) int
|
||||
}
|
||||
|
||||
class NoneExecutor {
|
||||
}
|
||||
|
||||
class DDPExecutor {
|
||||
+_prepare_model(model) nn.Module
|
||||
+_no_sync(model) context manager
|
||||
+unwrap_model(model) nn.Module
|
||||
}
|
||||
|
||||
class ExecutorFactory {
|
||||
+Registry _registry
|
||||
+register(name) decorator
|
||||
+create(parallel_mode, **kwargs) BaseExecutor
|
||||
}
|
||||
|
||||
class ParallelModel {
|
||||
+dist.ProcessGroup process_group
|
||||
+int rank
|
||||
@@ -868,8 +937,10 @@ classDiagram
|
||||
BaseFactory <|-- SchedulerFactory
|
||||
BaseFactory <|-- CallbackFactory
|
||||
BaseFactory <|-- StorageFactory
|
||||
BaseFactory <|-- ExecutorFactory
|
||||
BaseFactory <|-- ConfigFactory
|
||||
TrainCallback <|-- ValidationCallback
|
||||
BaseExecutor <|-- NoneExecutor
|
||||
BaseExecutor <|-- DDPExecutor
|
||||
ProtocolHandler <|-- OpenAIHandler
|
||||
ProtocolHandler <|-- AnthropicHandler
|
||||
|
||||
@@ -894,6 +965,9 @@ classDiagram
|
||||
MessagesRequest *-- AnthropicMessage
|
||||
AutoTokenizer *-- ChatTemplate
|
||||
BaseFactory *-- Registry
|
||||
BaseExecutor *-- GradientState
|
||||
AccumOptimizer o-- GradientState
|
||||
AccumScheduler o-- GradientState
|
||||
|
||||
%% --- Aggregation (weak ownership) ---
|
||||
AutoModel o-- BaseModelConfig
|
||||
@@ -901,6 +975,7 @@ classDiagram
|
||||
TrainContext o-- BaseStrategy
|
||||
TrainContext o-- BaseScheduler
|
||||
TrainContext o-- Checkpoint
|
||||
TrainContext o-- BaseExecutor
|
||||
KvcacheView o-- Storage
|
||||
SamplingPipeline o-- BaseSamplingStrategy
|
||||
BaseDataset o-- BaseStorage
|
||||
@@ -921,6 +996,9 @@ classDiagram
|
||||
StorageFactory ..> JSONStorage : creates
|
||||
ConfigFactory ..> AutoRegressiveLMConfig : creates
|
||||
ConfigFactory ..> EncoderConfig : creates
|
||||
ExecutorFactory ..> NoneExecutor : creates
|
||||
ExecutorFactory ..> DDPExecutor : creates
|
||||
TrainContextBuilder ..> ExecutorFactory : creates
|
||||
Trainer ..> TrainContextBuilder : uses
|
||||
TrainContextBuilder ..> TrainContext : creates
|
||||
Trainer ..> Functions : spawns
|
||||
@@ -963,15 +1041,16 @@ 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, BaseSamplingStrategy–SamplingPipeline, ProtocolHandler–AnthropicHandler, ChatMessage–MessagesRequest, app | Inference service |
|
||||
| **astrai.parallel** | spawn_parallel_fn, setup_parallel, get_rank/get_world_size/get_current_device, only_on_rank, ParallelModel, RowParallelLinear, ColumnParallelLinear | Distributed parallel |
|
||||
| **astrai.inference** | InferenceEngine, InferenceScheduler, Executor, KVCache–KvcacheView, Allocator–Storage, Task, TaskManager, TaskStatus, GenerationRequest, GenerateResult, BaseSamplingStrategy–SamplingPipeline, ProtocolHandler–AnthropicHandler, StopChecker, StreamContext, 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.factory** | Registry, BaseFactory[T] | Component registration |
|
||||
| **astrai.protocols** | OptimizerProtocol, SchedulerProtocol | Structural subtyping for optimizer/scheduler wrappers |
|
||||
|
||||
## Design Patterns
|
||||
|
||||
| Pattern | Classes | Purpose |
|
||||
|---------|---------|---------|
|
||||
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StorageFactory`, `ConfigFactory` | Decorator-based component creation |
|
||||
| **Factory** | `AttnFactory`, `FFNFactory`, `StrategyFactory`, `DatasetFactory`, `SchedulerFactory`, `CallbackFactory`, `StorageFactory`, `ConfigFactory`, `ExecutorFactory` | Decorator-based component creation |
|
||||
| **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 |
|
||||
@@ -980,20 +1059,23 @@ 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` | Gradient accumulation & model distribution |
|
||||
| **Storage** | `BaseStorage`, `H5Storage`, `JSONStorage` | Format-agnostic data access |
|
||||
| **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching |
|
||||
| **AutoModel Registry** | `AutoModel`, `AutoRegressiveLM`, `EmbeddingEncoder` | Model-type dynamic loading |
|
||||
|
||||
## Core Relationships
|
||||
|
||||
1. **Config → Training**: `TrainConfig` holds model, dataset, optimizer_fn, scheduler_fn
|
||||
2. **Training Flow**: `Trainer` → `TrainContextBuilder` → `TrainContext`, uses `BaseStrategy` for loss
|
||||
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. **Inference Flow**: `InferenceEngine` → `InferenceScheduler` → `AutoRegressiveLM`, backed by `KVCache` + `SamplingPipeline`
|
||||
5. **Distributed**: `spawn_parallel_fn` + `setup_parallel` for multi-process DDP
|
||||
6. **Dataset Loading**: `DatasetFactory` creates datasets, `BaseStorage` (H5Storage/JSONStorage) loads via `BaseSegmentFetcher` + `MultiSegmentFetcher`
|
||||
7. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only)
|
||||
8. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`
|
||||
9. **AutoModel**: `from_pretrained()` loads `config.json` + `model.safetensors`, `_disable_random_init` replaces `nn.init.*` with no-ops
|
||||
4. **Executor Selection**: `ExecutorFactory.create(parallel_mode, **executor_kwargs)` → `NoneExecutor` (single) / `DDPExecutor` (distributed)
|
||||
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`
|
||||
8. **Checkpoint**: `Checkpoint` saves/loads safetensors + metadata (rank-0 only), extra state saved as `{key}.pt`
|
||||
9. **Scheduler**: `SchedulerFactory` creates `CosineScheduler`/`SGDRScheduler`
|
||||
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-05-17
|
||||
> Document Update Time: 2026-05-24
|
||||
|
||||
@@ -53,7 +53,9 @@
|
||||
| Parameter | Description | Default |
|
||||
|-----------|-------------|---------|
|
||||
| `--nprocs` | Number of GPUs / processes | 1 |
|
||||
| `--parallel_mode` | Parallel strategy (`none` or `ddp`) | none |
|
||||
| `--device_type` | Device type | cuda |
|
||||
| `--start_method` | Multiprocessing start method (`spawn`, `fork`, `forkserver`) | spawn |
|
||||
|
||||
### Strategy-specific
|
||||
|
||||
@@ -94,4 +96,4 @@ nohup python scripts/tools/train.py \
|
||||
|
||||
---
|
||||
|
||||
> Document Update Time: 2026-05-17
|
||||
> Document Update Time: 2026-05-24
|
||||
+14
-13
@@ -72,17 +72,18 @@ on_train_begin
|
||||
on_epoch_begin
|
||||
for batch in dataloader:
|
||||
on_batch_begin
|
||||
loss = strategy(batch)
|
||||
(loss / grad_accum_steps).backward()
|
||||
iteration += 1
|
||||
with executor.accumulate(model):
|
||||
loss = strategy(batch)
|
||||
(loss / grad_accum_steps).backward()
|
||||
iteration += 1
|
||||
on_batch_end
|
||||
|
||||
if iteration % grad_accum_steps == 0:
|
||||
on_step_begin
|
||||
if executor.sync_gradients:
|
||||
on_optimizer_step
|
||||
optimizer.step()
|
||||
optimizer.zero_grad()
|
||||
on_step_end
|
||||
scheduler.step()
|
||||
|
||||
scheduler.step() # called every iteration
|
||||
on_epoch_end
|
||||
on_train_end
|
||||
```
|
||||
@@ -92,12 +93,11 @@ on_train_end
|
||||
| Hook | Fires | Default callback |
|
||||
|------|-------|-----------------|
|
||||
| `on_train_begin` | Before training starts | `GradientCheckpointingCallback` |
|
||||
| `on_step_begin` | Every accumulation window | `GradientClippingCallback` |
|
||||
| `on_optimizer_step` | Every accumulation window | `GradientClippingCallback`, `ValidationCallback` |
|
||||
| `on_batch_end` | Every batch | `CheckpointCallback`, `MetricLoggerCallback`, `ProgressBarCallback` |
|
||||
| `on_step_end` | Every accumulation window | `ValidationCallback` |
|
||||
| `on_train_end` | Training ends | `CheckpointCallback`, `MetricLoggerCallback` (final save) |
|
||||
|
||||
Default callbacks: `gradient_checkpointing` (activation checkpointing, optional), `progress_bar` (tqdm), `checkpoint` (safetensors, rank-0), `metric_logger` (JSONL, rank-0), `gradient_clipping`, `validation` (periodic validation on val_dataset).
|
||||
Default callbacks (in order): `gradient_checkpointing` (activation checkpointing, optional), `checkpoint` (safetensors, rank-0), `metric_logger` (JSONL, rank-0), `progress_bar` (tqdm), `gradient_clipping`, `validation` (periodic validation on val_dataset).
|
||||
|
||||
## Strategies
|
||||
|
||||
@@ -171,7 +171,7 @@ Callback wraps each `DecoderBlock.forward` with `torch.utils.checkpoint.checkpoi
|
||||
|
||||
```
|
||||
Checkpoint(state_dict, epoch, iteration, extra, meta)
|
||||
├── save(save_dir) rank-0 only: meta.json (includes training config) + state_dict.safetensors + optional extra.pt
|
||||
├── save(save_dir) rank-0 only: meta.json (includes training config) + state_dict.safetensors + optional optimizer.pt / scheduler.pt
|
||||
└── load(save_dir) broadcasts metadata from rank-0
|
||||
```
|
||||
|
||||
@@ -190,7 +190,8 @@ context = (
|
||||
```
|
||||
|
||||
- Loads checkpoint weights if provided
|
||||
- Wraps model with `parallel_wrapper` if `nprocs > 1`
|
||||
- Creates executor via `ExecutorFactory.create(parallel_mode, **executor_kwargs)`
|
||||
- Calls `executor.prepare(model, optimizer, dataloader, scheduler)` for model distribution (e.g. DDP) + gradient accumulation wrappers
|
||||
- Creates `ResumableDistributedSampler` for shuffle+resume
|
||||
- Builds strategy via `StrategyFactory.create(train_type, ...)`
|
||||
|
||||
@@ -222,4 +223,4 @@ nohup python scripts/tools/train.py \
|
||||
|
||||
Full parameter reference at [params.md](params.md).
|
||||
|
||||
> Document Update Time: 2026-05-17
|
||||
> Document Update Time: 2026-05-24
|
||||
|
||||
@@ -49,6 +49,7 @@ class AutoRegressiveLMConfig(BaseModelConfig):
|
||||
|
||||
max_len: Optional[int] = None
|
||||
rope_theta: Optional[float] = None
|
||||
rope_scaling: Optional[dict] = None
|
||||
|
||||
attn_type: str = "gqa"
|
||||
n_heads: Optional[int] = None
|
||||
@@ -80,6 +81,7 @@ class EncoderConfig(BaseModelConfig):
|
||||
|
||||
max_len: Optional[int] = None
|
||||
rope_theta: Optional[float] = None
|
||||
rope_scaling: Optional[dict] = None
|
||||
|
||||
n_heads: Optional[int] = None
|
||||
n_kv_heads: Optional[int] = None
|
||||
|
||||
@@ -7,6 +7,7 @@ from torch.optim.lr_scheduler import LRScheduler
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from astrai.config.base import BaseConfig
|
||||
from astrai.model.components.lora import LoRAConfig
|
||||
|
||||
|
||||
def required(**kw):
|
||||
@@ -56,6 +57,12 @@ class TrainConfig(BaseConfig):
|
||||
default=5000, metadata={"help": "Number of iterations between checkpoints."}
|
||||
)
|
||||
|
||||
# lora setting
|
||||
lora: Optional[LoRAConfig] = field(
|
||||
default=None,
|
||||
metadata={"help": "LoRA config. None means full fine-tuning."},
|
||||
)
|
||||
|
||||
# metric setting
|
||||
log_dir: str = field(
|
||||
default="./checkpoint/logs", metadata={"help": "Directory for metric logs."}
|
||||
@@ -95,11 +102,9 @@ class TrainConfig(BaseConfig):
|
||||
master_port: str = field(
|
||||
default="29500", metadata={"help": "Master port for distributed training."}
|
||||
)
|
||||
parallel_wrapper: Optional[Callable] = field(
|
||||
default=None, metadata={"help": "Parallel function for training."}
|
||||
)
|
||||
state_dict_fn: Optional[Callable] = field(
|
||||
default=None, metadata={"help": "Parallel function for state dict saving."}
|
||||
parallel_mode: str = field(
|
||||
default="none",
|
||||
metadata={"help": "Parallel strategy: none, ddp, fsdp."},
|
||||
)
|
||||
start_method: str = field(
|
||||
default="spawn",
|
||||
@@ -118,6 +123,10 @@ class TrainConfig(BaseConfig):
|
||||
metadata={"help": "Number of optimizer steps between validation runs."},
|
||||
)
|
||||
|
||||
executor_kwargs: dict = field(
|
||||
default_factory=dict,
|
||||
metadata={"help": "Extra kwargs passed to ExecutorFactory.create()."},
|
||||
)
|
||||
extra_kwargs: dict = field(
|
||||
default_factory=dict, metadata={"help": "Other arguments."}
|
||||
)
|
||||
|
||||
@@ -43,6 +43,7 @@ class ResumableDistributedSampler(Sampler[int]):
|
||||
offset = 0 if drop_last else self.num_replicas - 1
|
||||
self.num_samples_per_replica = (self.num_samples + offset) // self.num_replicas
|
||||
self.total_size = self.num_samples_per_replica * self.num_replicas
|
||||
self.iter = self.iter % self.num_samples_per_replica
|
||||
|
||||
self._indices = None
|
||||
|
||||
@@ -74,5 +75,10 @@ class ResumableDistributedSampler(Sampler[int]):
|
||||
self.epoch += 1
|
||||
self._indices = None
|
||||
|
||||
@property
|
||||
def _remaining(self):
|
||||
remaining = self.num_samples_per_replica - self.iter
|
||||
return max(remaining, 0)
|
||||
|
||||
def __len__(self):
|
||||
return self.num_samples_per_replica
|
||||
return self._remaining
|
||||
|
||||
@@ -1,25 +1,27 @@
|
||||
"""Inference module for continuous batching.
|
||||
|
||||
Layers:
|
||||
- core/: Core inference loop (cache, executor, scheduler, task)
|
||||
- api/: HTTP protocol handlers (OpenAI, Anthropic)
|
||||
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
|
||||
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy)
|
||||
- core/: Core inference loop (cache, executor, scheduler, task)
|
||||
- api/: HTTP orchestration (ProtocolHandler, server)
|
||||
- protocols/: Response builders (OpenAI, Anthropic)
|
||||
- transport/: SSE transport utilities
|
||||
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
|
||||
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy)
|
||||
"""
|
||||
|
||||
from astrai.inference.api import (
|
||||
AnthropicHandler,
|
||||
AnthropicMessage,
|
||||
ChatCompletionRequest,
|
||||
ChatMessage,
|
||||
GenContext,
|
||||
MessagesRequest,
|
||||
OpenAIHandler,
|
||||
ProtocolHandler,
|
||||
StopChecker,
|
||||
StreamContext,
|
||||
app,
|
||||
run_server,
|
||||
)
|
||||
from astrai.inference.api.anthropic import AnthropicResponseBuilder
|
||||
from astrai.inference.api.openai import OpenAIResponseBuilder
|
||||
from astrai.inference.core import (
|
||||
STOP,
|
||||
Allocator,
|
||||
@@ -36,10 +38,7 @@ from astrai.inference.core import (
|
||||
TaskTable,
|
||||
page_hash,
|
||||
)
|
||||
from astrai.inference.engine import (
|
||||
GenerationRequest,
|
||||
InferenceEngine,
|
||||
)
|
||||
from astrai.inference.engine import GenerationRequest, InferenceEngine
|
||||
from astrai.inference.sample import (
|
||||
BaseSamplingStrategy,
|
||||
SamplingPipeline,
|
||||
@@ -50,17 +49,14 @@ from astrai.inference.sample import (
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
# Engine / Requests
|
||||
"InferenceEngine",
|
||||
"GenerationRequest",
|
||||
# Core scheduler
|
||||
"InferenceScheduler",
|
||||
"Executor",
|
||||
"STOP",
|
||||
"Task",
|
||||
"TaskManager",
|
||||
"TaskStatus",
|
||||
# Core cache
|
||||
"Allocator",
|
||||
"KVCache",
|
||||
"KvcacheView",
|
||||
@@ -69,20 +65,17 @@ __all__ = [
|
||||
"Storage",
|
||||
"TaskTable",
|
||||
"page_hash",
|
||||
# Sampling (Strategy pattern)
|
||||
"sample",
|
||||
"BaseSamplingStrategy",
|
||||
"TemperatureStrategy",
|
||||
"TopKStrategy",
|
||||
"TopPStrategy",
|
||||
"SamplingPipeline",
|
||||
# Protocol
|
||||
"ProtocolHandler",
|
||||
"StopChecker",
|
||||
"StreamContext",
|
||||
"AnthropicHandler",
|
||||
"OpenAIHandler",
|
||||
# Server
|
||||
"GenContext",
|
||||
"OpenAIResponseBuilder",
|
||||
"AnthropicResponseBuilder",
|
||||
"ChatMessage",
|
||||
"ChatCompletionRequest",
|
||||
"AnthropicMessage",
|
||||
|
||||
@@ -1,12 +1,6 @@
|
||||
"""Inference API: protocol handlers and FastAPI server."""
|
||||
"""Inference API: protocol handler, stop checker, and FastAPI server."""
|
||||
|
||||
from astrai.inference.api.protocol import (
|
||||
AnthropicHandler,
|
||||
OpenAIHandler,
|
||||
ProtocolHandler,
|
||||
StopChecker,
|
||||
StreamContext,
|
||||
)
|
||||
from astrai.inference.api.protocol import GenContext, ProtocolHandler, StopChecker
|
||||
from astrai.inference.api.server import (
|
||||
AnthropicMessage,
|
||||
ChatCompletionRequest,
|
||||
@@ -17,11 +11,9 @@ from astrai.inference.api.server import (
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"AnthropicHandler",
|
||||
"OpenAIHandler",
|
||||
"ProtocolHandler",
|
||||
"StopChecker",
|
||||
"StreamContext",
|
||||
"GenContext",
|
||||
"AnthropicMessage",
|
||||
"ChatCompletionRequest",
|
||||
"ChatMessage",
|
||||
|
||||
@@ -0,0 +1,140 @@
|
||||
"""Anthropic message completion response builder."""
|
||||
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Tuple, Union
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from astrai.inference.api.protocol import (
|
||||
GenContext,
|
||||
ResponseBuilder,
|
||||
StopInfo,
|
||||
sse_event,
|
||||
)
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
|
||||
|
||||
def _extract_text(content: Union[str, List[Dict[str, Any]]]) -> str:
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
for block in content:
|
||||
if isinstance(block, dict) and block.get("type") == "text":
|
||||
return block.get("text", "")
|
||||
return ""
|
||||
|
||||
|
||||
class AnthropicResponseBuilder(ResponseBuilder):
|
||||
def prepare(
|
||||
self, request: BaseModel, engine: InferenceEngine
|
||||
) -> Tuple[str, GenContext, List[str]]:
|
||||
messages: List[Dict[str, str]] = []
|
||||
system = getattr(request, "system", None)
|
||||
if system:
|
||||
messages.append({"role": "system", "content": system})
|
||||
for m in request.messages:
|
||||
text = _extract_text(m.content)
|
||||
if text:
|
||||
messages.append({"role": m.role, "content": text})
|
||||
prompt = engine.tokenizer.apply_chat_template(messages, tokenize=False)
|
||||
ctx = GenContext(
|
||||
resp_id=f"msg_{uuid.uuid4().hex[:24]}",
|
||||
created=0,
|
||||
model=request.model,
|
||||
prompt_tokens=0,
|
||||
)
|
||||
stop_sequences = getattr(request, "stop_sequences", None) or []
|
||||
return prompt, ctx, stop_sequences
|
||||
|
||||
def format_stream_start(self, ctx: GenContext) -> List[str]:
|
||||
return [
|
||||
sse_event(
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": ctx.resp_id,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": ctx.model,
|
||||
"content": [],
|
||||
"usage": {"input_tokens": ctx.prompt_tokens},
|
||||
},
|
||||
},
|
||||
event="message_start",
|
||||
),
|
||||
sse_event(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 0,
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
},
|
||||
event="content_block_start",
|
||||
),
|
||||
]
|
||||
|
||||
def format_chunk(self, token: str) -> str:
|
||||
return sse_event(
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": token},
|
||||
},
|
||||
event="content_block_delta",
|
||||
)
|
||||
|
||||
def format_stream_end(self, ctx: GenContext, stop: StopInfo) -> List[str]:
|
||||
events: List[str] = []
|
||||
if stop.matched:
|
||||
trimmed = stop.body[: stop.body.rfind(stop.matched)]
|
||||
unyielded = trimmed[len(stop.yielded) :]
|
||||
if unyielded:
|
||||
events.append(
|
||||
sse_event(
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": unyielded},
|
||||
},
|
||||
event="content_block_delta",
|
||||
)
|
||||
)
|
||||
events.append(
|
||||
sse_event(
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
event="content_block_stop",
|
||||
)
|
||||
)
|
||||
events.append(
|
||||
sse_event(
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {
|
||||
"stop_reason": "stop_sequence" if stop.matched else "end_turn",
|
||||
"stop_sequence": stop.matched,
|
||||
},
|
||||
"usage": {"output_tokens": ctx.completion_tokens},
|
||||
},
|
||||
event="message_delta",
|
||||
)
|
||||
)
|
||||
events.append(sse_event({"type": "message_stop"}, event="message_stop"))
|
||||
return events
|
||||
|
||||
def format_response(
|
||||
self, ctx: GenContext, content: str, stop: StopInfo
|
||||
) -> Dict[str, Any]:
|
||||
if stop.matched:
|
||||
content = content[: content.rfind(stop.matched)]
|
||||
return {
|
||||
"id": ctx.resp_id,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": ctx.model,
|
||||
"content": [{"type": "text", "text": content}],
|
||||
"stop_reason": "stop_sequence" if stop.matched else "end_turn",
|
||||
"stop_sequence": stop.matched,
|
||||
"usage": {
|
||||
"input_tokens": ctx.prompt_tokens,
|
||||
"output_tokens": ctx.completion_tokens,
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,111 @@
|
||||
"""OpenAI chat completion response builder."""
|
||||
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Tuple
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from astrai.inference.api.protocol import (
|
||||
GenContext,
|
||||
ResponseBuilder,
|
||||
StopInfo,
|
||||
sse_event,
|
||||
)
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
|
||||
|
||||
class OpenAIResponseBuilder(ResponseBuilder):
|
||||
def prepare(
|
||||
self, request: BaseModel, engine: InferenceEngine
|
||||
) -> Tuple[str, GenContext, List[str]]:
|
||||
messages = [{"role": m.role, "content": m.content} for m in request.messages]
|
||||
prompt = engine.tokenizer.apply_chat_template(messages, tokenize=False)
|
||||
|
||||
self._resp_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
||||
self._model = request.model
|
||||
|
||||
ctx = GenContext(
|
||||
resp_id=self._resp_id,
|
||||
created=0,
|
||||
model=self._model,
|
||||
prompt_tokens=0,
|
||||
)
|
||||
stop = request.stop
|
||||
stop_sequences = (
|
||||
[] if stop is None else [stop] if isinstance(stop, str) else stop
|
||||
)
|
||||
return prompt, ctx, stop_sequences
|
||||
|
||||
def format_stream_start(self, ctx: GenContext) -> List[str]:
|
||||
return [
|
||||
sse_event(
|
||||
{
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": ctx.created,
|
||||
"model": self._model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"role": "assistant"},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
]
|
||||
|
||||
def format_chunk(self, token: str) -> str:
|
||||
return sse_event(
|
||||
{
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 0,
|
||||
"model": self._model,
|
||||
"choices": [
|
||||
{"index": 0, "delta": {"content": token}, "finish_reason": None}
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
def format_stream_end(self, ctx: GenContext, stop: StopInfo) -> List[str]:
|
||||
return [
|
||||
sse_event(
|
||||
{
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": ctx.created,
|
||||
"model": self._model,
|
||||
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
|
||||
}
|
||||
),
|
||||
sse_event(
|
||||
{
|
||||
"prompt_tokens": ctx.prompt_tokens,
|
||||
"completion_tokens": ctx.completion_tokens,
|
||||
"total_tokens": ctx.prompt_tokens + ctx.completion_tokens,
|
||||
}
|
||||
),
|
||||
]
|
||||
|
||||
def format_response(
|
||||
self, ctx: GenContext, content: str, stop: StopInfo
|
||||
) -> Dict[str, Any]:
|
||||
return {
|
||||
"id": self._resp_id,
|
||||
"object": "chat.completion",
|
||||
"created": ctx.created,
|
||||
"model": self._model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": content},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": ctx.prompt_tokens,
|
||||
"completion_tokens": ctx.completion_tokens,
|
||||
"total_tokens": ctx.prompt_tokens + ctx.completion_tokens,
|
||||
},
|
||||
}
|
||||
@@ -1,15 +1,13 @@
|
||||
"""Protocol handlers for OpenAI and Anthropic chat completion APIs.
|
||||
"""Orchestration layer: ProtocolHandler, StopChecker, GenContext, StopInfo, ResponseBuilder, SSE utils.
|
||||
|
||||
Template Method + Builder patterns eliminate the 45% code duplication between
|
||||
stream/non-stream branches and across protocol adapters.
|
||||
ProtocolHandler orchestrates the async generation loop and delegates
|
||||
protocol-specific formatting to a ResponseBuilder.
|
||||
"""
|
||||
|
||||
import json
|
||||
import time
|
||||
import uuid
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
from typing import Any, AsyncGenerator, Dict, List, Optional, Tuple, Union
|
||||
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import BaseModel
|
||||
@@ -17,7 +15,7 @@ from pydantic import BaseModel
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
|
||||
|
||||
def _sse_event(data: Dict[str, Any], event: Optional[str] = None) -> str:
|
||||
def sse_event(data: Dict[str, Any], event: Optional[str] = None) -> str:
|
||||
lines: List[str] = []
|
||||
if event:
|
||||
lines.append(f"event: {event}")
|
||||
@@ -26,22 +24,28 @@ def _sse_event(data: Dict[str, Any], event: Optional[str] = None) -> str:
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _sse_done() -> str:
|
||||
def sse_done() -> str:
|
||||
return "data: [DONE]\n\n"
|
||||
|
||||
|
||||
@dataclass
|
||||
class StreamContext:
|
||||
"""Shared state across the streaming generation lifecycle."""
|
||||
class GenContext:
|
||||
"""Per-generation metadata passed to builder format methods."""
|
||||
|
||||
resp_id: str
|
||||
created: int
|
||||
model: str
|
||||
prompt_tokens: int
|
||||
completion_tokens: int = 0
|
||||
accumulated: str = ""
|
||||
stop_matched: Optional[str] = None
|
||||
last_yield_trimmed: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class StopInfo:
|
||||
"""Stop-check result passed to format_stream_end / format_response."""
|
||||
|
||||
matched: Optional[str] = None
|
||||
body: str = ""
|
||||
yielded: str = ""
|
||||
|
||||
|
||||
class StopChecker:
|
||||
@@ -56,95 +60,60 @@ class StopChecker:
|
||||
return seq
|
||||
return None
|
||||
|
||||
def trim(self, text: str, matched: str) -> str:
|
||||
idx = text.rfind(matched)
|
||||
return text[:idx] if idx != -1 else text
|
||||
|
||||
@property
|
||||
def has_sequences(self) -> bool:
|
||||
return len(self._sequences) > 0
|
||||
class ResponseBuilder(ABC):
|
||||
"""Interface for protocol-specific response formatting.
|
||||
|
||||
|
||||
class ProtocolHandler(ABC):
|
||||
"""Template-method base for API protocol handlers.
|
||||
|
||||
Subclasses implement format hooks; the base class orchestrates the
|
||||
generate-async loop and SSE/JSON response construction.
|
||||
|
||||
Lifecycle::
|
||||
|
||||
handle()
|
||||
├─ build_prompt() # protocol-specific prompt assembly
|
||||
├─ create_response_id() # unique response identifier
|
||||
├─ [stream]
|
||||
│ ├─ format_stream_start()
|
||||
│ ├─ format_stream_token() × N
|
||||
│ │ └─ on_token() hook for stop-sequence interception
|
||||
│ └─ format_stream_end()
|
||||
└─ [non-stream]
|
||||
├─ (accumulate tokens)
|
||||
└─ format_non_stream_response()
|
||||
A new protocol requires one concrete builder implementing 6 methods.
|
||||
"""
|
||||
|
||||
request_model: type[BaseModel]
|
||||
@abstractmethod
|
||||
def prepare(
|
||||
self, request: BaseModel, engine: InferenceEngine
|
||||
) -> Tuple[str, GenContext, List[str]]:
|
||||
"""Return (prompt, ctx, stop_sequences) for a generation request."""
|
||||
|
||||
def __init__(self, request: BaseModel, engine: InferenceEngine):
|
||||
@abstractmethod
|
||||
def format_stream_start(self, ctx: GenContext) -> List[str]:
|
||||
"""SSE events that open the stream."""
|
||||
|
||||
@abstractmethod
|
||||
def format_chunk(self, token: str) -> str:
|
||||
"""SSE event for a single generated token."""
|
||||
|
||||
@abstractmethod
|
||||
def format_stream_end(self, ctx: GenContext, stop: StopInfo) -> List[str]:
|
||||
"""SSE events that close the stream."""
|
||||
|
||||
@abstractmethod
|
||||
def format_response(
|
||||
self, ctx: GenContext, content: str, stop: StopInfo
|
||||
) -> Dict[str, Any]:
|
||||
"""JSON response body for non-streaming mode."""
|
||||
|
||||
|
||||
class ProtocolHandler:
|
||||
"""Orchestrates the generation loop, delegates formatting to a builder.
|
||||
|
||||
Usage::
|
||||
|
||||
handler = ProtocolHandler(request, engine, OpenAIResponseBuilder())
|
||||
response = await handler.handle()
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, request: BaseModel, engine: InferenceEngine, builder: ResponseBuilder
|
||||
):
|
||||
self.request = request
|
||||
self.engine = engine
|
||||
|
||||
@abstractmethod
|
||||
def build_prompt(self) -> str:
|
||||
"""Build the full prompt string from the request messages."""
|
||||
|
||||
@abstractmethod
|
||||
def create_response_id(self) -> str:
|
||||
"""Generate a unique response ID following the protocol convention."""
|
||||
|
||||
@abstractmethod
|
||||
def format_stream_start(self, ctx: StreamContext) -> List[str]:
|
||||
"""Yield SSE events that open the stream (role marker, metadata)."""
|
||||
|
||||
@abstractmethod
|
||||
def format_stream_token(self, ctx: StreamContext, token: str) -> str:
|
||||
"""Yield an SSE event for a single generated token."""
|
||||
|
||||
@abstractmethod
|
||||
def format_stream_end(self, ctx: StreamContext) -> List[str]:
|
||||
"""Yield SSE events that close the stream (finish reason, usage stats)."""
|
||||
|
||||
@abstractmethod
|
||||
def format_non_stream_response(
|
||||
self, ctx: StreamContext, content: str
|
||||
) -> Dict[str, Any]:
|
||||
"""Build the JSON response body for non-streaming mode."""
|
||||
|
||||
def get_stop_sequences(self) -> List[str]:
|
||||
return []
|
||||
|
||||
def create_stop_checker(self) -> StopChecker:
|
||||
return StopChecker(self.get_stop_sequences())
|
||||
|
||||
def on_token(
|
||||
self, ctx: StreamContext, token: str, stop_checker: StopChecker
|
||||
) -> Optional[str]:
|
||||
"""Hook after each token is appended to accumulated.
|
||||
|
||||
Return a matched stop-sequence string to break the loop,
|
||||
or None to continue.
|
||||
|
||||
"""
|
||||
return None
|
||||
self.builder = builder
|
||||
|
||||
async def handle(self) -> Union[StreamingResponse, Dict[str, Any]]:
|
||||
ctx = StreamContext(
|
||||
resp_id=self.create_response_id(),
|
||||
created=int(time.time()),
|
||||
model=self.request.model,
|
||||
prompt_tokens=self._count_prompt_tokens(),
|
||||
)
|
||||
prompt, ctx, stop_sequences = self.builder.prepare(self.request, self.engine)
|
||||
ctx.prompt_tokens = len(self.engine.tokenizer.encode(prompt))
|
||||
|
||||
agen = self.engine.generate_async(
|
||||
prompt=self.build_prompt(),
|
||||
prompt=prompt,
|
||||
max_tokens=self.request.max_tokens,
|
||||
temperature=self.request.temperature,
|
||||
top_p=self.request.top_p,
|
||||
@@ -152,33 +121,37 @@ class ProtocolHandler(ABC):
|
||||
)
|
||||
|
||||
if self.request.stream:
|
||||
return self._handle_stream(agen, ctx)
|
||||
return self._handle_stream(agen, ctx, stop_sequences)
|
||||
else:
|
||||
return await self._handle_non_stream(agen, ctx)
|
||||
return await self._handle_non_stream(agen, ctx, stop_sequences)
|
||||
|
||||
def _count_prompt_tokens(self) -> int:
|
||||
return len(self.engine.tokenizer.encode(self.build_prompt()))
|
||||
|
||||
def _handle_stream(self, agen, ctx: StreamContext) -> StreamingResponse:
|
||||
stop_checker = self.create_stop_checker()
|
||||
def _handle_stream(
|
||||
self, agen: AsyncGenerator, ctx: GenContext, stop_sequences: List[str]
|
||||
) -> StreamingResponse:
|
||||
checker = StopChecker(stop_sequences)
|
||||
|
||||
async def event_stream():
|
||||
for event in self.format_stream_start(ctx):
|
||||
for event in self.builder.format_stream_start(ctx):
|
||||
yield event
|
||||
|
||||
body = ""
|
||||
yielded = ""
|
||||
matched = None
|
||||
async for token in agen:
|
||||
ctx.completion_tokens += 1
|
||||
ctx.accumulated += token
|
||||
body += token
|
||||
|
||||
matched = self.on_token(ctx, token, stop_checker)
|
||||
matched = checker.check(body)
|
||||
if matched:
|
||||
break
|
||||
|
||||
yield self.format_stream_token(ctx, token)
|
||||
yield self.builder.format_chunk(token)
|
||||
yielded += token
|
||||
|
||||
for event in self.format_stream_end(ctx):
|
||||
stop = StopInfo(matched=matched, body=body, yielded=yielded)
|
||||
for event in self.builder.format_stream_end(ctx, stop):
|
||||
yield event
|
||||
yield _sse_done()
|
||||
yield sse_done()
|
||||
|
||||
return StreamingResponse(
|
||||
event_stream(),
|
||||
@@ -186,260 +159,23 @@ class ProtocolHandler(ABC):
|
||||
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
||||
)
|
||||
|
||||
async def _handle_non_stream(self, agen, ctx: StreamContext) -> Dict[str, Any]:
|
||||
stop_checker = self.create_stop_checker()
|
||||
async def _handle_non_stream(
|
||||
self, agen: AsyncGenerator, ctx: GenContext, stop_sequences: List[str]
|
||||
) -> Dict[str, Any]:
|
||||
checker = StopChecker(stop_sequences)
|
||||
chunks: List[str] = []
|
||||
body = ""
|
||||
matched = None
|
||||
|
||||
async for token in agen:
|
||||
ctx.completion_tokens += 1
|
||||
ctx.accumulated += token
|
||||
chunks.append(token)
|
||||
body += token
|
||||
|
||||
matched = self.on_token(ctx, token, stop_checker)
|
||||
matched = checker.check(body)
|
||||
if matched:
|
||||
break
|
||||
|
||||
content = "".join(chunks)
|
||||
return self.format_non_stream_response(ctx, content)
|
||||
|
||||
|
||||
def _extract_text_content(content: Union[str, List[Dict[str, Any]]]) -> str:
|
||||
"""Extract plain text from an Anthropic content block (string or list)."""
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
for block in content:
|
||||
if isinstance(block, dict) and block.get("type") == "text":
|
||||
return block.get("text", "")
|
||||
return ""
|
||||
|
||||
|
||||
class OpenAIHandler(ProtocolHandler):
|
||||
"""OpenAI-compatible /v1/chat/completions handler."""
|
||||
|
||||
def build_prompt(self) -> str:
|
||||
messages = [
|
||||
{"role": m.role, "content": m.content} for m in self.request.messages
|
||||
]
|
||||
return self.engine.tokenizer.apply_chat_template(messages, tokenize=False)
|
||||
|
||||
def create_response_id(self) -> str:
|
||||
return f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
||||
|
||||
def get_stop_sequences(self) -> List[str]:
|
||||
stop = self.request.stop
|
||||
if stop is None:
|
||||
return []
|
||||
return [stop] if isinstance(stop, str) else stop
|
||||
|
||||
def on_token(
|
||||
self, ctx: StreamContext, token: str, stop_checker: StopChecker
|
||||
) -> Optional[str]:
|
||||
return stop_checker.check(ctx.accumulated)
|
||||
|
||||
def format_stream_start(self, ctx: StreamContext) -> List[str]:
|
||||
return [
|
||||
_sse_event(
|
||||
{
|
||||
"id": ctx.resp_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": ctx.created,
|
||||
"model": ctx.model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"role": "assistant"},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
]
|
||||
|
||||
def format_stream_token(self, ctx: StreamContext, token: str) -> str:
|
||||
return _sse_event(
|
||||
{
|
||||
"id": ctx.resp_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": ctx.created,
|
||||
"model": ctx.model,
|
||||
"choices": [
|
||||
{"index": 0, "delta": {"content": token}, "finish_reason": None}
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
def format_stream_end(self, ctx: StreamContext) -> List[str]:
|
||||
return [
|
||||
_sse_event(
|
||||
{
|
||||
"id": ctx.resp_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": ctx.created,
|
||||
"model": ctx.model,
|
||||
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
|
||||
}
|
||||
),
|
||||
_sse_event(
|
||||
{
|
||||
"prompt_tokens": ctx.prompt_tokens,
|
||||
"completion_tokens": ctx.completion_tokens,
|
||||
"total_tokens": ctx.prompt_tokens + ctx.completion_tokens,
|
||||
}
|
||||
),
|
||||
]
|
||||
|
||||
def format_non_stream_response(
|
||||
self, ctx: StreamContext, content: str
|
||||
) -> Dict[str, Any]:
|
||||
return {
|
||||
"id": ctx.resp_id,
|
||||
"object": "chat.completion",
|
||||
"created": ctx.created,
|
||||
"model": ctx.model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": content},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": ctx.prompt_tokens,
|
||||
"completion_tokens": ctx.completion_tokens,
|
||||
"total_tokens": ctx.prompt_tokens + ctx.completion_tokens,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class AnthropicHandler(ProtocolHandler):
|
||||
"""Anthropic-compatible /v1/messages handler."""
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self._yielded = ""
|
||||
|
||||
def build_prompt(self) -> str:
|
||||
messages: List[Dict[str, str]] = []
|
||||
system = getattr(self.request, "system", None)
|
||||
if system:
|
||||
messages.append({"role": "system", "content": system})
|
||||
for m in self.request.messages:
|
||||
content = _extract_text_content(m.content)
|
||||
if content:
|
||||
messages.append({"role": m.role, "content": content})
|
||||
return self.engine.tokenizer.apply_chat_template(messages, tokenize=False)
|
||||
|
||||
def create_response_id(self) -> str:
|
||||
return f"msg_{uuid.uuid4().hex[:24]}"
|
||||
|
||||
def get_stop_sequences(self) -> List[str]:
|
||||
return getattr(self.request, "stop_sequences", None) or []
|
||||
|
||||
def on_token(
|
||||
self, ctx: StreamContext, token: str, stop_checker: StopChecker
|
||||
) -> Optional[str]:
|
||||
matched = stop_checker.check(ctx.accumulated)
|
||||
if not matched:
|
||||
return None
|
||||
|
||||
ctx.stop_matched = matched
|
||||
trimmed = ctx.accumulated[: ctx.accumulated.rfind(matched)]
|
||||
unyielded = trimmed[len(self._yielded) :]
|
||||
if unyielded:
|
||||
ctx.last_yield_trimmed = unyielded
|
||||
return matched
|
||||
|
||||
def format_stream_start(self, ctx: StreamContext) -> List[str]:
|
||||
return [
|
||||
_sse_event(
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": ctx.resp_id,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": ctx.model,
|
||||
"content": [],
|
||||
"usage": {"input_tokens": ctx.prompt_tokens},
|
||||
},
|
||||
},
|
||||
event="message_start",
|
||||
),
|
||||
_sse_event(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 0,
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
},
|
||||
event="content_block_start",
|
||||
),
|
||||
]
|
||||
|
||||
def format_stream_token(self, ctx: StreamContext, token: str) -> str:
|
||||
self._yielded += token
|
||||
return _sse_event(
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": token},
|
||||
},
|
||||
event="content_block_delta",
|
||||
)
|
||||
|
||||
def format_stream_end(self, ctx: StreamContext) -> List[str]:
|
||||
matched = ctx.stop_matched
|
||||
events: List[str] = []
|
||||
last_yielded = ctx.last_yield_trimmed
|
||||
if last_yielded:
|
||||
events.append(
|
||||
_sse_event(
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": last_yielded},
|
||||
},
|
||||
event="content_block_delta",
|
||||
)
|
||||
)
|
||||
events.append(
|
||||
_sse_event(
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
event="content_block_stop",
|
||||
)
|
||||
)
|
||||
events.append(
|
||||
_sse_event(
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {
|
||||
"stop_reason": "stop_sequence" if matched else "end_turn",
|
||||
"stop_sequence": matched,
|
||||
},
|
||||
"usage": {"output_tokens": ctx.completion_tokens},
|
||||
},
|
||||
event="message_delta",
|
||||
)
|
||||
)
|
||||
events.append(_sse_event({"type": "message_stop"}, event="message_stop"))
|
||||
return events
|
||||
|
||||
def format_non_stream_response(
|
||||
self, ctx: StreamContext, content: str
|
||||
) -> Dict[str, Any]:
|
||||
matched = ctx.stop_matched
|
||||
if matched:
|
||||
content = content[: content.rfind(matched)]
|
||||
return {
|
||||
"id": ctx.resp_id,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": ctx.model,
|
||||
"content": [{"type": "text", "text": content}],
|
||||
"stop_reason": "stop_sequence" if matched else "end_turn",
|
||||
"stop_sequence": matched,
|
||||
"usage": {
|
||||
"input_tokens": ctx.prompt_tokens,
|
||||
"output_tokens": ctx.completion_tokens,
|
||||
},
|
||||
}
|
||||
stop = StopInfo(matched=matched, body=body)
|
||||
return self.builder.format_response(ctx, content, stop)
|
||||
|
||||
@@ -15,7 +15,9 @@ import uvicorn
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from astrai.inference.api.protocol import AnthropicHandler, OpenAIHandler
|
||||
from astrai.inference.api.anthropic import AnthropicResponseBuilder
|
||||
from astrai.inference.api.openai import OpenAIResponseBuilder
|
||||
from astrai.inference.api.protocol import ProtocolHandler
|
||||
from astrai.inference.engine import InferenceEngine
|
||||
from astrai.model import AutoModel
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
@@ -133,14 +135,14 @@ async def get_stats():
|
||||
@app.post("/v1/chat/completions")
|
||||
async def chat_completion(request: ChatCompletionRequest):
|
||||
engine = _get_engine()
|
||||
handler = OpenAIHandler(request, engine)
|
||||
handler = ProtocolHandler(request, engine, OpenAIResponseBuilder())
|
||||
return await handler.handle()
|
||||
|
||||
|
||||
@app.post("/v1/messages")
|
||||
async def create_message(request: MessagesRequest):
|
||||
engine = _get_engine()
|
||||
handler = AnthropicHandler(request, engine)
|
||||
handler = ProtocolHandler(request, engine, AnthropicResponseBuilder())
|
||||
return await handler.handle()
|
||||
|
||||
|
||||
|
||||
@@ -108,7 +108,10 @@ class InferenceScheduler:
|
||||
continue
|
||||
|
||||
to_prefill = [
|
||||
t for t in self._task_mgr.get_active_tasks() if t.output_tokens == 0
|
||||
t
|
||||
for t in self._task_mgr.get_active_tasks()
|
||||
if t.output_tokens == 0
|
||||
and self._page_cache.task_cached(t.task_id) < len(t.prompt_ids)
|
||||
]
|
||||
if to_prefill:
|
||||
for t in to_prefill:
|
||||
@@ -156,11 +159,15 @@ class InferenceScheduler:
|
||||
t.output_ids.append(ntok)
|
||||
t.output_tokens += 1
|
||||
pos = t.input_tokens + t.output_tokens
|
||||
self._page_cache.task_extend(t.task_id, pos)
|
||||
extend_ok = self._page_cache.task_extend(t.task_id, pos)
|
||||
if t.stream_callback:
|
||||
t.stream_callback(
|
||||
self._task_mgr.tokenizer.decode([ntok])
|
||||
)
|
||||
if not extend_ok:
|
||||
t.status = TaskStatus.ABORTED
|
||||
if t.stream_callback:
|
||||
t.stream_callback(STOP)
|
||||
|
||||
for t in valid:
|
||||
if t.is_finished(stop_ids):
|
||||
@@ -173,6 +180,9 @@ class InferenceScheduler:
|
||||
if task.stream_callback:
|
||||
task.stream_callback(STOP)
|
||||
self._page_cache.task_free(task.task_id)
|
||||
for task in self._task_mgr.get_waiting_tasks():
|
||||
if task.stream_callback:
|
||||
task.stream_callback(STOP)
|
||||
self._task_mgr.clear_queues()
|
||||
raise
|
||||
|
||||
|
||||
@@ -193,6 +193,10 @@ class TaskManager:
|
||||
with self._lock:
|
||||
return list(self.active_tasks)
|
||||
|
||||
def get_waiting_tasks(self) -> List[Task]:
|
||||
with self._lock:
|
||||
return list(self.waiting_queue)
|
||||
|
||||
def clear_queues(self) -> None:
|
||||
with self._lock:
|
||||
self.waiting_queue.clear()
|
||||
|
||||
@@ -13,17 +13,6 @@ from astrai.inference.core.task import STOP
|
||||
from astrai.tokenize import AutoTokenizer
|
||||
|
||||
|
||||
def _validate_sampling_params(
|
||||
top_k: int, top_p: float, temperature: float, max_tokens: Optional[int] = None
|
||||
):
|
||||
if not (isinstance(top_k, int) and top_k >= 0):
|
||||
raise ValueError("top_k must be a non-negative integer")
|
||||
if not (0.0 <= top_p <= 1.0):
|
||||
raise ValueError("top_p must be a float between 0.0 and 1.0")
|
||||
if not (isinstance(temperature, (int, float)) and temperature >= 0):
|
||||
raise ValueError("temperature must be a non-negative number")
|
||||
|
||||
|
||||
class GenerateResult:
|
||||
"""Thread-safe token accumulator for streaming and non-streaming modes."""
|
||||
|
||||
@@ -86,7 +75,12 @@ class GenerationRequest:
|
||||
max_tokens: Optional[int] = None,
|
||||
stream: bool = False,
|
||||
):
|
||||
_validate_sampling_params(top_k, top_p, temperature, max_tokens)
|
||||
if not (isinstance(top_k, int) and top_k >= 0):
|
||||
raise ValueError("top_k must be a non-negative integer")
|
||||
if not (0.0 <= top_p <= 1.0):
|
||||
raise ValueError("top_p must be a float between 0.0 and 1.0")
|
||||
if not (isinstance(temperature, (int, float)) and temperature >= 0):
|
||||
raise ValueError("temperature must be a non-negative number")
|
||||
|
||||
self.messages = messages
|
||||
self.top_k = top_k
|
||||
@@ -137,7 +131,6 @@ class InferenceEngine:
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
) -> Union[Generator, str, List[str]]:
|
||||
_validate_sampling_params(top_k, top_p, temperature, max_tokens)
|
||||
is_batch = isinstance(prompt, list)
|
||||
prompts = prompt if is_batch else [prompt]
|
||||
|
||||
@@ -158,7 +151,6 @@ class InferenceEngine:
|
||||
top_p: float = 1.0,
|
||||
top_k: int = 50,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
_validate_sampling_params(top_k, top_p, temperature, max_tokens)
|
||||
sync_gen = self._generate_streaming(
|
||||
[prompt], False, max_tokens, temperature, top_p, top_k
|
||||
)
|
||||
|
||||
@@ -2,6 +2,13 @@ from astrai.model.automodel import AutoModel
|
||||
from astrai.model.components.attention import GQA
|
||||
from astrai.model.components.decoder_block import DecoderBlock
|
||||
from astrai.model.components.linear import Linear
|
||||
from astrai.model.components.lora import (
|
||||
LoRAConfig,
|
||||
inject_lora,
|
||||
load_lora,
|
||||
merge_lora,
|
||||
save_lora,
|
||||
)
|
||||
from astrai.model.components.mlp import MLP
|
||||
from astrai.model.components.norm import RMSNorm
|
||||
from astrai.model.encoder import EmbeddingEncoder
|
||||
@@ -18,4 +25,10 @@ __all__ = [
|
||||
"AutoRegressiveLM",
|
||||
"EmbeddingEncoder",
|
||||
"AutoModel",
|
||||
# LoRA
|
||||
"LoRAConfig",
|
||||
"inject_lora",
|
||||
"merge_lora",
|
||||
"save_lora",
|
||||
"load_lora",
|
||||
]
|
||||
|
||||
+12
-19
@@ -2,16 +2,15 @@
|
||||
AutoModel base class for model loading and saving.
|
||||
"""
|
||||
|
||||
import json
|
||||
from contextlib import contextmanager
|
||||
from pathlib import Path
|
||||
from typing import Self, Union
|
||||
|
||||
import safetensors.torch as st
|
||||
import torch.nn as nn
|
||||
|
||||
from astrai.config.model_config import BaseModelConfig, ConfigFactory
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.serialization import load_model_config, load_model_weights, save_model
|
||||
|
||||
|
||||
@contextmanager
|
||||
@@ -60,25 +59,22 @@ class AutoModel(BaseFactory["AutoModel"], nn.Module):
|
||||
|
||||
model_path = Path(path)
|
||||
|
||||
# Load config
|
||||
config_path = model_path / "config.json"
|
||||
if config_path.exists():
|
||||
with open(config_path, "r") as f:
|
||||
raw = json.load(f)
|
||||
config = ConfigFactory.load(raw)
|
||||
model_type = config.model_type or "autoregressive_lm"
|
||||
else:
|
||||
if not config_path.exists():
|
||||
raise FileNotFoundError(f"Config file not found: {config_path}")
|
||||
|
||||
raw = load_model_config(str(model_path))
|
||||
config = ConfigFactory.load(raw)
|
||||
model_type = config.model_type or "autoregressive_lm"
|
||||
|
||||
actual_cls = AutoModel.get_component_class(model_type)
|
||||
|
||||
with _disable_random_init(enable=disable_random_init):
|
||||
model = actual_cls(config)
|
||||
|
||||
# Load weights
|
||||
weights_path = model_path / "model.safetensors"
|
||||
if weights_path.exists():
|
||||
state_dict = st.load_file(str(weights_path))
|
||||
state_dict = load_model_weights(str(model_path))
|
||||
model.load_state_dict(state_dict, strict=strict)
|
||||
|
||||
return model
|
||||
@@ -87,14 +83,11 @@ class AutoModel(BaseFactory["AutoModel"], nn.Module):
|
||||
self,
|
||||
save_directory: Union[str, Path],
|
||||
) -> None:
|
||||
save_path = Path(save_directory)
|
||||
save_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Save config
|
||||
self.config.to_file(str(save_path / "config.json"))
|
||||
|
||||
# Save weights
|
||||
st.save_file(self.state_dict(), str(save_path / "model.safetensors"))
|
||||
save_model(
|
||||
config=self.config.to_dict(),
|
||||
state_dict=self.state_dict(),
|
||||
save_directory=str(save_directory),
|
||||
)
|
||||
|
||||
def to(self, *args, **kwargs) -> Self:
|
||||
"""Move model to device/dtype."""
|
||||
|
||||
@@ -0,0 +1,194 @@
|
||||
import logging
|
||||
from dataclasses import asdict, dataclass
|
||||
from pathlib import Path
|
||||
from typing import Optional, Set
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from astrai.model.components.linear import Linear
|
||||
from astrai.serialization import (
|
||||
load_json,
|
||||
load_safetensors,
|
||||
save_json,
|
||||
save_safetensors,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
TARGET_MODULES_ATTN = {"q_proj", "k_proj", "v_proj", "o_proj"}
|
||||
TARGET_MODULES_FFN = {"up", "gate", "down"}
|
||||
|
||||
|
||||
@dataclass
|
||||
class LoRAConfig:
|
||||
r: int = 16
|
||||
alpha: int = 32
|
||||
target_modules: tuple = ("q_proj", "v_proj")
|
||||
|
||||
|
||||
class LoRALinear(nn.Module):
|
||||
def __init__(self, base: Linear, r: int = 16, alpha: int = 32):
|
||||
super().__init__()
|
||||
self.register_parameter("weight", base.weight)
|
||||
self.weight.requires_grad_(False)
|
||||
self.bias = base.bias
|
||||
if self.bias is not None:
|
||||
self.bias.requires_grad_(False)
|
||||
|
||||
self.r = r
|
||||
self.scaling = alpha / r
|
||||
self.lora_A = nn.Parameter(torch.randn(r, self.weight.shape[1]) / r)
|
||||
self.lora_B = nn.Parameter(torch.zeros(self.weight.shape[0], r))
|
||||
self._merged = False
|
||||
|
||||
def forward(self, x):
|
||||
out = F.linear(x, self.weight, self.bias)
|
||||
if not self._merged:
|
||||
out += (F.linear(x, self.lora_A) @ self.lora_B.T) * self.scaling
|
||||
return out
|
||||
|
||||
def merge(self):
|
||||
if self._merged:
|
||||
return
|
||||
self.weight.data += (self.lora_B @ self.lora_A) * self.scaling
|
||||
self._merged = True
|
||||
del self.lora_A
|
||||
del self.lora_B
|
||||
|
||||
|
||||
def _collect_lora_info(model: nn.Module) -> dict:
|
||||
names = {}
|
||||
for n, m in model.named_modules():
|
||||
if isinstance(m, Linear):
|
||||
_, _, child = n.rpartition(".")
|
||||
names.setdefault(child, []).append(n)
|
||||
return names
|
||||
|
||||
|
||||
def _get_lora_count(model: nn.Module) -> int:
|
||||
return sum(1 for m in model.modules() if isinstance(m, LoRALinear))
|
||||
|
||||
|
||||
def inject_lora(
|
||||
model: nn.Module,
|
||||
r: int = 16,
|
||||
alpha: int = 32,
|
||||
target_modules: Optional[Set[str]] = None,
|
||||
) -> LoRAConfig:
|
||||
if target_modules is None:
|
||||
target_modules = TARGET_MODULES_ATTN
|
||||
|
||||
available = _collect_lora_info(model)
|
||||
injected = 0
|
||||
|
||||
for name, module in list(model.named_modules()):
|
||||
if not isinstance(module, Linear):
|
||||
continue
|
||||
parent_name, _, child_name = name.rpartition(".")
|
||||
if child_name not in target_modules:
|
||||
continue
|
||||
parent = model.get_submodule(parent_name) if parent_name else model
|
||||
setattr(parent, child_name, LoRALinear(module, r=r, alpha=alpha))
|
||||
injected += 1
|
||||
|
||||
if injected == 0:
|
||||
logger.warning(
|
||||
"No LoRA layers injected. Available Linear child names: %s. "
|
||||
"target_modules: %s. Check model type and target_modules.",
|
||||
sorted(available),
|
||||
sorted(target_modules),
|
||||
)
|
||||
else:
|
||||
logger.info("LoRA injected: %d layers (r=%d, alpha=%d)", injected, r, alpha)
|
||||
|
||||
return LoRAConfig(r=r, alpha=alpha, target_modules=tuple(target_modules))
|
||||
|
||||
|
||||
def merge_lora(model: nn.Module):
|
||||
n = 0
|
||||
for module in model.modules():
|
||||
if isinstance(module, LoRALinear):
|
||||
module.merge()
|
||||
n += 1
|
||||
if n == 0:
|
||||
logger.warning("No LoRA layers to merge.")
|
||||
else:
|
||||
logger.info("Merged %d LoRA layers", n)
|
||||
|
||||
|
||||
def save_lora(model: nn.Module, save_dir: str, config: LoRAConfig):
|
||||
lora_sd = {
|
||||
k: v
|
||||
for k, v in model.state_dict().items()
|
||||
if k.endswith((".lora_A", ".lora_B"))
|
||||
}
|
||||
if not lora_sd:
|
||||
raise RuntimeError(
|
||||
"No LoRA parameters found in model. "
|
||||
"The model may not have been injected or was already merged."
|
||||
)
|
||||
|
||||
path = Path(save_dir)
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
save_safetensors(lora_sd, path / "adapter_model.safetensors")
|
||||
save_json(asdict(config), path / "adapter_config.json")
|
||||
logger.info("LoRA adapter saved to %s (%d keys)", save_dir, len(lora_sd))
|
||||
|
||||
|
||||
def load_lora(model: nn.Module, load_dir: str) -> LoRAConfig:
|
||||
path = Path(load_dir)
|
||||
raw = load_json(path / "adapter_config.json")
|
||||
config = LoRAConfig(
|
||||
r=raw["r"], alpha=raw["alpha"], target_modules=tuple(raw["target_modules"])
|
||||
)
|
||||
|
||||
existing = _get_lora_count(model)
|
||||
if existing > 0:
|
||||
logger.warning(
|
||||
"Model already has %d LoRA layers. Skipping injection, "
|
||||
"loading weights onto existing layers only.",
|
||||
existing,
|
||||
)
|
||||
else:
|
||||
inject_lora(
|
||||
model,
|
||||
r=config.r,
|
||||
alpha=config.alpha,
|
||||
target_modules=set(config.target_modules),
|
||||
)
|
||||
|
||||
weights = load_safetensors(path / "adapter_model.safetensors")
|
||||
try:
|
||||
missing, unexpected = model.load_state_dict(weights, strict=False)
|
||||
except RuntimeError as e:
|
||||
msg = str(e)
|
||||
if "size mismatch" in msg:
|
||||
raise RuntimeError(
|
||||
f"LoRA weight shapes do not match the model. "
|
||||
f"The adapter config (r={config.r}) may not match the injected layers. "
|
||||
f"Original error: {msg}"
|
||||
) from e
|
||||
raise
|
||||
|
||||
injected = _get_lora_count(model)
|
||||
if injected == 0:
|
||||
raise RuntimeError(
|
||||
"No LoRA layers found after loading. "
|
||||
"Inject LoRA before calling load_lora, or check the adapter config."
|
||||
)
|
||||
|
||||
if missing:
|
||||
lora_missing = [k for k in missing if "lora" in k]
|
||||
if lora_missing:
|
||||
raise RuntimeError(
|
||||
f"LoRA weight keys not found in model: {lora_missing}. "
|
||||
f"The adapter config (r={config.r}) may not match the model."
|
||||
)
|
||||
logger.debug("LoRA load: %d missing base-weight keys (expected)", len(missing))
|
||||
if unexpected:
|
||||
logger.warning("LoRA load: %d unexpected keys", len(unexpected))
|
||||
|
||||
logger.info("LoRA adapter loaded from %s", load_dir)
|
||||
return config
|
||||
@@ -1,4 +1,4 @@
|
||||
from typing import Optional
|
||||
from typing import Dict, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
@@ -19,6 +19,10 @@ def get_rotary_emb(
|
||||
return torch.complex(cos, sin)
|
||||
|
||||
|
||||
def ntk_base(base: float, dim: int, factor: float) -> float:
|
||||
return base * (factor ** (dim / (dim - 2)))
|
||||
|
||||
|
||||
def apply_rotary_emb(x: torch.Tensor, freqs_cis: Tensor) -> Tensor:
|
||||
dtype = x.dtype
|
||||
x_ = x.float().reshape(*x.shape[:-1], -1, 2)
|
||||
@@ -30,11 +34,25 @@ def apply_rotary_emb(x: torch.Tensor, freqs_cis: Tensor) -> Tensor:
|
||||
|
||||
|
||||
class RotaryEmbedding(nn.Module):
|
||||
def __init__(self, dim: int, max_len: int, base: float = 10000):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
max_len: int,
|
||||
base: float = 10000,
|
||||
rope_scaling: Optional[Dict] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.max_len = max_len
|
||||
self.base = base
|
||||
self.rope_scaling = rope_scaling
|
||||
|
||||
if rope_scaling is not None:
|
||||
scaling_type = rope_scaling.get("type", "ntk")
|
||||
factor = rope_scaling.get("factor", 1.0)
|
||||
if scaling_type == "ntk":
|
||||
self.base = ntk_base(base, dim, factor)
|
||||
|
||||
self._set_rotary_buffer(self.max_len)
|
||||
|
||||
def _set_rotary_buffer(self, max_len: int):
|
||||
|
||||
@@ -20,7 +20,9 @@ class EmbeddingEncoder(AutoModel):
|
||||
self.config = config
|
||||
rope_dim = config.dim // config.n_heads
|
||||
rope_base = config.rope_theta if config.rope_theta is not None else 10000
|
||||
self.rotary_embedding = RotaryEmbedding(rope_dim, config.max_len, rope_base)
|
||||
self.rotary_embedding = RotaryEmbedding(
|
||||
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling
|
||||
)
|
||||
self.embed_tokens = Embedding(config.vocab_size, config.dim)
|
||||
|
||||
self.layers = nn.ModuleList(
|
||||
|
||||
@@ -59,7 +59,9 @@ class AutoRegressiveLM(AutoModel):
|
||||
else config.dim // config.n_heads
|
||||
)
|
||||
rope_base = config.rope_theta if config.rope_theta is not None else 10000
|
||||
self.rotary_embedding = RotaryEmbedding(rope_dim, config.max_len, rope_base)
|
||||
self.rotary_embedding = RotaryEmbedding(
|
||||
rope_dim, config.max_len, rope_base, rope_scaling=config.rope_scaling
|
||||
)
|
||||
self.embed_tokens = Embedding(config.vocab_size, config.dim)
|
||||
|
||||
self.layers = nn.ModuleList(
|
||||
|
||||
@@ -1,3 +1,13 @@
|
||||
from astrai.parallel.executor import (
|
||||
AccumOptimizer,
|
||||
AccumScheduler,
|
||||
BaseExecutor,
|
||||
DDPExecutor,
|
||||
ExecutorFactory,
|
||||
FSDPExecutor,
|
||||
GradientState,
|
||||
NoneExecutor,
|
||||
)
|
||||
from astrai.parallel.module import ColumnParallelLinear, RowParallelLinear
|
||||
from astrai.parallel.setup import (
|
||||
get_current_device,
|
||||
@@ -17,4 +27,12 @@ __all__ = [
|
||||
"spawn_parallel_fn",
|
||||
"RowParallelLinear",
|
||||
"ColumnParallelLinear",
|
||||
"ExecutorFactory",
|
||||
"BaseExecutor",
|
||||
"GradientState",
|
||||
"AccumOptimizer",
|
||||
"AccumScheduler",
|
||||
"NoneExecutor",
|
||||
"DDPExecutor",
|
||||
"FSDPExecutor",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,231 @@
|
||||
"""Unified training executor — parallel strategy + gradient accumulation."""
|
||||
|
||||
import contextlib
|
||||
import logging
|
||||
from contextlib import contextmanager
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||
from torch.optim import Optimizer
|
||||
from torch.optim.lr_scheduler import LRScheduler
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
from astrai.parallel.setup import get_rank, get_world_size
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class GradientState:
|
||||
def __init__(self, grad_accum_steps: int = 1):
|
||||
self.num_steps = max(grad_accum_steps, 1)
|
||||
self._step: int = 0
|
||||
self._sync_gradients: bool = True
|
||||
|
||||
@property
|
||||
def sync_gradients(self) -> bool:
|
||||
return self._sync_gradients
|
||||
|
||||
def _do_sync(self):
|
||||
self._step += 1
|
||||
self._sync_gradients = self._step % self.num_steps == 0
|
||||
|
||||
|
||||
class AccumOptimizer:
|
||||
def __init__(self, optimizer: Optimizer, gradient_state: GradientState):
|
||||
self.optimizer = optimizer
|
||||
self.gradient_state = gradient_state
|
||||
|
||||
def step(self, closure=None):
|
||||
if self.gradient_state.sync_gradients:
|
||||
self.optimizer.step(closure)
|
||||
|
||||
def zero_grad(self):
|
||||
if self.gradient_state.sync_gradients:
|
||||
self.optimizer.zero_grad()
|
||||
|
||||
@property
|
||||
def param_groups(self):
|
||||
return self.optimizer.param_groups
|
||||
|
||||
def state_dict(self):
|
||||
return self.optimizer.state_dict()
|
||||
|
||||
def load_state_dict(self, d):
|
||||
self.optimizer.load_state_dict(d)
|
||||
|
||||
|
||||
class AccumScheduler:
|
||||
def __init__(self, scheduler: LRScheduler, gradient_state: GradientState):
|
||||
self.scheduler = scheduler
|
||||
self.gradient_state = gradient_state
|
||||
|
||||
def step(self):
|
||||
if self.gradient_state.sync_gradients:
|
||||
self.scheduler.step()
|
||||
|
||||
def state_dict(self):
|
||||
return self.scheduler.state_dict()
|
||||
|
||||
def load_state_dict(self, d):
|
||||
self.scheduler.load_state_dict(d)
|
||||
|
||||
def get_last_lr(self):
|
||||
return self.scheduler.get_last_lr()
|
||||
|
||||
|
||||
class BaseExecutor:
|
||||
def __init__(self, grad_accum_steps: int = 1):
|
||||
self.gradient_state = GradientState(grad_accum_steps)
|
||||
|
||||
def prepare(
|
||||
self,
|
||||
model: nn.Module,
|
||||
optimizer: Optional[Optimizer] = None,
|
||||
dataloader: Optional[DataLoader] = None,
|
||||
scheduler: Optional[LRScheduler] = None,
|
||||
) -> Tuple[
|
||||
nn.Module, Optional[Optimizer], Optional[DataLoader], Optional[LRScheduler]
|
||||
]:
|
||||
model = self._prepare_model(model)
|
||||
if optimizer is not None:
|
||||
optimizer = AccumOptimizer(optimizer, self.gradient_state)
|
||||
if scheduler is not None:
|
||||
scheduler = AccumScheduler(scheduler, self.gradient_state)
|
||||
return model, optimizer, dataloader, scheduler
|
||||
|
||||
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||
return model
|
||||
|
||||
def _no_sync(self, model: nn.Module):
|
||||
return contextlib.nullcontext()
|
||||
|
||||
@contextmanager
|
||||
def accumulate(self, model: nn.Module):
|
||||
self.gradient_state._do_sync()
|
||||
if not self.gradient_state.sync_gradients:
|
||||
with self._no_sync(model):
|
||||
yield
|
||||
else:
|
||||
yield
|
||||
|
||||
def backward(self, loss: torch.Tensor):
|
||||
loss.backward()
|
||||
|
||||
def unwrap_model(self, model: nn.Module) -> nn.Module:
|
||||
return model
|
||||
|
||||
@property
|
||||
def use_distributed(self) -> bool:
|
||||
return get_world_size() > 1
|
||||
|
||||
@property
|
||||
def sync_gradients(self) -> bool:
|
||||
return self.gradient_state.sync_gradients
|
||||
|
||||
@property
|
||||
def grad_accum_steps(self) -> int:
|
||||
return self.gradient_state.num_steps
|
||||
|
||||
|
||||
class ExecutorFactory(BaseFactory[BaseExecutor]):
|
||||
pass
|
||||
|
||||
|
||||
@ExecutorFactory.register("none")
|
||||
class NoneExecutor(BaseExecutor):
|
||||
pass
|
||||
|
||||
|
||||
@ExecutorFactory.register("ddp")
|
||||
class DDPExecutor(BaseExecutor):
|
||||
def __init__(
|
||||
self,
|
||||
grad_accum_steps: int = 1,
|
||||
dim: int = 0,
|
||||
broadcast_buffers: bool = True,
|
||||
init_sync: bool = True,
|
||||
process_group=None,
|
||||
bucket_cap_mb: int = 25,
|
||||
find_unused_parameters: bool = False,
|
||||
check_reduction: bool = False,
|
||||
gradient_as_bucket_view: bool = False,
|
||||
static_graph: bool = False,
|
||||
delay_all_reduce_named_params=None,
|
||||
param_to_hook_all_reduce=None,
|
||||
mixed_precision=None,
|
||||
device_mesh=None,
|
||||
):
|
||||
super().__init__(grad_accum_steps=grad_accum_steps)
|
||||
self._ddp_kwargs = dict(
|
||||
dim=dim,
|
||||
broadcast_buffers=broadcast_buffers,
|
||||
init_sync=init_sync,
|
||||
process_group=process_group,
|
||||
bucket_cap_mb=bucket_cap_mb,
|
||||
find_unused_parameters=find_unused_parameters,
|
||||
check_reduction=check_reduction,
|
||||
gradient_as_bucket_view=gradient_as_bucket_view,
|
||||
static_graph=static_graph,
|
||||
delay_all_reduce_named_params=delay_all_reduce_named_params,
|
||||
param_to_hook_all_reduce=param_to_hook_all_reduce,
|
||||
mixed_precision=mixed_precision,
|
||||
device_mesh=device_mesh,
|
||||
)
|
||||
|
||||
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||
if not self.use_distributed:
|
||||
logger.warning("DDP backend selected but world_size=1, model not wrapped")
|
||||
return model
|
||||
local_rank = get_rank()
|
||||
model = DDP(
|
||||
model,
|
||||
device_ids=[local_rank],
|
||||
output_device=local_rank,
|
||||
**self._ddp_kwargs,
|
||||
)
|
||||
logger.info("Model wrapped with DDP (world_size=%d)", get_world_size())
|
||||
return model
|
||||
|
||||
def _no_sync(self, model: nn.Module):
|
||||
if isinstance(model, DDP):
|
||||
return model.no_sync()
|
||||
return contextlib.nullcontext()
|
||||
|
||||
def unwrap_model(self, model: nn.Module) -> nn.Module:
|
||||
if isinstance(model, DDP):
|
||||
return model.module
|
||||
return model
|
||||
|
||||
|
||||
@ExecutorFactory.register("fsdp")
|
||||
class FSDPExecutor(BaseExecutor):
|
||||
def __init__(self, grad_accum_steps: int = 1, **fsdp_kwargs):
|
||||
super().__init__(grad_accum_steps=grad_accum_steps)
|
||||
self._fsdp_kwargs = fsdp_kwargs
|
||||
self._original_model: Optional[nn.Module] = None
|
||||
|
||||
def _prepare_model(self, model: nn.Module) -> nn.Module:
|
||||
if not self.use_distributed:
|
||||
logger.warning("FSDP backend selected but world_size=1, model not wrapped")
|
||||
return model
|
||||
self._original_model = model
|
||||
device_id = torch.device("cuda", get_rank())
|
||||
model = FSDP(model, device_id=device_id, **self._fsdp_kwargs)
|
||||
logger.info("Model wrapped with FSDP (world_size=%d)", get_world_size())
|
||||
return model
|
||||
|
||||
def _no_sync(self, model: nn.Module):
|
||||
if isinstance(model, FSDP):
|
||||
return model.no_sync()
|
||||
return contextlib.nullcontext()
|
||||
|
||||
def unwrap_model(self, model: nn.Module) -> nn.Module:
|
||||
if self._original_model is not None:
|
||||
return self._original_model
|
||||
if isinstance(model, FSDP):
|
||||
return model._fsdp_wrapped_module
|
||||
return model
|
||||
@@ -0,0 +1,21 @@
|
||||
"""Training component protocols — structural subtyping for optimizer/scheduler wrappers."""
|
||||
|
||||
from typing import Any, Protocol, runtime_checkable
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class OptimizerProtocol(Protocol):
|
||||
def step(self, closure=None): ...
|
||||
def zero_grad(self): ...
|
||||
@property
|
||||
def param_groups(self) -> Any: ...
|
||||
def state_dict(self) -> dict: ...
|
||||
def load_state_dict(self, d: dict): ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class SchedulerProtocol(Protocol):
|
||||
def step(self): ...
|
||||
def state_dict(self) -> dict: ...
|
||||
def load_state_dict(self, d: dict): ...
|
||||
def get_last_lr(self): ...
|
||||
+75
-48
@@ -1,7 +1,8 @@
|
||||
import json
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any, Dict
|
||||
|
||||
import safetensors.torch as st
|
||||
import torch
|
||||
@@ -9,75 +10,101 @@ import torch.distributed as dist
|
||||
|
||||
from astrai.parallel.setup import get_rank
|
||||
|
||||
_META_FILE = "meta.json"
|
||||
_WEIGHTS_FILE = "model.safetensors"
|
||||
_MODEL_CONFIG_FILE = "config.json"
|
||||
|
||||
|
||||
def save_safetensors(state_dict: dict, path: str | Path) -> None:
|
||||
st.save_file(state_dict, str(path))
|
||||
|
||||
|
||||
def load_safetensors(path: str | Path) -> dict:
|
||||
return st.load_file(str(path))
|
||||
|
||||
|
||||
def save_json(data: dict, path: str | Path) -> None:
|
||||
with open(str(path), "w") as f:
|
||||
json.dump(data, f, indent=2)
|
||||
|
||||
|
||||
def load_json(path: str | Path) -> dict:
|
||||
with open(str(path), "r") as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def save_torch(obj: Any, path: str | Path) -> None:
|
||||
torch.save(obj, str(path))
|
||||
|
||||
|
||||
def load_torch(path: str | Path) -> Any:
|
||||
return torch.load(str(path), map_location="cpu", weights_only=False)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Checkpoint:
|
||||
def __init__(
|
||||
self,
|
||||
state_dict: Dict[str, Any],
|
||||
epoch: int = 0,
|
||||
iteration: int = 0,
|
||||
extra: Optional[Dict[str, Any]] = None,
|
||||
meta: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
self.state_dict = state_dict
|
||||
self.epoch = epoch
|
||||
self.iteration = iteration
|
||||
self.extra = extra or {}
|
||||
self.meta = meta or {}
|
||||
|
||||
def save(
|
||||
self,
|
||||
save_dir: str,
|
||||
) -> None:
|
||||
state_dict: Dict[str, Any] = field(default_factory=dict)
|
||||
epoch: int = 0
|
||||
iteration: int = 0
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
meta: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def save(self, save_dir: str) -> None:
|
||||
save_path = Path(save_dir)
|
||||
save_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
rank = get_rank()
|
||||
if rank == 0:
|
||||
meta = {
|
||||
"epoch": self.epoch,
|
||||
"iteration": self.iteration,
|
||||
"timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"),
|
||||
}
|
||||
meta.update(self.meta)
|
||||
with open(save_path / "meta.json", "w") as f:
|
||||
json.dump(meta, f, indent=2)
|
||||
if get_rank() != 0:
|
||||
return
|
||||
|
||||
st.save_file(self.state_dict, save_path / "state_dict.safetensors")
|
||||
if self.extra:
|
||||
for key, value in self.extra.items():
|
||||
torch.save(value, save_path / f"{key}.pt")
|
||||
meta = {
|
||||
"epoch": self.epoch,
|
||||
"iteration": self.iteration,
|
||||
"timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"),
|
||||
**self.meta,
|
||||
}
|
||||
save_json(meta, save_path / _META_FILE)
|
||||
save_safetensors(self.state_dict, save_path / _WEIGHTS_FILE)
|
||||
for key, value in self.extra.items():
|
||||
save_torch(value, save_path / f"{key}.pt")
|
||||
|
||||
@classmethod
|
||||
def load(
|
||||
cls,
|
||||
save_dir: str,
|
||||
) -> "Checkpoint":
|
||||
|
||||
rank = get_rank()
|
||||
def load(cls, save_dir: str) -> "Checkpoint":
|
||||
save_path = Path(save_dir)
|
||||
|
||||
meta = {}
|
||||
if rank == 0:
|
||||
with open(Path(save_dir) / "meta.json", "r") as f:
|
||||
meta = json.load(f)
|
||||
if get_rank() == 0:
|
||||
meta = load_json(save_path / _META_FILE)
|
||||
|
||||
if dist.is_initialized():
|
||||
meta_list = [meta]
|
||||
dist.broadcast_object_list(meta_list, src=0)
|
||||
meta = meta_list[0]
|
||||
|
||||
state_dict = st.load_file(save_path / "state_dict.safetensors")
|
||||
state_dict = load_safetensors(save_path / _WEIGHTS_FILE)
|
||||
|
||||
extra = {}
|
||||
for f in save_path.iterdir():
|
||||
if f.suffix == ".pt" and f.stem not in ("meta",):
|
||||
extra[f.stem] = torch.load(f, map_location="cpu", weights_only=False)
|
||||
if f.suffix == ".pt":
|
||||
extra[f.stem] = load_torch(f)
|
||||
|
||||
return cls(
|
||||
state_dict=state_dict,
|
||||
epoch=meta["epoch"],
|
||||
iteration=meta["iteration"],
|
||||
extra=extra or None,
|
||||
epoch=meta.get("epoch", 0),
|
||||
iteration=meta.get("iteration", 0),
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
|
||||
def save_model(config: dict, state_dict: dict, save_directory: str) -> None:
|
||||
save_path = Path(save_directory)
|
||||
save_path.mkdir(parents=True, exist_ok=True)
|
||||
save_json(config, save_path / _MODEL_CONFIG_FILE)
|
||||
save_safetensors(state_dict, save_path / _WEIGHTS_FILE)
|
||||
|
||||
|
||||
def load_model_config(save_directory: str) -> dict:
|
||||
return load_json(Path(save_directory) / _MODEL_CONFIG_FILE)
|
||||
|
||||
|
||||
def load_model_weights(save_directory: str) -> dict:
|
||||
return load_safetensors(Path(save_directory) / _WEIGHTS_FILE)
|
||||
|
||||
+64
-34
@@ -4,17 +4,17 @@ from torch.optim import Optimizer
|
||||
|
||||
def _zeropower_via_newtonschulz(G: torch.Tensor, steps: int = 5):
|
||||
assert G.ndim == 2
|
||||
X = G.bfloat16()
|
||||
X = G
|
||||
scale = max(1, G.size(0) / G.size(1)) ** 0.5
|
||||
X = X / (X.norm() + 1e-7) * scale
|
||||
if steps == 0:
|
||||
return X.type_as(G)
|
||||
return X
|
||||
a, b, c = (3.4445, -4.7750, 2.0315)
|
||||
for _ in range(steps):
|
||||
A = X @ X.T
|
||||
B = A @ X
|
||||
X = a * X + b * B + c * (A @ B)
|
||||
return X.type_as(G)
|
||||
return X
|
||||
|
||||
|
||||
class Muon(Optimizer):
|
||||
@@ -50,64 +50,94 @@ class Muon(Optimizer):
|
||||
if closure is not None:
|
||||
with torch.enable_grad():
|
||||
loss = closure()
|
||||
|
||||
for group in self.param_groups:
|
||||
params_2d, params_1d = [], []
|
||||
grads_2d, grads_1d = [], []
|
||||
|
||||
for p in group["params"]:
|
||||
if p.grad is None:
|
||||
continue
|
||||
grad = p.grad
|
||||
if grad.is_sparse:
|
||||
if p.grad.is_sparse:
|
||||
raise RuntimeError("Muon does not support sparse gradients")
|
||||
if p.ndim >= 2:
|
||||
self._muon_update(p, grad, group)
|
||||
params_2d.append(p)
|
||||
grads_2d.append(p.grad)
|
||||
else:
|
||||
self._adamw_update(p, grad, group)
|
||||
params_1d.append(p)
|
||||
grads_1d.append(p.grad)
|
||||
|
||||
if params_2d:
|
||||
self._muon_update_foreach(params_2d, grads_2d, group)
|
||||
if params_1d:
|
||||
self._adamw_update_foreach(params_1d, grads_1d, group)
|
||||
|
||||
return loss
|
||||
|
||||
def _muon_update(self, p, grad, group):
|
||||
def _muon_update_foreach(self, params_2d, grads_2d, group):
|
||||
lr = group["lr"]
|
||||
momentum = group["momentum"]
|
||||
wd = group["weight_decay"]
|
||||
nesterov = group["nesterov"]
|
||||
ns_steps = group["ns_steps"]
|
||||
state = self.state[p]
|
||||
|
||||
p.mul_(1 - lr * wd)
|
||||
if wd != 0:
|
||||
torch._foreach_mul_(params_2d, 1 - lr * wd)
|
||||
|
||||
if nesterov:
|
||||
grad = grad.add(p, alpha=wd)
|
||||
grads_2d = torch._foreach_add(grads_2d, params_2d, alpha=wd)
|
||||
|
||||
if "momentum_buffer" not in state:
|
||||
state["momentum_buffer"] = torch.zeros_like(grad)
|
||||
buf = state["momentum_buffer"]
|
||||
buf.lerp_(grad, 1 - momentum)
|
||||
bufs = []
|
||||
for p, grad in zip(params_2d, grads_2d):
|
||||
state = self.state[p]
|
||||
if "momentum_buffer" not in state:
|
||||
state["momentum_buffer"] = torch.zeros_like(grad)
|
||||
bufs.append(state["momentum_buffer"])
|
||||
|
||||
update = _zeropower_via_newtonschulz(buf, steps=ns_steps)
|
||||
scale = max(1, p.size(0) / p.size(1)) ** 0.5
|
||||
p.add_(update, alpha=-lr * scale)
|
||||
torch._foreach_lerp_(bufs, grads_2d, 1 - momentum)
|
||||
|
||||
def _adamw_update(self, p, grad, group):
|
||||
for p, buf in zip(params_2d, bufs):
|
||||
update = _zeropower_via_newtonschulz(buf, steps=ns_steps)
|
||||
scale = max(1, p.size(0) / p.size(1)) ** 0.5
|
||||
p.add_(update, alpha=-lr * scale)
|
||||
|
||||
def _adamw_update_foreach(self, params_1d, grads_1d, group):
|
||||
lr = group["adamw_lr"]
|
||||
betas = group["adamw_betas"]
|
||||
eps = group["adamw_eps"]
|
||||
wd = group["adamw_wd"]
|
||||
state = self.state[p]
|
||||
|
||||
if not state:
|
||||
state["step"] = 0
|
||||
state["exp_avg"] = torch.zeros_like(p)
|
||||
state["exp_avg_sq"] = torch.zeros_like(p)
|
||||
steps: list[int] = []
|
||||
exp_avgs, exp_avg_sqs = [], []
|
||||
has_state = []
|
||||
for p in params_1d:
|
||||
state = self.state[p]
|
||||
if not state:
|
||||
state["step"] = 0
|
||||
state["exp_avg"] = torch.zeros_like(p)
|
||||
state["exp_avg_sq"] = torch.zeros_like(p)
|
||||
has_state.append(False)
|
||||
else:
|
||||
has_state.append(True)
|
||||
state["step"] += 1
|
||||
steps.append(state["step"])
|
||||
exp_avgs.append(state["exp_avg"])
|
||||
exp_avg_sqs.append(state["exp_avg_sq"])
|
||||
|
||||
state["step"] += 1
|
||||
exp_avg, exp_avg_sq = state["exp_avg"], state["exp_avg_sq"]
|
||||
beta1, beta2 = betas
|
||||
|
||||
exp_avg.lerp_(grad, 1 - beta1)
|
||||
exp_avg_sq.lerp_(grad.square(), 1 - beta2)
|
||||
torch._foreach_lerp_(exp_avgs, grads_1d, 1 - beta1)
|
||||
grads_sq = torch._foreach_mul(grads_1d, grads_1d)
|
||||
torch._foreach_lerp_(exp_avg_sqs, grads_sq, 1 - beta2)
|
||||
|
||||
step = state["step"]
|
||||
bias1 = 1 - beta1**step
|
||||
bias2 = 1 - beta2**step
|
||||
bias_correction1 = [1 - beta1**s for s in steps]
|
||||
bias_correction2 = [1 - beta2**s for s in steps]
|
||||
|
||||
p.mul_(1 - lr * wd)
|
||||
denom = exp_avg_sq.sqrt().div_(bias2**0.5).add_(eps)
|
||||
p.addcdiv_(exp_avg / bias1, denom, value=-lr)
|
||||
if wd != 0:
|
||||
torch._foreach_mul_(params_1d, 1 - lr * wd)
|
||||
|
||||
exp_avg_corrected = torch._foreach_div(exp_avgs, bias_correction1)
|
||||
denom = torch._foreach_div(exp_avg_sqs, bias_correction2)
|
||||
denom = torch._foreach_sqrt(denom)
|
||||
torch._foreach_add_(denom, eps)
|
||||
torch._foreach_addcdiv_(params_1d, exp_avg_corrected, denom, value=-lr)
|
||||
|
||||
@@ -8,15 +8,17 @@ import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||
|
||||
from astrai.factory import BaseFactory
|
||||
|
||||
|
||||
def unwrap_model(model: nn.Module) -> nn.Module:
|
||||
"""Unwrap DDP wrapper if present to get the original model."""
|
||||
if isinstance(model, DDP):
|
||||
return model.module
|
||||
if isinstance(model, FSDP):
|
||||
return model._fsdp_wrapped_module
|
||||
return model
|
||||
|
||||
|
||||
|
||||
@@ -51,18 +51,15 @@ class TrainCallback(Protocol):
|
||||
def on_epoch_end(self, context: TrainContext):
|
||||
"""Called at the end of each epoch."""
|
||||
|
||||
def on_step_begin(self, context: TrainContext):
|
||||
"""Called at the beginning of each step."""
|
||||
|
||||
def on_step_end(self, context: TrainContext):
|
||||
"""Called at the end of each step."""
|
||||
|
||||
def on_batch_begin(self, context: TrainContext):
|
||||
"""Called at the beginning of each batch."""
|
||||
|
||||
def on_batch_end(self, context: TrainContext):
|
||||
"""Called at the end of each batch."""
|
||||
|
||||
def on_optimizer_step(self, context: TrainContext):
|
||||
"""Called on every optimizer step (sync step only)."""
|
||||
|
||||
def on_error(self, context: TrainContext):
|
||||
"""Called when an error occurs during training."""
|
||||
|
||||
@@ -88,7 +85,7 @@ class GradientClippingCallback(TrainCallback):
|
||||
def __init__(self, max_grad_norm: float):
|
||||
self.max_grad_norm = max_grad_norm
|
||||
|
||||
def on_step_begin(self, context: TrainContext):
|
||||
def on_optimizer_step(self, context: TrainContext):
|
||||
clip_grad_norm_(context.model.parameters(), self.max_grad_norm)
|
||||
|
||||
|
||||
@@ -213,7 +210,7 @@ class ProgressBarCallback(TrainCallback):
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, num_epoch: int, log_interval: int = 100, file: IO[str] = sys.stdout
|
||||
self, num_epoch: int, log_interval: int = 100, file: Optional[IO[str]] = None
|
||||
):
|
||||
self.num_epoch = num_epoch
|
||||
self.log_interval = log_interval
|
||||
@@ -226,7 +223,7 @@ class ProgressBarCallback(TrainCallback):
|
||||
context.dataloader,
|
||||
desc=f"Epoch {context.epoch + 1}/{self.num_epoch}",
|
||||
dynamic_ncols=True,
|
||||
file=self.file,
|
||||
file=self.file or sys.stdout,
|
||||
)
|
||||
|
||||
@only_on_rank(0)
|
||||
@@ -344,7 +341,7 @@ class ValidationCallback(TrainCallback):
|
||||
f"Epoch {context.epoch + 1}, Step {step_count}, Val Loss: {avg_loss:.4f}"
|
||||
)
|
||||
|
||||
def on_step_end(self, context: TrainContext):
|
||||
def on_optimizer_step(self, context: TrainContext):
|
||||
if context.val_dataloader is None:
|
||||
return
|
||||
cfg = context.config
|
||||
|
||||
@@ -2,13 +2,14 @@ from dataclasses import dataclass, field
|
||||
from typing import Optional, Self
|
||||
|
||||
import torch.nn as nn
|
||||
from torch.optim import Optimizer
|
||||
from torch.optim.lr_scheduler import LRScheduler
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
from astrai.config.train_config import TrainConfig
|
||||
from astrai.dataset import ResumableDistributedSampler
|
||||
from astrai.model.components.lora import inject_lora
|
||||
from astrai.parallel.executor import BaseExecutor, ExecutorFactory
|
||||
from astrai.parallel.setup import get_current_device, get_rank, get_world_size
|
||||
from astrai.protocols import OptimizerProtocol, SchedulerProtocol
|
||||
from astrai.serialization import Checkpoint
|
||||
from astrai.trainer.strategy import BaseStrategy, StrategyFactory
|
||||
|
||||
@@ -18,10 +19,11 @@ class TrainContext:
|
||||
model: nn.Module = field(default=None)
|
||||
strategy: BaseStrategy = field(default=None)
|
||||
dataloader: DataLoader = field(default=None)
|
||||
optimizer: Optimizer = field(default=None)
|
||||
scheduler: LRScheduler = field(default=None)
|
||||
optimizer: OptimizerProtocol = field(default=None)
|
||||
scheduler: SchedulerProtocol = field(default=None)
|
||||
checkpoint: Checkpoint = field(default=None)
|
||||
config: TrainConfig = field(default=None)
|
||||
executor: BaseExecutor = field(default=None)
|
||||
|
||||
epoch: int = field(default=0)
|
||||
iteration: int = field(default=0)
|
||||
@@ -47,33 +49,47 @@ class TrainContextBuilder:
|
||||
return self
|
||||
|
||||
def build(self) -> TrainContext:
|
||||
cfg = self.config
|
||||
device = get_current_device()
|
||||
|
||||
executor = ExecutorFactory.create(
|
||||
cfg.parallel_mode,
|
||||
grad_accum_steps=cfg.grad_accum_steps,
|
||||
**cfg.executor_kwargs,
|
||||
)
|
||||
|
||||
context = TrainContext(
|
||||
model=self.config.model,
|
||||
model=cfg.model,
|
||||
world_size=get_world_size(),
|
||||
rank=get_rank(),
|
||||
config=self.config,
|
||||
config=cfg,
|
||||
executor=executor,
|
||||
)
|
||||
|
||||
device = get_current_device()
|
||||
context.model = context.model.to(device=device)
|
||||
|
||||
if self.config.nprocs > 1 and self.config.parallel_wrapper:
|
||||
context.model = self.config.parallel_wrapper(context.model)
|
||||
|
||||
if self._checkpoint is not None:
|
||||
context.epoch = max(self._checkpoint.epoch, self.config.start_epoch)
|
||||
context.iteration = max(self._checkpoint.iteration, self.config.start_batch)
|
||||
context.model.load_state_dict(self._checkpoint.state_dict)
|
||||
context.epoch = max(self._checkpoint.epoch, cfg.start_epoch)
|
||||
context.iteration = max(self._checkpoint.iteration, cfg.start_batch)
|
||||
if self._checkpoint.state_dict:
|
||||
context.model.load_state_dict(self._checkpoint.state_dict)
|
||||
context.checkpoint = self._checkpoint
|
||||
else:
|
||||
context.checkpoint = Checkpoint(
|
||||
state_dict=context.model.state_dict(),
|
||||
)
|
||||
|
||||
context.optimizer = self.config.optimizer_fn(context.model)
|
||||
context.scheduler = self.config.scheduler_fn(context.optimizer)
|
||||
if cfg.lora is not None:
|
||||
inject_lora(
|
||||
context.model,
|
||||
r=cfg.lora.r,
|
||||
alpha=cfg.lora.alpha,
|
||||
target_modules=set(cfg.lora.target_modules),
|
||||
)
|
||||
|
||||
context.optimizer = cfg.optimizer_fn(context.model)
|
||||
context.scheduler = cfg.scheduler_fn(context.optimizer)
|
||||
|
||||
cfg = self.config
|
||||
sampler_offset = context.iteration * cfg.batch_per_device
|
||||
sampler = ResumableDistributedSampler(
|
||||
data_source=cfg.dataset,
|
||||
@@ -107,11 +123,20 @@ class TrainContextBuilder:
|
||||
prefetch_factor=cfg.prefetch_factor,
|
||||
)
|
||||
|
||||
context.model, context.optimizer, context.dataloader, context.scheduler = (
|
||||
executor.prepare(
|
||||
context.model,
|
||||
context.optimizer,
|
||||
context.dataloader,
|
||||
context.scheduler,
|
||||
)
|
||||
)
|
||||
|
||||
context.strategy = StrategyFactory.create(
|
||||
model=context.model,
|
||||
train_type=self.config.strategy,
|
||||
train_type=cfg.strategy,
|
||||
device=device,
|
||||
**self.config.extra_kwargs,
|
||||
**cfg.extra_kwargs,
|
||||
)
|
||||
|
||||
return context
|
||||
|
||||
+17
-16
@@ -34,7 +34,6 @@ class Trainer:
|
||||
"checkpoint",
|
||||
cfg.ckpt_dir,
|
||||
cfg.ckpt_interval,
|
||||
state_dict_fn=cfg.state_dict_fn,
|
||||
),
|
||||
CallbackFactory.create(
|
||||
"metric_logger",
|
||||
@@ -56,32 +55,34 @@ class Trainer:
|
||||
method(context)
|
||||
|
||||
def _trainer_loop(self, checkpoint: Optional[Checkpoint] = None):
|
||||
cfg = self.train_config
|
||||
context = TrainContextBuilder(cfg).with_checkpoint(checkpoint).build()
|
||||
context = (
|
||||
TrainContextBuilder(self.train_config).with_checkpoint(checkpoint).build()
|
||||
)
|
||||
executor = context.executor
|
||||
self._call_callbacks("on_train_begin", context)
|
||||
|
||||
try:
|
||||
context.model.train()
|
||||
grad_accum_steps = cfg.grad_accum_steps
|
||||
|
||||
for epoch in range(context.epoch, cfg.n_epoch):
|
||||
for epoch in range(context.epoch, context.config.n_epoch):
|
||||
context.epoch = epoch
|
||||
self._call_callbacks("on_epoch_begin", context)
|
||||
|
||||
for batch in context.dataloader:
|
||||
self._call_callbacks("on_batch_begin", context)
|
||||
loss = context.strategy(batch)
|
||||
context.loss = loss.item()
|
||||
stand_loss = loss / grad_accum_steps
|
||||
stand_loss.backward()
|
||||
context.iteration += 1
|
||||
self._call_callbacks("on_batch_end", context)
|
||||
|
||||
if context.iteration % grad_accum_steps == 0:
|
||||
self._call_callbacks("on_step_begin", context)
|
||||
context.optimizer.step()
|
||||
context.optimizer.zero_grad()
|
||||
self._call_callbacks("on_step_end", context)
|
||||
with executor.accumulate(context.model):
|
||||
loss = context.strategy(batch)
|
||||
context.loss = loss.item()
|
||||
stand_loss = loss / executor.grad_accum_steps
|
||||
executor.backward(stand_loss)
|
||||
context.iteration += 1
|
||||
self._call_callbacks("on_batch_end", context)
|
||||
|
||||
if executor.sync_gradients:
|
||||
self._call_callbacks("on_optimizer_step", context)
|
||||
context.optimizer.step()
|
||||
context.optimizer.zero_grad()
|
||||
|
||||
if context.scheduler:
|
||||
context.scheduler.step()
|
||||
|
||||
+25
-37
@@ -2,16 +2,13 @@ import argparse
|
||||
import os
|
||||
from functools import partial
|
||||
|
||||
import safetensors.torch as st
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.optim as optim
|
||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||
|
||||
from astrai.config import AutoRegressiveLMConfig, TrainConfig
|
||||
from astrai.dataset import DatasetFactory
|
||||
from astrai.model import AutoRegressiveLM
|
||||
from astrai.parallel import get_rank
|
||||
from astrai.serialization import Checkpoint
|
||||
from astrai.trainer import SchedulerFactory, Trainer
|
||||
|
||||
|
||||
@@ -146,6 +143,13 @@ def parse_args() -> argparse.Namespace:
|
||||
)
|
||||
|
||||
parser.add_argument("--nprocs", type=int, default=1, help="Number of GPUs to use.")
|
||||
parser.add_argument(
|
||||
"--parallel_mode",
|
||||
type=str,
|
||||
default="none",
|
||||
choices=["none", "ddp"],
|
||||
help="Parallel training strategy.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--device_type", type=str, default="cuda", help="Device type to use."
|
||||
)
|
||||
@@ -162,21 +166,7 @@ def parse_args() -> argparse.Namespace:
|
||||
return args
|
||||
|
||||
|
||||
def ddp_wrap(model: nn.Module):
|
||||
local_rank = get_rank()
|
||||
ddp_model = DDP(
|
||||
model,
|
||||
device_ids=[local_rank],
|
||||
output_device=local_rank,
|
||||
static_graph=True,
|
||||
find_unused_parameters=False,
|
||||
gradient_as_bucket_view=True,
|
||||
broadcast_buffers=False,
|
||||
)
|
||||
return ddp_model
|
||||
|
||||
|
||||
def create_optimizer(model: nn.Module, **kwargs) -> optim.Optimizer:
|
||||
def create_optimizer(model, **kwargs) -> optim.Optimizer:
|
||||
return optim.AdamW(model.parameters(), fused=True, **kwargs)
|
||||
|
||||
|
||||
@@ -186,12 +176,6 @@ def create_scheduler(
|
||||
return SchedulerFactory.create(optimizer, **kwargs)
|
||||
|
||||
|
||||
def prepare_checkpoint(model: nn.Module) -> dict:
|
||||
if isinstance(model, DDP):
|
||||
return model.module.state_dict()
|
||||
return model.state_dict()
|
||||
|
||||
|
||||
def compute_total_steps(
|
||||
dataset_len: int,
|
||||
n_epoch: int,
|
||||
@@ -238,6 +222,7 @@ def train(
|
||||
window_size: int,
|
||||
stride: int,
|
||||
nprocs: int,
|
||||
parallel_mode: str,
|
||||
device_type: str,
|
||||
start_method: str,
|
||||
):
|
||||
@@ -251,16 +236,14 @@ def train(
|
||||
if window_size is None:
|
||||
window_size = config.max_len
|
||||
|
||||
# Create bare AutoRegressiveLM (for training, no tokenizer needed)
|
||||
model = AutoRegressiveLM(config)
|
||||
# Create model and load full checkpoint (state_dict + optimizer + scheduler + meta)
|
||||
checkpoint = Checkpoint.load(param_path)
|
||||
model = AutoRegressiveLM(config).to(dtype=torch.bfloat16)
|
||||
model.load_state_dict(checkpoint.state_dict, strict=False)
|
||||
|
||||
# Load weights if available
|
||||
weights_path = os.path.join(param_path, "model.safetensors")
|
||||
if os.path.exists(weights_path):
|
||||
state_dict = st.load_file(weights_path)
|
||||
model.load_state_dict(state_dict, strict=False)
|
||||
|
||||
model = model.to(dtype=torch.bfloat16)
|
||||
# Strip state_dict to avoid pickling ~7GB through mp.spawn pipe
|
||||
# (model weights already loaded into model above)
|
||||
checkpoint.state_dict = {}
|
||||
|
||||
strategy_kwargs = {
|
||||
"beta": dpo_beta,
|
||||
@@ -271,6 +254,11 @@ def train(
|
||||
"sync_interval": grpo_sync_interval,
|
||||
}
|
||||
|
||||
executor_kwargs = {
|
||||
"gradient_as_bucket_view": True,
|
||||
"broadcast_buffers": False,
|
||||
}
|
||||
|
||||
dataset = DatasetFactory.load(
|
||||
train_type=train_type,
|
||||
load_path=data_root_path,
|
||||
@@ -319,15 +307,15 @@ def train(
|
||||
num_workers=num_workers,
|
||||
pin_memory=pin_memory,
|
||||
nprocs=nprocs,
|
||||
parallel_wrapper=ddp_wrap,
|
||||
state_dict_fn=prepare_checkpoint,
|
||||
parallel_mode=parallel_mode,
|
||||
device_type=device_type,
|
||||
start_method=start_method,
|
||||
executor_kwargs=executor_kwargs,
|
||||
extra_kwargs=strategy_kwargs,
|
||||
)
|
||||
|
||||
trainer = Trainer(train_config)
|
||||
trainer.train()
|
||||
trainer.train(checkpoint=checkpoint)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
import torch
|
||||
@@ -36,7 +37,6 @@ def test_single_process():
|
||||
|
||||
|
||||
def test_checkpoint_with_extra():
|
||||
"""Verify extra keys are saved as individual .pt files and loaded back."""
|
||||
model = torch.nn.Linear(10, 5)
|
||||
optimizer = AdamW(model.parameters(), lr=1e-3)
|
||||
optimizer.step()
|
||||
@@ -52,8 +52,6 @@ def test_checkpoint_with_extra():
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
checkpoint.save(tmpdir)
|
||||
|
||||
import os
|
||||
|
||||
assert os.path.exists(os.path.join(tmpdir, "optimizer.pt"))
|
||||
assert os.path.exists(os.path.join(tmpdir, "scheduler.pt"))
|
||||
|
||||
|
||||
@@ -0,0 +1,286 @@
|
||||
"""Unit tests for protocol builders, StopChecker, GenContext, StopInfo."""
|
||||
|
||||
import json
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from astrai.inference.api.anthropic import AnthropicResponseBuilder
|
||||
from astrai.inference.api.openai import OpenAIResponseBuilder
|
||||
from astrai.inference.api.protocol import GenContext, StopChecker, StopInfo
|
||||
from astrai.inference.engine import GenerationRequest
|
||||
|
||||
|
||||
def _make_ctx(**kwargs):
|
||||
defaults = {
|
||||
"resp_id": "test-123",
|
||||
"created": 1000,
|
||||
"model": "test-model",
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 5,
|
||||
}
|
||||
defaults.update(kwargs)
|
||||
return GenContext(**defaults)
|
||||
|
||||
|
||||
def _sse_payloads(events):
|
||||
payloads = []
|
||||
for chunk in events:
|
||||
for line in chunk.strip().split("\n"):
|
||||
if line.startswith("data: "):
|
||||
try:
|
||||
payloads.append(json.loads(line[6:]))
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
return payloads
|
||||
|
||||
|
||||
class TestStopChecker:
|
||||
def test_check_finds_match(self):
|
||||
sc = StopChecker(["stop", "end"])
|
||||
assert sc.check("hello stop world") == "stop"
|
||||
|
||||
def test_check_returns_none_when_no_match(self):
|
||||
sc = StopChecker(["stop"])
|
||||
assert sc.check("hello world") is None
|
||||
|
||||
def test_check_empty_sequences(self):
|
||||
sc = StopChecker([])
|
||||
assert sc.check("hello") is None
|
||||
|
||||
|
||||
class TestGenContext:
|
||||
def test_defaults(self):
|
||||
ctx = GenContext(resp_id="a", created=1, model="m", prompt_tokens=10)
|
||||
assert ctx.completion_tokens == 0
|
||||
|
||||
def test_fields_mutable(self):
|
||||
ctx = GenContext(resp_id="a", created=1, model="m", prompt_tokens=10)
|
||||
ctx.completion_tokens = 42
|
||||
assert ctx.completion_tokens == 42
|
||||
|
||||
|
||||
class TestStopInfo:
|
||||
def test_defaults(self):
|
||||
s = StopInfo()
|
||||
assert s.matched is None
|
||||
assert s.body == ""
|
||||
assert s.yielded == ""
|
||||
|
||||
def test_with_values(self):
|
||||
s = StopInfo(matched="stop", body="hello stop", yielded="hello ")
|
||||
assert s.matched == "stop"
|
||||
assert s.body == "hello stop"
|
||||
assert s.yielded == "hello "
|
||||
|
||||
|
||||
class TestOpenAIResponseBuilder:
|
||||
@pytest.fixture
|
||||
def builder(self):
|
||||
builder = OpenAIResponseBuilder()
|
||||
req = MagicMock()
|
||||
req.messages = [MagicMock(role="user", content="Hello")]
|
||||
req.stop = None
|
||||
req.model = "astrai"
|
||||
engine = MagicMock()
|
||||
engine.tokenizer.apply_chat_template.return_value = "Hello"
|
||||
builder.prepare(req, engine)
|
||||
return builder
|
||||
|
||||
def test_prepare_returns_prompt_ctx_stops(self, builder):
|
||||
req = MagicMock()
|
||||
req.messages = [MagicMock(role="user", content="Hi")]
|
||||
req.stop = ["END"]
|
||||
req.model = "gpt"
|
||||
engine = MagicMock()
|
||||
engine.tokenizer.apply_chat_template.return_value = "Hi"
|
||||
prompt, ctx, stops = builder.prepare(req, engine)
|
||||
assert prompt == "Hi"
|
||||
assert ctx.model == "gpt"
|
||||
assert ctx.prompt_tokens == 0
|
||||
assert stops == ["END"]
|
||||
|
||||
def test_prepare_no_stop_returns_empty_list(self, builder):
|
||||
req = MagicMock()
|
||||
req.messages = []
|
||||
req.stop = None
|
||||
req.model = "x"
|
||||
engine = MagicMock()
|
||||
engine.tokenizer.apply_chat_template.return_value = ""
|
||||
_, _, stops = builder.prepare(req, engine)
|
||||
assert stops == []
|
||||
|
||||
def test_format_stream_start(self, builder):
|
||||
ctx = _make_ctx()
|
||||
events = builder.format_stream_start(ctx)
|
||||
payloads = _sse_payloads(events)
|
||||
assert len(payloads) == 1
|
||||
p = payloads[0]
|
||||
assert p["object"] == "chat.completion.chunk"
|
||||
assert p["choices"][0]["delta"]["role"] == "assistant"
|
||||
assert p["choices"][0]["finish_reason"] is None
|
||||
|
||||
def test_format_chunk(self, builder):
|
||||
event = builder.format_chunk("hello")
|
||||
payload = json.loads(event.split("data: ", 1)[1])
|
||||
assert payload["choices"][0]["delta"]["content"] == "hello"
|
||||
assert payload["choices"][0]["finish_reason"] is None
|
||||
|
||||
def test_format_stream_end(self, builder):
|
||||
ctx = _make_ctx(completion_tokens=5)
|
||||
stop = StopInfo(matched="stop")
|
||||
events = builder.format_stream_end(ctx, stop)
|
||||
payloads = _sse_payloads(events)
|
||||
finish = payloads[0]
|
||||
assert finish["choices"][0]["finish_reason"] == "stop"
|
||||
usage = payloads[1]
|
||||
assert usage["completion_tokens"] == 5
|
||||
assert usage["total_tokens"] == 15
|
||||
|
||||
def test_format_response(self, builder):
|
||||
ctx = _make_ctx()
|
||||
stop = StopInfo()
|
||||
resp = builder.format_response(ctx, "hello", stop)
|
||||
assert resp["object"] == "chat.completion"
|
||||
assert resp["choices"][0]["message"]["content"] == "hello"
|
||||
assert resp["usage"]["prompt_tokens"] == 10
|
||||
|
||||
|
||||
class TestAnthropicResponseBuilder:
|
||||
@pytest.fixture
|
||||
def builder(self):
|
||||
builder = AnthropicResponseBuilder()
|
||||
req = MagicMock()
|
||||
req.messages = [MagicMock(role="user", content="Hello")]
|
||||
req.model = "claude"
|
||||
engine = MagicMock()
|
||||
engine.tokenizer.apply_chat_template.return_value = "Hello"
|
||||
req.system = None
|
||||
builder.prepare(req, engine)
|
||||
return builder
|
||||
|
||||
def test_prepare_messages(self, builder):
|
||||
req = MagicMock()
|
||||
req.messages = [MagicMock(role="user", content="Hi")]
|
||||
req.model = "claude"
|
||||
req.system = None
|
||||
req.stop_sequences = None
|
||||
engine = MagicMock()
|
||||
engine.tokenizer.apply_chat_template.return_value = "Hi"
|
||||
prompt, ctx, stops = builder.prepare(req, engine)
|
||||
assert prompt == "Hi"
|
||||
assert stops == []
|
||||
|
||||
def test_prepare_with_stop_sequences(self, builder):
|
||||
req = MagicMock()
|
||||
req.messages = []
|
||||
req.model = "x"
|
||||
req.stop_sequences = ["stop", "end"]
|
||||
req.system = None
|
||||
engine = MagicMock()
|
||||
engine.tokenizer.apply_chat_template.return_value = ""
|
||||
_, _, stops = builder.prepare(req, engine)
|
||||
assert stops == ["stop", "end"]
|
||||
|
||||
def test_format_stream_start(self, builder):
|
||||
ctx = _make_ctx(prompt_tokens=3)
|
||||
events = builder.format_stream_start(ctx)
|
||||
payloads = _sse_payloads(events)
|
||||
assert len(payloads) == 2
|
||||
assert payloads[0]["type"] == "message_start"
|
||||
assert payloads[0]["message"]["usage"]["input_tokens"] == 3
|
||||
assert payloads[1]["type"] == "content_block_start"
|
||||
|
||||
def test_format_chunk(self, builder):
|
||||
event = builder.format_chunk("tok")
|
||||
payload = json.loads(event.split("data: ", 1)[1])
|
||||
assert payload["type"] == "content_block_delta"
|
||||
assert payload["delta"]["text"] == "tok"
|
||||
|
||||
def test_format_stream_end_no_stop(self, builder):
|
||||
ctx = _make_ctx(completion_tokens=3)
|
||||
stop = StopInfo()
|
||||
events = builder.format_stream_end(ctx, stop)
|
||||
payloads = _sse_payloads(events)
|
||||
# content_block_stop, message_delta, message_stop
|
||||
types = [p["type"] for p in payloads]
|
||||
assert types == ["content_block_stop", "message_delta", "message_stop"]
|
||||
assert payloads[1]["delta"]["stop_reason"] == "end_turn"
|
||||
|
||||
def test_format_stream_end_with_stop_trims_and_emits_remaining(self, builder):
|
||||
ctx = _make_ctx(completion_tokens=7)
|
||||
stop = StopInfo(
|
||||
matched="END",
|
||||
body="Hello world END extra",
|
||||
yielded="Hello ",
|
||||
)
|
||||
events = builder.format_stream_end(ctx, stop)
|
||||
payloads = _sse_payloads(events)
|
||||
# unyielded delta, content_block_stop, message_delta, message_stop
|
||||
types = [p["type"] for p in payloads]
|
||||
assert types == [
|
||||
"content_block_delta",
|
||||
"content_block_stop",
|
||||
"message_delta",
|
||||
"message_stop",
|
||||
]
|
||||
assert payloads[0]["delta"]["text"] == "world "
|
||||
assert payloads[2]["delta"]["stop_reason"] == "stop_sequence"
|
||||
assert payloads[2]["delta"]["stop_sequence"] == "END"
|
||||
|
||||
def test_format_stream_end_stop_trimmed_already_yielded(self, builder):
|
||||
ctx = _make_ctx()
|
||||
stop = StopInfo(
|
||||
matched="END",
|
||||
body="Hello END",
|
||||
yielded="Hello ",
|
||||
)
|
||||
events = builder.format_stream_end(ctx, stop)
|
||||
payloads = _sse_payloads(events)
|
||||
# No unyielded delta (everything already sent)
|
||||
types = [p["type"] for p in payloads]
|
||||
assert types == ["content_block_stop", "message_delta", "message_stop"]
|
||||
|
||||
def test_format_response_with_stop_trims_content(self, builder):
|
||||
ctx = _make_ctx()
|
||||
stop = StopInfo(matched="STOP", body="text STOP extra", yielded="text ")
|
||||
resp = builder.format_response(ctx, "text STOP extra", stop)
|
||||
assert resp["content"][0]["text"] == "text "
|
||||
assert resp["stop_reason"] == "stop_sequence"
|
||||
assert resp["stop_sequence"] == "STOP"
|
||||
|
||||
def test_format_response_no_stop(self, builder):
|
||||
ctx = _make_ctx()
|
||||
stop = StopInfo()
|
||||
resp = builder.format_response(ctx, "full text", stop)
|
||||
assert resp["content"][0]["text"] == "full text"
|
||||
assert resp["stop_reason"] == "end_turn"
|
||||
|
||||
|
||||
class TestGenerationRequestValidation:
|
||||
def test_valid_params(self):
|
||||
gr = GenerationRequest(
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
top_k=50,
|
||||
top_p=0.9,
|
||||
temperature=0.7,
|
||||
)
|
||||
assert gr.top_k == 50
|
||||
|
||||
def test_invalid_top_p_raises(self):
|
||||
with pytest.raises(ValueError, match="top_p"):
|
||||
GenerationRequest(messages=[{"role": "user", "content": "hi"}], top_p=1.5)
|
||||
|
||||
def test_invalid_top_k_raises(self):
|
||||
with pytest.raises(ValueError, match="top_k"):
|
||||
GenerationRequest(messages=[{"role": "user", "content": "hi"}], top_k=-1)
|
||||
|
||||
def test_invalid_temperature_raises(self):
|
||||
with pytest.raises(ValueError, match="temperature"):
|
||||
GenerationRequest(
|
||||
messages=[{"role": "user", "content": "hi"}], temperature=-0.1
|
||||
)
|
||||
|
||||
def test_top_k_zero_valid(self):
|
||||
gr = GenerationRequest(messages=[{"role": "user", "content": "hi"}], top_k=0)
|
||||
assert gr.top_k == 0
|
||||
@@ -173,3 +173,21 @@ def test_scheduler_concurrent_get_stats(mock_model_and_tokenizer):
|
||||
for stats in results["stats"]:
|
||||
assert "total_tasks" in stats
|
||||
assert stats["total_tasks"] >= 0
|
||||
|
||||
|
||||
def test_prefill_skips_fully_cached_tasks(mock_model_and_tokenizer):
|
||||
"""Tasks whose entire prompt is cached skip the prefill phase."""
|
||||
mock_model, mock_tokenizer = mock_model_and_tokenizer
|
||||
|
||||
with patch("astrai.inference.core.scheduler.AutoModel"):
|
||||
with patch("astrai.inference.core.scheduler.AutoTokenizer"):
|
||||
scheduler = InferenceScheduler(
|
||||
model=mock_model,
|
||||
tokenizer=mock_tokenizer,
|
||||
max_batch_size=4,
|
||||
device="cpu",
|
||||
)
|
||||
|
||||
task_id = scheduler.add_task("short prompt", stream_callback=lambda t: None)
|
||||
scheduler.stop()
|
||||
assert task_id.startswith("task_")
|
||||
|
||||
@@ -0,0 +1,355 @@
|
||||
import tempfile
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||
from astrai.model import AutoRegressiveLM
|
||||
from astrai.model.components.linear import Linear
|
||||
from astrai.model.components.lora import (
|
||||
LoRAConfig,
|
||||
LoRALinear,
|
||||
_collect_lora_info,
|
||||
_get_lora_count,
|
||||
inject_lora,
|
||||
load_lora,
|
||||
merge_lora,
|
||||
save_lora,
|
||||
)
|
||||
|
||||
MODEL_KWARGS = dict(
|
||||
vocab_size=1000,
|
||||
dim=64,
|
||||
n_heads=4,
|
||||
n_kv_heads=2,
|
||||
dim_ffn=128,
|
||||
n_layers=2,
|
||||
max_len=32,
|
||||
norm_eps=1e-5,
|
||||
)
|
||||
|
||||
|
||||
def _make_model(**kwargs):
|
||||
kw = {**MODEL_KWARGS, **kwargs}
|
||||
config = AutoRegressiveLMConfig(**kw)
|
||||
model = AutoRegressiveLM(config)
|
||||
model.eval()
|
||||
return model
|
||||
|
||||
|
||||
def test_loralinear_init():
|
||||
base = Linear(64, 128)
|
||||
lora = LoRALinear(base, r=8, alpha=16)
|
||||
|
||||
assert lora.weight is base.weight
|
||||
assert not lora.weight.requires_grad
|
||||
assert lora.lora_A.shape == (8, 64)
|
||||
assert lora.lora_B.shape == (128, 8)
|
||||
assert lora.scaling == 2.0
|
||||
assert not lora._merged
|
||||
assert lora.lora_A.requires_grad
|
||||
assert lora.lora_B.requires_grad
|
||||
|
||||
|
||||
def test_loralinear_forward_init_zero_delta():
|
||||
base = Linear(4, 4)
|
||||
with torch.no_grad():
|
||||
base.weight.zero_()
|
||||
|
||||
x = torch.randn(2, 4)
|
||||
lora = LoRALinear(base, r=2, alpha=2)
|
||||
base_out = base(x)
|
||||
lora_out = lora(x)
|
||||
|
||||
assert torch.allclose(base_out, lora_out)
|
||||
|
||||
|
||||
def test_loralinear_forward_with_delta():
|
||||
base = Linear(4, 4)
|
||||
with torch.no_grad():
|
||||
base.weight.zero_()
|
||||
|
||||
x = torch.randn(2, 4)
|
||||
lora = LoRALinear(base, r=2, alpha=2)
|
||||
base_out = base(x)
|
||||
|
||||
with torch.no_grad():
|
||||
lora.lora_B.fill_(1.0)
|
||||
|
||||
lora_out = lora(x)
|
||||
assert not torch.allclose(base_out, lora_out)
|
||||
|
||||
|
||||
def test_loralinear_merge():
|
||||
base = Linear(4, 4)
|
||||
with torch.no_grad():
|
||||
base.weight.zero_()
|
||||
|
||||
x = torch.randn(2, 4)
|
||||
lora = LoRALinear(base, r=2, alpha=2)
|
||||
with torch.no_grad():
|
||||
lora.lora_B.fill_(1.0)
|
||||
|
||||
out_before = lora(x).clone()
|
||||
lora.merge()
|
||||
out_after = lora(x)
|
||||
|
||||
torch.testing.assert_close(out_before, out_after)
|
||||
assert lora._merged
|
||||
assert not hasattr(lora, "lora_A")
|
||||
|
||||
|
||||
def test_loralinear_merge_is_idempotent():
|
||||
base = Linear(4, 4)
|
||||
with torch.no_grad():
|
||||
base.weight.zero_()
|
||||
|
||||
lora = LoRALinear(base, r=2, alpha=2)
|
||||
with torch.no_grad():
|
||||
lora.lora_B.fill_(1.0)
|
||||
|
||||
lora.merge()
|
||||
lora.merge()
|
||||
|
||||
|
||||
def test_inject_lora_default_target():
|
||||
model = _make_model()
|
||||
n_before = sum(1 for m in model.modules() if isinstance(m, Linear))
|
||||
|
||||
inject_lora(model, r=4, alpha=8)
|
||||
|
||||
lora_count = _get_lora_count(model)
|
||||
assert lora_count > 0
|
||||
assert lora_count < n_before
|
||||
|
||||
|
||||
def test_inject_lora_ffn():
|
||||
model = _make_model()
|
||||
from astrai.model.components.lora import TARGET_MODULES_FFN
|
||||
|
||||
inject_lora(model, r=4, alpha=8, target_modules=TARGET_MODULES_FFN)
|
||||
assert _get_lora_count(model) > 0
|
||||
|
||||
|
||||
def test_inject_lora_returns_config():
|
||||
model = _make_model()
|
||||
cfg = inject_lora(model, r=8, alpha=32)
|
||||
assert isinstance(cfg, LoRAConfig)
|
||||
assert cfg.r == 8
|
||||
assert cfg.alpha == 32
|
||||
|
||||
|
||||
def test_inject_lora_no_matching_targets_warns(caplog):
|
||||
model = _make_model()
|
||||
inject_lora(model, r=4, alpha=8, target_modules={"nonexistent"})
|
||||
assert "No LoRA layers injected" in caplog.text
|
||||
|
||||
|
||||
def test_inject_lora_preserves_base_output():
|
||||
model = _make_model()
|
||||
x = torch.randint(0, 1000, (2, 16))
|
||||
|
||||
with torch.no_grad():
|
||||
out_before = model(x)["logits"].clone()
|
||||
|
||||
inject_lora(model, r=4, alpha=8)
|
||||
|
||||
with torch.no_grad():
|
||||
out_after = model(x)["logits"]
|
||||
|
||||
torch.testing.assert_close(out_before, out_after)
|
||||
|
||||
|
||||
def test_inject_lora_does_not_reinject():
|
||||
model = _make_model()
|
||||
inject_lora(model, r=4, alpha=8, target_modules={"q_proj"})
|
||||
first_count = _get_lora_count(model)
|
||||
|
||||
inject_lora(model, r=2, alpha=4, target_modules={"q_proj"})
|
||||
assert _get_lora_count(model) == first_count
|
||||
|
||||
|
||||
def test_inject_lora_adds_new_modules():
|
||||
model = _make_model()
|
||||
inject_lora(model, r=4, alpha=8, target_modules={"q_proj"})
|
||||
first = _get_lora_count(model)
|
||||
|
||||
inject_lora(model, r=4, alpha=8, target_modules={"v_proj"})
|
||||
assert _get_lora_count(model) > first
|
||||
|
||||
|
||||
def test_inject_lora_on_mla_model():
|
||||
model = _make_model(
|
||||
attn_type="mla", kv_lora_rank=16, qk_nope_head_dim=16, qk_rope_head_dim=16
|
||||
)
|
||||
inject_lora(model, r=4, alpha=8, target_modules={"q_proj", "o_proj"})
|
||||
assert _get_lora_count(model) > 0
|
||||
|
||||
|
||||
def test_inject_lora_on_moe_model():
|
||||
model = _make_model(
|
||||
ffn_type="moe",
|
||||
n_routed_experts=4,
|
||||
n_shared_experts=1,
|
||||
n_activated_experts=2,
|
||||
dim_ffn=32,
|
||||
)
|
||||
inject_lora(model, r=4, alpha=8, target_modules={"up", "gate", "down"})
|
||||
assert _get_lora_count(model) > 0
|
||||
|
||||
|
||||
def test_state_dict_key_format():
|
||||
model = _make_model()
|
||||
inject_lora(model, r=4, alpha=8, target_modules={"q_proj"})
|
||||
|
||||
sd = model.state_dict()
|
||||
assert "layers.0.attention.q_proj.weight" in sd
|
||||
assert "layers.0.attention.q_proj.lora_A" in sd
|
||||
assert "layers.0.attention.q_proj.lora_B" in sd
|
||||
|
||||
|
||||
def test_only_lora_params_trainable():
|
||||
model = _make_model()
|
||||
inject_lora(model, r=4, alpha=8, target_modules={"q_proj", "v_proj"})
|
||||
|
||||
for name, param in model.named_parameters():
|
||||
if isinstance(name.split(".")[-1], str) and "lora" in name:
|
||||
assert param.requires_grad, f"lora param should be trainable: {name}"
|
||||
elif any(name.endswith(f".{t}.weight") for t in ("q_proj", "v_proj")):
|
||||
assert not param.requires_grad, f"injected weight should be frozen: {name}"
|
||||
|
||||
|
||||
def test_state_dict_after_inject_consistent_with_original():
|
||||
model = _make_model()
|
||||
sd_before = {k: v for k, v in model.state_dict().items()}
|
||||
|
||||
inject_lora(model, r=4, alpha=8, target_modules={"q_proj"})
|
||||
sd_after = model.state_dict()
|
||||
|
||||
# original keys unchanged
|
||||
for k in sd_before:
|
||||
assert k in sd_after
|
||||
assert sd_before[k].shape == sd_after[k].shape
|
||||
|
||||
# new lora keys present
|
||||
lora_keys = [k for k in sd_after if "lora" in k]
|
||||
assert len(lora_keys) > 0
|
||||
|
||||
|
||||
def test_save_load_roundtrip():
|
||||
model = _make_model()
|
||||
cfg = inject_lora(model, r=4, alpha=8, target_modules={"q_proj"})
|
||||
|
||||
with torch.no_grad():
|
||||
for m in model.modules():
|
||||
if isinstance(m, LoRALinear):
|
||||
m.lora_B.fill_(0.5)
|
||||
|
||||
x = torch.randint(0, 1000, (2, 16))
|
||||
with torch.no_grad():
|
||||
out_src = model(x)["logits"].clone()
|
||||
|
||||
tmpdir = tempfile.mkdtemp()
|
||||
save_lora(model, tmpdir, cfg)
|
||||
|
||||
model2 = _make_model()
|
||||
model2.load_state_dict(model.state_dict(), strict=False)
|
||||
load_lora(model2, tmpdir)
|
||||
|
||||
with torch.no_grad():
|
||||
out_dst = model2(x)["logits"]
|
||||
|
||||
torch.testing.assert_close(out_src, out_dst)
|
||||
|
||||
|
||||
def test_save_after_merge_raises():
|
||||
model = _make_model()
|
||||
cfg = inject_lora(model, r=4, alpha=8, target_modules={"q_proj"})
|
||||
|
||||
with torch.no_grad():
|
||||
for m in model.modules():
|
||||
if isinstance(m, LoRALinear):
|
||||
m.lora_B.fill_(0.5)
|
||||
|
||||
tmpdir = tempfile.mkdtemp()
|
||||
save_lora(model, tmpdir, cfg)
|
||||
merge_lora(model)
|
||||
|
||||
tmpdir2 = tempfile.mkdtemp()
|
||||
with pytest.raises(RuntimeError, match="No LoRA parameters"):
|
||||
save_lora(model, tmpdir2, cfg)
|
||||
|
||||
|
||||
def test_load_lora_on_already_injected():
|
||||
model = _make_model()
|
||||
inject_lora(model, r=4, alpha=8, target_modules={"q_proj"})
|
||||
|
||||
with torch.no_grad():
|
||||
for m in model.modules():
|
||||
if isinstance(m, LoRALinear):
|
||||
m.lora_B.fill_(0.5)
|
||||
|
||||
tmpdir = tempfile.mkdtemp()
|
||||
save_lora(model, tmpdir, LoRAConfig(r=4, alpha=8, target_modules=("q_proj",)))
|
||||
|
||||
model2 = _make_model()
|
||||
model2.load_state_dict(model.state_dict(), strict=False)
|
||||
inject_lora(model2, r=4, alpha=8, target_modules={"q_proj"})
|
||||
|
||||
# load onto already-injected model
|
||||
load_lora(model2, tmpdir)
|
||||
assert _get_lora_count(model2) > 0
|
||||
|
||||
|
||||
def test_load_lora_mismatched_r_raises():
|
||||
model = _make_model()
|
||||
cfg = inject_lora(model, r=8, alpha=16, target_modules={"q_proj"})
|
||||
|
||||
with torch.no_grad():
|
||||
for m in model.modules():
|
||||
if isinstance(m, LoRALinear):
|
||||
m.lora_B.fill_(0.5)
|
||||
|
||||
tmpdir = tempfile.mkdtemp()
|
||||
save_lora(model, tmpdir, cfg)
|
||||
|
||||
model2 = _make_model()
|
||||
model2.load_state_dict(model.state_dict(), strict=False)
|
||||
inject_lora(model2, r=4, alpha=8, target_modules={"q_proj"})
|
||||
|
||||
with pytest.raises(RuntimeError, match="size mismatch"):
|
||||
load_lora(model2, tmpdir) # strict=False, only lora keys
|
||||
|
||||
|
||||
def test_merge_preserves_output():
|
||||
model = _make_model()
|
||||
inject_lora(model, r=4, alpha=8, target_modules={"q_proj"})
|
||||
|
||||
with torch.no_grad():
|
||||
for m in model.modules():
|
||||
if isinstance(m, LoRALinear):
|
||||
m.lora_B.fill_(0.5)
|
||||
|
||||
x = torch.randint(0, 1000, (2, 16))
|
||||
with torch.no_grad():
|
||||
out_before = model(x)["logits"].clone()
|
||||
|
||||
merge_lora(model)
|
||||
|
||||
with torch.no_grad():
|
||||
out_after = model(x)["logits"]
|
||||
torch.testing.assert_close(out_before, out_after)
|
||||
|
||||
|
||||
def test_merge_no_lora_warns(caplog):
|
||||
model = _make_model()
|
||||
merge_lora(model)
|
||||
assert "No LoRA layers to merge" in caplog.text
|
||||
|
||||
|
||||
def test_collect_lora_info():
|
||||
model = _make_model()
|
||||
info = _collect_lora_info(model)
|
||||
assert "q_proj" in info
|
||||
assert "o_proj" in info
|
||||
assert "q_proj" in info # each layer has one
|
||||
@@ -1,3 +1,5 @@
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
@@ -73,6 +75,7 @@ def create_train_config(
|
||||
optimizer_fn=optimizer_fn,
|
||||
scheduler_fn=scheduler_fn,
|
||||
ckpt_dir=test_dir,
|
||||
log_dir=os.path.join(test_dir, "logs"),
|
||||
n_epoch=n_epoch,
|
||||
batch_per_device=batch_per_device,
|
||||
ckpt_interval=ckpt_interval,
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.config.train_config import TrainConfig
|
||||
@@ -110,6 +112,7 @@ def test_gradient_checkpointing_trainer_integration(base_test_env, random_datase
|
||||
optimizer_fn=optimizer_fn,
|
||||
scheduler_fn=scheduler_fn,
|
||||
ckpt_dir=base_test_env["test_dir"],
|
||||
log_dir=os.path.join(base_test_env["test_dir"], "logs"),
|
||||
n_epoch=1,
|
||||
batch_per_device=2,
|
||||
ckpt_interval=3,
|
||||
@@ -143,6 +146,7 @@ def test_callback_integration(base_test_env, random_dataset):
|
||||
optimizer_fn=optimizer_fn,
|
||||
scheduler_fn=scheduler_fn,
|
||||
ckpt_dir=base_test_env["test_dir"],
|
||||
log_dir=os.path.join(base_test_env["test_dir"], "logs"),
|
||||
n_epoch=1,
|
||||
batch_per_device=2,
|
||||
ckpt_interval=3,
|
||||
|
||||
@@ -27,6 +27,7 @@ def test_early_stopping_simulation(base_test_env, early_stopping_dataset):
|
||||
model=base_test_env["model"],
|
||||
dataset=early_stopping_dataset,
|
||||
ckpt_dir=base_test_env["test_dir"],
|
||||
log_dir=os.path.join(base_test_env["test_dir"], "logs"),
|
||||
n_epoch=2,
|
||||
batch_per_device=2,
|
||||
ckpt_interval=1,
|
||||
|
||||
Reference in New Issue
Block a user