refactor : 并行启动 Strategy 模式重构,local_rank 解耦

- setup_parallel 接收 local_rank 参数,不再读环境变量推导
- TorchrunStrategy 从 env 读取 LOCAL_RANK,LocalStrategy 用 rank
- _detect_launcher() 分级检测替代内联 RANK 检查
- _run_single_rank 统一入口,消除 _run_single/_run_multi 重复
- 优雅退出:except BaseException 终止子进程并 re-join
- gradient_checkpointing_modules 判定提取到外部变量
This commit is contained in:
2026-06-02 11:22:24 +08:00
parent d6899100ac
commit 9b416c1bbb
2 changed files with 129 additions and 67 deletions
+3 -3
View File
@@ -315,6 +315,8 @@ def train(
},
)
grad_ckpt_modules = [DecoderBlock] if gradient_checkpointing else []
train_config = TrainConfig(
model_fn=model_fn,
strategy=train_type,
@@ -332,9 +334,6 @@ def train(
random_seed=random_seed,
num_workers=num_workers,
pin_memory=pin_memory,
gradient_checkpointing_modules=[DecoderBlock]
if gradient_checkpointing
else [],
nprocs=nprocs,
backend=backend,
master_addr=master_addr,
@@ -342,6 +341,7 @@ def train(
parallel_mode=parallel_mode,
device_type=device_type,
start_method=start_method,
gradient_checkpointing_modules=grad_ckpt_modules,
executor_kwargs=executor_kwargs,
extra_kwargs=strategy_kwargs,
)