跳转至

sft_trainer.py — 这是 SFT (Supervised Fine-Tuning) 单机训练器

文件概述

模块路径: verl.trainer.sft_trainer

这是 SFT (Supervised Fine-Tuning) 单机训练器,使用 SPMD (Single Program, Multiple Data) 模式进行分布式训练。它直接使用 PyTorch 的 torch.distributed 进行多 GPU 并行,不依赖 Ray。

SFT 是 RLHF 流程的第一步:在进行强化学习之前,先用监督学习微调基础模型,使其具备基本的对话和遵循指令能力。

在训练流程中的位置

SFT 训练是 RLHF 流程的前置步骤,它发生在 PPO 训练之前。

预训练模型 --> [SFT 训练] --> SFT 模型 --> [PPO 训练] --> RL 对齐模型
                 ↑                                 ↑
            (本文件)                         (main_ppo.py)

关键代码讲解

1. SFTTrainer 初始化

class SFTTrainer:
    def __init__(self, config):
        self.config = config
        self.rank = torch.distributed.get_rank()

        self._build_config()      # 解析配置为 dataclass
        self._build_dataset()     # 构建数据集
        self._build_engine()      # 构建训练引擎
        self._build_dataloader()  # 构建数据加载器
        self._init_engine()       # 初始化引擎(设置总步数等)
        self._build_ckpt_handler()  # 构建 checkpoint 处理器

        # 加载 checkpoint(如果有的话)
        self.resume_global_step = self.ckpt_handler.load_checkpoint()

初始化过程非常清晰,按顺序构建各个组件。

2. 训练引擎构建

    def _build_engine(self):
        from verl.workers.engine_workers import TrainingWorkerConfig
        from verl.workers.utils.losses import sft_loss

        self.loss_fn = partial(sft_loss, config=None)

        config = TrainingWorkerConfig(
            model_type="language_model",
            model_config=self.model_config,
            engine_config=self.engine_config,
            optimizer_config=self.optimizer_config,
            ...
        )

        self.training_client = TrainingWorker(config=config)
        self.training_client.set_loss_fn(loss_fn=self.loss_fn)
        self.engine = self.training_client.engine

使用 TrainingWorker 作为训练引擎,SFT 的损失函数是标准的交叉熵损失(sft_loss)。

3. 数据加载器构建

    def _build_dataloader(self):
        dp_rank = self.engine.get_data_parallel_rank()
        dp_size = self.engine.get_data_parallel_size()

        self.train_sampler = DistributedSampler(
            self.train_dataset, shuffle=True, num_replicas=dp_size, rank=dp_rank, drop_last=True
        )

        self.train_dataloader = StatefulDataLoader(
            dataset=self.train_dataset,
            batch_size=self.train_batch_size_per_dp,
            sampler=self.train_sampler,
            collate_fn=self.collate_fn,
            ...
        )

使用 DistributedSampler 将数据分配到不同的 data parallel rank 上。StatefulDataLoader 支持断点恢复——即使训练中断,也能从上次停止的位置继续。

4. 训练循环

    def fit(self):
        tracking = Tracking(...)
        global_step = self.resume_global_step

        for epoch in range(start_epoch, self.config.trainer.total_epochs):
            self.train_sampler.set_epoch(epoch=epoch)
            for step_in_epoch, data in enumerate(tqdm(self.train_dataloader, ...)):
                global_step += 1

                # 构造 tensordict
                data = tu.get_tensordict(tensor_dict=data, non_tensor_dict=meta_info)
                batch_seqlens = self._get_batch_seqlens(data=data)
                tu.assign_non_tensor(data, update_lr_scheduler=True, ...)

                # 训练一个 batch
                output = self.training_client.train_batch(data=data)

                # 记录 metrics
                if self.engine.is_mp_src_rank_with_outputs():
                    metrics = tu.get(output, "metrics")
                    tracking.log(data=metrics, step=global_step)

                # 验证
                if is_valid_step:
                    val_losses = []
                    for val_data in self.val_dataloader:
                        output = self.training_client.infer_batch(val_data)
                        val_losses.append(metrics["loss"])
                    val_loss = torch.mean(torch.tensor(val_losses, ...))

                # 保存 checkpoint
                if is_save_step:
                    self.ckpt_handler.save_checkpoint(step=global_step)

训练循环是标准的 epoch-based 训练,支持: - 学习率调度(通过 update_lr_scheduler 标记) - 定期验证 - 定期保存 checkpoint - 性能分析(profiling)

5. 批次序列长度计算

    def _get_batch_seqlens(self, data):
        is_nested = data["input_ids"].is_nested
        if is_nested:
            # Nested Tensor 格式
            batch_seqlens = data["input_ids"].offsets().diff()
        else:
            # 传统 padding 格式
            batch_seqlens = data["attention_mask"].sum(dim=-1)
        batch_seqlens = batch_seqlens.to(self.device_name)

        # 跨 DP group 收集所有序列长度
        dp_group = self.engine.get_data_parallel_group()
        torch.distributed.all_gather_into_tensor(output_tensor, batch_seqlens, group=dp_group)
        return output_tensor.tolist()

支持两种数据格式:Nested Tensor(去除 padding 的高效格式)和传统的 padding + attention_mask 格式。

核心类/函数列表

名称 类型 作用
SFTTrainer class SFT 单机训练器主类
SFTTrainer._build_config() method 将 OmegaConf 配置转为 dataclass
SFTTrainer._build_engine() method 构建 TrainingWorker 训练引擎
SFTTrainer._build_dataset() method 构建训练/验证数据集
SFTTrainer._build_dataloader() method 构建分布式数据加载器
SFTTrainer.fit() method 执行完整训练循环
run_sft(config) function 初始化分布式环境并启动训练
create_sft_dataset(...) function 创建 SFT 数据集(支持自定义数据集类)

数据流和调用关系

main()
  |
  +-- auto_set_device()
  +-- run_sft(config)
        |
        +-- initialize_global_process_group()
        +-- SFTTrainer(config)
        |       |
        |       +-- _build_config()
        |       +-- _build_dataset() --> create_sft_dataset()
        |       +-- _build_engine() --> TrainingWorker
        |       +-- _build_dataloader() --> StatefulDataLoader
        |       +-- _init_engine()
        |       +-- _build_ckpt_handler()
        |
        +-- trainer.fit()
        |       |
        |       +-- for epoch:
        |       |     for batch:
        |       |       train_batch() --> output
        |       |       [validate]
        |       |       [save_checkpoint]
        |
        +-- destroy_global_process_group()

小结

sft_trainer.py 实现了一个清晰的 SFT 训练流程:

  1. SPMD 模式:直接使用 PyTorch 分布式,每个进程运行相同代码
  2. 断点恢复:通过 StatefulDataLoader 和 CheckpointHandler 实现
  3. 灵活数据格式:支持 Nested Tensor 和传统 padding 两种格式
  4. 模块化设计:构建过程清晰分离(config -> dataset -> engine -> dataloader)

与 PPO 训练器相比,SFT 训练器更加简单直接,因为它只需要一个模型(不需要 Critic、Ref Policy 等)。