- backend.py → executor.py,BaseTrainingBackend → BaseExecutor - 新增 NoneExecutor(单卡)和 DDPExecutor(DDP,world_size=1 自动降级) - 新增 GradientState 分离梯度同步状态,AccumOptimizer/AccumScheduler 包裹拦截 - 新增 astrai/protocols.py:OptimizerProtocol/SchedulerProtocol 结构子类型 - TrainContext.backend → executor,TrainConfig 移除 parallel_wrapper/state_dict_fn,新增 parallel_mode/executor_kwargs - 训练循环用 accumulate() 包裹,on_optimizer_step 命名约定=gate - scripts/tools/train.py 移除 ddp_wrap/prepare_checkpoint,新增 --parallel_mode
22 lines
598 B
Python
22 lines
598 B
Python
"""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): ...
|