异步GRPO结合LoRA在HF Jobs上的实践:用存储桶和代理替代NCCL

1 阅读

架构概览:利用HF Jobs和存储桶

最近,TRL库的AsyncGRPOTrainer增加了对LoRA的支持(v1.14版本),使得异步强化学习训练可以只同步几MB大小的适配器,而非完整的3GB模型权重。这一特性让我们能在Hugging Face Jobs上构建一个完全分离的训练-推理架构:

  • 训练Job:运行带LoRA的AsyncGRPOTrainer(使用FSDP)
  • 两个vLLM推理Job:各占1个GPU,动态加载最新适配器
  • 共享存储桶:通过FUSE挂载到所有Job,作为适配器传输通道
  • 轻量代理:处理认证、路由和适配器广播

关键突破在于利用HF Storage Buckets替代了传统集群中的NCCL通信。由于LoRA适配器仅几MB大小,通过存储桶同步的延迟完全可以接受。每个Job通过hf-mount将同一存储桶挂载到容器内的相同路径(如/lora),训练Job写入适配器,推理Job读取使用。

存储桶还用于保存检查点,确保训练Job被抢占后能从中断处恢复,而不会丢失任何进度。

三个核心组件

vLLM推理实例

每个推理Job使用官方vllm/vllm-openai:v0.27.1镜像,关键配置包括:

# 启用运行时LoRA加载
-e VLLM_ALLOW_RUNTIME_LORA_UPDATING=1 
# 开发模式(需/pause等端点)
-e VLLM_SERVER_DEV_MODE=1
# 适配器槽位 = max_staleness + 2
--max-loras 6

这里max_staleness=4意味着训练器会同时保留5个版本的适配器(当前版+前4版)。vLLM需要额外1个槽位用于原子切换,故设为6。若设为5,每次同步时vLLM会意外驱逐仍有推理任务在用的旧适配器。

特别要注意的是,我们为每个适配器使用唯一版本号(如trl-policy-v3)。如果始终用相同名称覆盖,vLLM的KV缓存可能混用不同版本权重计算的结果——因为缓存键只认名称不认内容。版本化命名彻底规避了这个问题。

数据集选择:Sanity测试集

实验采用sail/Sanity-Test-R1D-1.5B数据集,源自论文《Defeating the Training-Inference Mismatch via FP16》。该数据集包含1,460道MATH问题,每题有40个模型生成的答案,且筛选条件为原始成功率在20%-80%之间。这种设计确保了:

  1. 问题既非已完全解决也非完全无解
  2. 模型能获得清晰的早期训练信号
  3. 小规模数据集可在2小时内完成一轮训练

超参数直接沿用论文中的LoRA配置:Qwen2.5-Math-1.5B基座模型、rank=1(alpha=2)、学习率4e-5、每提示8个样本、每步128个完成、最大生成3,000 token。

训练器配置

训练Job同样基于vLLM镜像,额外安装TRL库。关键配置如下:

AsyncGRPOConfig(
    output_dir="/lora/sanity-lora-r1",  # 存储桶路径
    vllm_server_base_url="http://localhost:8000",  # 指向本地代理
    max_staleness=4,
    weight_sync_steps=4,  # 每4步同步一次适配器
    save_steps=50,        # 检查点也存入存储桶
)
peft_config=LoraConfig(r=1, lora_alpha=2, target_modules="all-linear")

Figure 1

初始化时,TRL会调用/server_info检测vLLM是否支持适配器模式。日志中出现"Adapter-only vLLM sync enabled"即表示配置成功。

Figure 2

代理的核心作用

Figure 3

代理解决了两个关键问题:

Figure 4

  1. 认证处理:HF Jobs的公开端点需要Authorization: Bearer <token>头,代理自动添加
  2. 多副本协调:vLLM的/load_lora_adapter端点只影响单个DP rank,而Jobs间无NCCL通信

Figure 5

代理运行在训练Job的127.0.0.1:8000,对TRL表现为单个vLLM服务器,实际执行两项操作:

Figure 6

  • 智能路由:将同一提示的多个rollout请求发送到已缓存其KV前缀的副本
  • 广播状态变更:适配器加载、暂停/恢复等操作同步到所有副本

Figure 7

KV前缀路由机制

vLLM以16-token为单位管理KV缓存块。对于同一提示的8个rollout请求,理想情况是全部路由到首个处理该提示的副本,后续7个请求可复用预填充结果。我们的路由策略包含以下步骤:

  1. 分块哈希:将提示按16-token分块,计算链式哈希(以适配器名作种子)
  2. 识别公共前缀:系统提示等所有请求共有的部分被标记为"common",不参与路由决策
  3. 副本选择
    • 若某副本持有该提示特有块且负载未超标(默认差值≤8请求),则路由过去(affinity hit)
    • 若目标副本过载,则发往负载最低副本(spill)
    • 全新提示发往负载最低副本(unmatched)

实际运行数据显示:84.5%的请求命中缓存(affinity),仅1.3%因负载不均溢出(spill),14.2%为全新提示(理论最小值12.5%)。

适配器广播可靠性

适配器加载采用"全有或全无"策略:

  1. 并行向所有副本发送加载请求
  2. 对"No adapter found"错误重试(等待存储桶同步)
  3. 任一副本失败则回滚所有已加载副本
# 伪代码:适配器加载广播
async def broadcast_load(adapter_path):
    results = await gather(*[load_to_replica(r, adapter_path) for r in replicas])
    if any_failed(results):
        await gather(*[unload_from_replica(r) for r in successful_replicas])
        raise Exception("Partial load rolled back")

Async GRPO with LoRA across Hugging Face Jobs. The trainer Job runs AsyncGRPOTrainer and the proxy, two vLLM Jobs serve the base model plus the latest adapter, and a Storage Bucket is mounted at /lora in all three.

实测126次同步(252次加载)全部成功,其中246次在第三次重试时成功——反映存储桶同步存在约2秒延迟。

性能优化历程

通过五轮实验逐步优化,500步训练时间从3小时27分钟降至53分钟:

初始瓶颈:训练器算力浪费

首版配置(r1-dp2)存在严重问题:

  • 每步耗时22.9秒,其中前向+反向占21.9秒
  • 微批次仅1序列(per_device_train_batch_size=1
  • GPU MFU仅3.9%
  • 推理队列持续满载(476/512),但第二副本几乎闲置

根本原因是小批次导致GPU计算单元大量空闲,而推理速度远超训练消耗能力。

优化1:Token-Budget批处理

启用token-budget批处理(token_budget=16384):

  • 每行打包12.7个样本(原为1个)
  • 微批次从64降至6
  • 前向+反向时间降至5.6秒
  • GPU MFU升至19%
  • 推理吞吐从4.6k tok/s跃升至25k tok/s

关键洞察:推理吞吐提升并非因修改vLLM,而是训练器不再阻塞推理队列。

优化2:关闭梯度检查点

默认开启的梯度检查点导致反向计算耗时异常(前向1.34秒 vs 反向4.26秒)。关闭后:

  • 前向+反向降至4.6秒(MFU达23%)
  • 推理队列从420降至60
  • 训练器开始等待推理结果(rollout_wait从0.02s升至0.5s)

此时瓶颈成功转移至推理端,同时暴露新问题:每4步的8.5秒权重同步占总耗时25%。

优化3:增加推理副本

添加第三推理副本并调整代理重试间隔(2s→0.5s):

  • 权重同步降至5.8秒
  • 但推理吞吐仅微增至26k tok/s

排查发现max_inflight_tasks=128限制了总并发量,三副本实际各处理约43请求——与双副本时单副本负载相当。

优化4:提升并发上限

将并发上限提至384(max_inflight_tasks=384):

  • 每副本稳定处理128请求
  • 推理队列回升至690/768
  • 训练重新成为瓶颈(每步4.8秒)
  • 最终500步耗时53分钟(提速3.9倍)

尽管平均样本延迟从1.5增至2.0版本,但仍低于max_staleness=4的阈值,奖励曲线保持稳定(最终reward 0.416 vs 原始0.438)。

实践建议

  1. 监控指标组合

    • perf/step_sperf/fwd_bwd_s → 训练瓶颈
    • rollout_queue_size满 + rollout/backpressure_s高 → 训练瓶颈
    • rollout_wait_s上升 + 队列空 → 推理瓶颈
  2. 批处理策略:小模型务必启用token-budget批处理,避免GPU利用率低下

  3. 梯度检查点:1.5B模型在H200上无需开启,保留激活值更高效

  4. 并发调优:初始可设max_inflight_tasks=128*num_replicas,根据队列水位调整

完整代码已在GitHub开源,通过环境变量即可复现各轮实验配置。