refactor: 重构训练后端为 Executor 模式

- 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
This commit is contained in:
2026-05-24 20:35:44 +08:00
parent 8cbf3f36e2
commit 3ab4f237e5
8 changed files with 153 additions and 104 deletions
+18 -26
View File
@@ -4,14 +4,11 @@ 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.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,
):
@@ -271,6 +256,13 @@ def train(
"sync_interval": grpo_sync_interval,
}
executor_kwargs = {
"static_graph": True,
"find_unused_parameters": False,
"gradient_as_bucket_view": True,
"broadcast_buffers": False,
}
dataset = DatasetFactory.load(
train_type=train_type,
load_path=data_root_path,
@@ -319,10 +311,10 @@ 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,
)