Google 的 Tunix 系统解决了阻碍大规模智能体强化学习 (RL) 高效利用 TPU 的瓶颈问题。通过将生成交互数据的工作与更新策略的工作分离,Tunix 将 TPU 利用率从个位数提升到了接近满载的状态,大幅减少了计算资源的浪费。

智能体 RL 的瓶颈

智能体 RL 与更为常见的“下一 token”语言模型训练不同。智能体必须发送 API 调用、执行代码或在模拟环境中进行步进,然后对结果做出反应。因此,训练循环是同步的:模型产生一个动作,环境运行,结果返回,然后模型才能接收到梯度更新。当单个环境步进需要数秒时间时,昂贵的 TPU 硬件就会处于空闲状态,报告的利用率可能会降至 10% 以下。这种低效率直接转化为更高的云端账单和更慢的研究周期。

Tunix 的解耦架构

Tunix 通过将两个阶段——轨迹生成 (trajectory generation) 和策略优化 (policy optimization)——分配到独立的硬件池中来解决这个问题。

  • 异步执行器 (Asynchronous actors) 运行在廉价的 CPU 或 GPU 上。每个执行器持续与其分配的环境进行交互,记录动作和观测结果,并将生成的轨迹流式传输到共享存储中。
  • 持续学习器 (Continuous learners) 占用专用的 TPU Pod。学习器从中央缓冲区提取批次数据,并进行梯度更新,无需等待任何单个执行器完成 rollout。
  • 高吞吐量缓冲区 (High-throughput buffer) 位于中间,作为轨迹的中转区域。由于学习器的读取速度可以与缓冲区的供应速度相匹配,因此 TPU 永远不会停顿。

最终效果是,训练流水线中的 TPU 几乎始终保持忙碌,将利用率推向 100%。

技术障碍及 Tunix 的应对方案

变长回合与 XLA 重编译

JAX 的 XLA 编译器针对固定张量形状进行优化。然而,智能体任务会产生长度不同的序列,这通常会触发昂贵的重编译。Tunix 将较短的序列打包在一起,并将长度相似的回合分桶 (buckets) 处理,使形状保持足够稳定,以便 XLA 重用已编译的内核。其结果是实现了稳定的吞吐量,而不会产生可能破坏性能的编译器开销。

在多个 TPU 芯片上扩展大规模模型

训练参数超过 700 亿的模型需要将权重和数据分布到多个 TPU 节点上。Tunix 使用 JAX 的 ShardMap 原语对模型参数和激活值进行分片,使学习器能够在保持高速喂入数据的同时,将整个模型保留在内存中。这种分片策略使得训练以前单个 TPU Pod 无法触及的模型成为可能。

解耦流水线带来的过期梯度

当执行器运行领先于学习器时,它们提供的数据相对于当前策略可能会变得“过期” (stale)。Tunix 通过两种机制来缓解这种漂移:使用重要性采样 (importance-sampling) 对旧样本进行重新加权以反映其相关性,以及使用可配置的过期阈值丢弃超过预设时长的轨迹。两者共同作用,即使在流水线异步运行时也能保持学习的稳定性。

采用者需要注意的事项

  • 延迟审计 – 解耦带来的收益取决于环境的响应时间。团队应测量端到端延迟,并确保执行器池的规模足以保持缓冲区的充盈。
  • 工作池设计 – 廉价的 CPU 或 GPU 可以承载许多执行器,但过度订阅可能会导致网络或存储的竞争。建立一个与缓冲区摄取速率相匹配的平衡工作池至关重要。
  • 缓冲区鲁棒性 – 中央存储必须能够处理高写入和高读取速率,而不会成为新的瓶颈。选择具有低长尾延迟和足够带宽的存储系统是架构中不可逾越的要求。

潜在缺点

分离式架构引入了更多的变动环节:独立的硬件集群、持久化缓冲区以及用于执行过期限制的协调逻辑。

总结

Tunix 表明,智能体 RL 中的主要成本不在于模型本身,而在于同步交互循环导致的空闲时间。通过将 rollout 工作卸载到廉价硬件上,并利用高吞吐量缓冲区为持续学习的 TPU Pod 提供数据,Google 已将利用率低于 10% 的问题转变为接近满载的工作流。