使用 TensorFlow 和 XLA 加速文本生成

TL;DR

Hugging Face 已在 TensorFlow 的 transformers 库中为文本生成启用 XLA(Accelerated Linear Algebra)编译。此优化可将生成速度提升至最高 100 倍,并且在许多基准测试中,文本生成任务的性能优于 PyTorch。

使用 XLA 加速 TensorFlow

XLA 是一种旨在加速 TensorFlow 模型的编译器,也是 JAX 和某些 PyTorch 实现的基础。在使用 Eager Execution 以提升透明度和调试体验的 TensorFlow 2 中,图模式的一些性能优势会丢失。为恢复这些优势,用户可以将函数包装在 tf.function 中,它会将代码转换为图。

通过在 tf.functiontf.keras.Model.compile 中添加 jit_compile=True 参数,用户可以触发 XLA 编译。虽然首次调用 XLA 编译的函数会因编译过程而较慢,但随后使用相同张量形状和类型的调用会显著加速。

XLA 文本生成的实现要求

XLA 依赖即时(JIT)编译和多态性。为避免在文本生成过程中产生代价高昂的重新编译(追踪),必须满足以下技术要求:

输入填充

由于 XLA 在遇到不同的张量形状、类型或非张量参数时会触发新的编译步骤,输入提示必须填充到统一长度。Hugging Face 建议在分词器类中使用 pad_to_multiple_of 参数,以在保持输入灵活性的同时限制可能的形状数量。

代码库向量化

自回归文本生成本质上是动态的,常常会扩展张量并使用动态切片,这对 XLA 并不友好。为实现 XLA 支持,Hugging Face 重写了 TensorFlow 文本生成代码库,使操作向量化并使用带填充的固定大小结构。此外,还对 NLP 模型进行了修改,以确保位置嵌入在这些填充结构下能够正常工作。

Transformers 中的文本生成能力

transformers 库中的 generate 函数支持多种解码策略:

  • 贪婪解码: 默认的确定性方法(do_sample=False),在每一步选择最可能的 token。
  • 采样: 一种随机方法(do_sample=True),可通过 temperature 参数控制随机性。较低的值倾向于高概率 token,而较高的值会增加熵。
  • 束搜索:num_beams 大于 1 时触发,此方法探索高概率序列,以提升相较于贪婪解码的输出质量。

性能基准

在多个 GPU 型号上比较 TensorFlow 与 PyTorch 的基准测试显示了两个主要结果:

  1. 巨大的加速: 在使用 XLA 时,TensorFlow 文本生成显著更快,某些情况下加速超过 100 倍。
  2. 框架对比: 在绝大多数情况下,使用 XLA 的 TensorFlow 是最快的选项,有时在文本生成任务上比 PyTorch 快至 9 倍。

Sources