Hugging Face PyTorch / XLA TPU 集成
Hugging Face 已将 PyTorch / XLA 集成,使用户能够在保持标准 Hugging Face Trainer 接口的同时,在 Cloud TPU 上训练和扩展 transformer 模型。此集成利用 PyTorch / XLA 库将 PyTorch 框架与 XLA(加速线性代数)设备(包括 Cloud TPU)连接起来。
PyTorch / XLA 技术实现
该集成将 xla 设备类型引入 PyTorch,使得张量可以在 TPU 硬件上创建和管理。Hugging Face Trainer 模块利用 TrainingArguments 数据类自动检测并在 is_torch_tpu_available() 为 true 时返回 TPU 设备。
梯度合并和优化器步骤
由于 Cloud TPU 设备通常由多个核心组成(例如,单个设备可能有 8 个核心),梯度必须在数据并行副本之间进行交换。集成使用 xm.optimizer_step(optimizer) 来处理梯度合并和随后的优化器步骤,确保 TPU 核心之间的同步。
输入流水线
为了防止宿主 CPU 和 TPU 加速器在相互等待时空闲,PyTorch / XLA 实现了一个输入流水线。通过使用 pl.MpDeviceLoader,系统可以在步骤 $n$ 仍在执行时重叠追踪步骤 $n+1$,从而优化模型的数据馈送。
检查点管理
为了确保可移植性并避免特定于设备的加载问题,张量在检查点之前会被移动到 CPU。使用 xm.save() API 确保只有一个进程(主序数)写入存储位置,防止在多进程环境中出现文件损坏。
PyTorch / XLA 的工作原理
惰性张量执行
与 CPU 和 CUDA 张量不同,后者会急切地执行操作,XLA 张量是惰性的。它们会在结果需要之前将操作记录在图中。这种延迟执行使得 XLA 编译器能够将多个独立操作融合为一个优化后的操作。
Trace-Compile-Execute 循环
PyTorch / XLA 遵循特定的执行流程以优化 TPU 性能:
- 追踪:随着前向和反向传递的运行,中间表示(IR)图会实时追踪。
- 截断:当调用
xm.mark_step()(通常通过MpDeviceLoader间接调用)时,活动图会被截断。 - 编译:IR 图被降低为 XLA 高级操作(HLO),编译为 TPU 二进制文件并执行。
- 缓存:为了避免重新编译的高昂成本,编译后的 TPU 二进制文件会存储在一个缓存中,缓存键为 HLO 图的唯一哈希。
为了最大化缓存命中率并最小化编译开销,建议保持张量形状静态。Hugging Face 模型通常通过适当填充输入令牌来保持静态形状。
性能基准
在使用 v3-8 Cloud TPU 系统(4 个 TPU v3 芯片)的 WikiText103 数据集上训练 bert-large-uncased 得到以下结果:
| 名称 | 全局批量大小 | 精度 | 训练时间(分钟) |
|---|---|---|---|
| bert-large-uncased | 64 | FP32 | 178.4 |
| bert-large-uncased | 128 | BF16 | 106.4 |
这些基准是在使用 n1-standard-96 CPU 配置的情况下进行的,以确保工作负载不受宿主 CPU 限制。