sft_trainer_ray.py — 这是 SFT 的 Ray 分布式版本训练器¶
文件概述¶
模块路径: verl.trainer.sft_trainer_ray
这是 SFT 的 Ray 分布式版本训练器。与 sft_trainer.py(单机 SPMD 模式)不同,这个版本使用 Ray 框架来管理分布式训练,通过 RayWorkerGroup 将训练任务分发到远程 Worker 上执行。
在训练流程中的位置¶
与 sft_trainer.py 功能相同(都是 SFT 训练),但部署方式不同。本文件适用于需要 Ray 集群管理的场景。
关键代码讲解¶
1. 与单机版本的主要区别 - 引擎构建¶
def _build_engine(self):
from verl.single_controller.ray import RayClassWithInitArgs, RayResourcePool, RayWorkerGroup
n_gpus_per_node = self.config.trainer.n_gpus_per_node
nnodes = self.config.trainer.nnodes
# 创建 Ray 资源池
self.resource_pool = RayResourcePool(process_on_nodes=[n_gpus_per_node] * nnodes)
# 将 TrainingWorker 包装为 Ray Remote Actor
ray_cls_with_init = RayClassWithInitArgs(ray.remote(TrainingWorker), config=config)
# 创建 RayWorkerGroup(管理多个 Worker 实例)
self.training_client = RayWorkerGroup(
resource_pool=self.resource_pool,
ray_cls_with_init=ray_cls_with_init,
device_name=self.config.trainer.device,
)
self.training_client.set_loss_fn(loss_fn=self.loss_fn)
self.training_client.reset()
关键区别:单机版直接创建 TrainingWorker,而 Ray 版本创建 RayWorkerGroup,后者管理一组通过 Ray 远程调度的 Worker。
2. 数据加载器差异¶
def _build_dataloader(self):
# Ray 模式下,dp_rank 和 dp_size 默认为 0 和 1
# 因为实际的数据并行由 RayWorkerGroup 内部处理
dp_rank = 0
dp_size = 1
在 Ray 模式下,数据并行的切分由 Worker Group 内部处理,驱动进程看到的是完整数据。
3. 训练循环差异¶
def fit(self):
for step_in_epoch, data in enumerate(tqdm(...)):
# 训练
output = self.training_client.train_batch(data)
output = output.get() # Ray 模式需要 .get() 获取远程结果
# 验证
output = self.training_client.infer_batch(val_data)
output = output.get() # 同样需要 .get()
注意 .get() 调用:Ray 模式下 train_batch 返回的是一个 future,需要 .get() 来获取实际结果。
4. Checkpoint 处理差异¶
def _build_ckpt_handler(self):
self.ckpt_handler = CheckpointHandler(
engine=self.training_client, # 注意:这里传的是 RayWorkerGroup
...
mode=OrchestrationMode.RAY, # 指定 Ray 编排模式
)
Checkpoint 保存需要适配 Ray 的编排模式。
5. 序列长度计算简化¶
def _get_batch_seqlens(self, data):
is_nested = data["input_ids"].is_nested
if is_nested:
batch_seqlens = data["input_ids"].offsets().diff()
else:
batch_seqlens = data["attention_mask"].sum(dim=-1)
return batch_seqlens # 不需要 all_gather,因为在驱动进程上计算
由于数据在驱动进程上是完整的,不需要像 SPMD 模式那样做 all_gather。
核心类/函数列表¶
| 名称 | 类型 | 作用 |
|---|---|---|
SFTTrainer |
class | SFT Ray 分布式训练器 |
run_sft(config) |
function | 初始化 Ray 并启动训练 |
create_sft_dataset(...) |
function | 创建 SFT 数据集 |
单机版 vs Ray 版对比¶
| 特性 | sft_trainer.py (SPMD) | sft_trainer_ray.py (Ray) |
|---|---|---|
| 分布式框架 | torch.distributed | Ray |
| 训练引擎 | TrainingWorker (本地) | RayWorkerGroup (远程) |
| 数据分发 | DistributedSampler | 由 WorkerGroup 内部处理 |
| 结果获取 | 直接返回 | 需要 .get() |
| Checkpoint 模式 | SPMD | RAY |
| 适用场景 | 单机多卡 | 多机多卡集群 |
数据流和调用关系¶
main()
+-- run_sft(config)
+-- ray.init()
+-- SFTTrainer(config)
| +-- _build_config()
| +-- _build_dataset()
| +-- _build_dataloader()
| +-- _build_engine()
| | +-- RayResourcePool(...)
| | +-- RayWorkerGroup(...)
| +-- _build_ckpt_handler(mode=RAY)
|
+-- trainer.fit()
+-- training_client.train_batch(data).get()
+-- training_client.infer_batch(data).get()
+-- ckpt_handler.save_checkpoint()
小结¶
sft_trainer_ray.py 是 sft_trainer.py 的 Ray 分布式版本。主要差异在于:
- 使用
RayWorkerGroup而非直接的TrainingWorker - 远程调用需要
.get()获取结果 - 数据并行由 Worker Group 内部管理
- Checkpoint 使用 Ray 编排模式
两者的训练逻辑完全一致,只是"编排方式"不同。