Stable Diffusion JAX 与 Flax 集成
Hugging Face 已在 diffusers 库中从版本 0.5.1 开始集成了 Flax 支持,使 Stable Diffusion 能够在 Google TPU 上高效运行。此集成让用户能够利用 TPU 服务器的并行处理能力——通常配备八个加速器——在生成一张图像的时间内同时生成多张图像。
使用 JAX 与 Flax 的高速 TPU 推理
通过使用 JAX 和 Flax,对 TPU 上的 Stable Diffusion 推理进行了优化,与标准 GPU 实现相比可实现显著加速。在 TPU v2-8 上,首次编译后的后续推理大约耗时 7 秒。
关键技术优化包括:
- bfloat16 精度:TPU 设备支持
bfloat16,这是一种高效的半精度浮点类型,能够在保持性能的同时降低内存开销。 - JIT 编译:通过向 Flax pipeline 传入
jit=True,JAX 将模型编译为高效的表示。首次运行需要进行编译(在 TPU v2-8 上超过一分钟),但所有后续调用都显著更快。 - 无状态模型:由于 Flax 是函数式框架,模型是无状态的,参数存储在模型之外。
通过 SPMD 实现并行化
diffusers 的 Flax pipeline 使用单程序多数据(Single-Program, Multiple-Data,SPMD)并行化来最大化 TPU 硬件利用率。这主要通过 jax.pmap 函数实现。
并行化的实现方式
jax.pmap 执行两个关键功能:编译代码(类似于 jax.jit())并确保编译后的代码在所有可用设备上并行运行。
为了实现并行执行,pipeline 按以下步骤进行:
- 复制:使用
flax.jax_utils.replicate将模型参数复制到所有设备上。 - 分片:使用
shard对输入数据(例如标记化的提示 ID)进行分片。例如,若有 8 台设备,提示数组会被拆分,使每台设备收到输入的特定部分。 - PRNG 处理:为确保生成图像的可复现性和多样性,创建随机数生成器(RNG)并将其拆分为多个生成器——每个设备一个。
该架构使得 pipeline 能够同时生成八张不同的图像(或同一图像的八个副本),因为每个设备独立处理一个批次项。
模型获取与许可
Flax 版的 Stable Diffusion 权重可在 Hugging Face Hub 的 CompVis/stable-diffusion-v1-4 仓库中获取。访问需要接受 CreativeML OpenRAIL-M 许可证,其中包括以下条款:
- 用户不得使用该模型故意生成或分享非法或有害内容。
- 用户保留对其生成输出的权利,并对其使用负责。
- 允许商业使用和权重再分发,前提是向所有用户共享相同的使用限制以及 CreativeML OpenRAIL-M 许可证的副本。