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 训练之前。
关键代码讲解¶
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 训练流程:
- SPMD 模式:直接使用 PyTorch 分布式,每个进程运行相同代码
- 断点恢复:通过
StatefulDataLoader和CheckpointHandler实现 - 灵活数据格式:支持 Nested Tensor 和传统 padding 两种格式
- 模块化设计:构建过程清晰分离(config -> dataset -> engine -> dataloader)
与 PPO 训练器相比,SFT 训练器更加简单直接,因为它只需要一个模型(不需要 Critic、Ref Policy 等)。