跳转至

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 分布式版本。主要差异在于:

  1. 使用 RayWorkerGroup 而非直接的 TrainingWorker
  2. 远程调用需要 .get() 获取结果
  3. 数据并行由 Worker Group 内部管理
  4. Checkpoint 使用 Ray 编排模式

两者的训练逻辑完全一致,只是"编排方式"不同。