在 Cloud TPU v5e 上使用 JAX 加速 Stable Diffusion XL 推理
Hugging Face 将 JAX 支持集成到 Diffusers 库中,以在 Cloud TPU v5e 上实现 Stable Diffusion XL (SDXL) 的高性能、成本效益推理。此集成通过利用 JAX 的即时 (JIT) 编译和 XLA 驱动的并行性,解决了 SDXL 的计算挑战,其 UNet 大约是前身的三倍大小。
通过 JAX 和 TPU v5e 进行的技术优化
在 Cloud TPU v5e 上服务 SDXL 通过两种主要的软硬件机制实现高效率:JIT 编译和 SPMD 并行性。
静态形状的 JIT 编译
JAX 使用即时 (JIT) 编译在初始执行期间追踪代码,并为后续调用生成优化的 TPU 二进制文件。此过程需要静态的输入、中间和输出形状。SDXL 与 JIT 编译高度兼容,原因如下:
- 恒定的输出形状: 图像生成通常使用固定数量的图像和一致的尺寸。
- 固定形状嵌入: Stable Diffusion 和 SDXL 使用固定形状的嵌入向量(带填充)来处理文本提示。
虽然初始编译需要几分钟(在提供的示例中大约为三分钟),但后续的推理调用将显著加速。
XLA 并行性和吞吐量
JAX 的 pmap 使得单程多数据(SPMD)执行成为可能,允许在多个 XLA 设备之间扩展工作负载。这使得图像生成可以线性扩展:例如,拥有 8 个芯片的 TPU 在单个芯片生成一张图像所需的时间内可以生成 8 张图像。Cloud TPU v5e 实例提供多种配置(从 1 到 256 个芯片),通过超快速 ICI 链路相连,使用户能够根据特定的吞吐量需求进行扩展。
在 JAX 中的实现流水线
使用 JAX 运行 SDXL 推理涉及一种功能方法,其中模型参数与管道分开处理。关键实现步骤包括:
- 模型加载: 使用
FlaxStableDiffusionXLPipeline.from_pretrained加载基础 SDXL 1.0 模型。 - 精度管理: 将模型参数转换为
bfloat16以减少内存使用并提高速度,同时将调度器状态保持在float32以防止导致低质量或黑色图像的精度错误。 - 输入准备: 使用
prepare_inputs确保提示在各次调用中具有一致的维度,这是 JIT 编译所必需的。 - 设备复制: 在可用的 TPU 芯片之间复制参数和输入(例如,对 TPU v5e-4 使用
replicate),并为每个芯片分配唯一的随机种子,以确保图像输出的多样性。 - 执行: 使用
jit=True调用管道以触发 XLA 编译过程。
性能基准测试
在使用 Euler Discrete 调度器进行 20 步的 SDXL 1.0 基础模型基准测试表明,TPU v5e 在成本效益方面优于 TPU v4。
| 硬件 | 批量大小 | 延迟 | 每美元性能 |
|---|---|---|---|
| TPU v5e-4 (JAX) | 4 | 2.33s | 21.46 |
| TPU v5e-4 (JAX) | 8 | 4.99s | 20.04 |
| TPU v4-8 (JAX) | 4 | 2.16s | 9.05 |
| TPU v4-8 (JAX) | 8 | 4.17s | 8.98 |
TPU v5e 相比 TPU v4 能够实现最高 2.4 倍的每美元性能。性能通过计算吞吐量(每芯片的批量大小除以延迟)并将该数值除以硬件的列表价格来衡量。
部署架构
当前实现使用一个负载均衡服务器,将用户请求随机路由到运行在预分配的 Cloud TPU v5e-4 实例上的后端服务器。每个实例大约在 4 秒内生成四张 1024×1024 的图像(包括前端处理和通信),实际生成时间大约为 2.3 秒。