refactor: 优化并行训练配置与启动管理
- 配置新增 start_method 支持 spawn/fork/forkserver 选择 - 启动方式 mp.spawn 改为 mp.start_processes,支持 daemon=True - validate() 改为基于 metadata 的反射式校验,不再硬编码字段列表 - CLI 新增 --start_method 参数
This commit is contained in:
@@ -149,6 +149,13 @@ def parse_args() -> argparse.Namespace:
|
||||
parser.add_argument(
|
||||
"--device_type", type=str, default="cuda", help="Device type to use."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--start_method",
|
||||
type=str,
|
||||
default="spawn",
|
||||
choices=["spawn", "fork", "forkserver"],
|
||||
help="Multiprocessing start method.",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
@@ -232,6 +239,7 @@ def train(
|
||||
stride: int,
|
||||
nprocs: int,
|
||||
device_type: str,
|
||||
start_method: str,
|
||||
):
|
||||
assert train_type in ["seq", "sft", "dpo", "grpo"]
|
||||
assert os.path.exists(param_path)
|
||||
@@ -314,6 +322,7 @@ def train(
|
||||
parallel_wrapper=ddp_wrap,
|
||||
state_dict_fn=prepare_checkpoint,
|
||||
device_type=device_type,
|
||||
start_method=start_method,
|
||||
extra_kwargs=strategy_kwargs,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user