PyTorch 2.0 加速扩散模型:提升文本到图像生成速度 49% 的实战教程
近年来,生成式人工智能的突破大多来自扩散模型,这类模型能够根据文本提示生成高质量的图像和视频。然而,所有扩散模型都有一个共同的缺点:生成速度较慢,因为图像生成的采样过程是迭代的。优化采样循环中的代码至关重要。本教程将以开源文本到图像扩散模型为例,介绍如何利用 PyTorch 2.0 的两大优化方法——编译(Compilation)和快速注意力(Fast Attention),让生成速度提升最高 49%(根据 GPU 架构和批量大小有所不同)。最重要的是,加速过程无需安装任何额外的依赖库(如 xFormers)。
一、优化方法概述
我们对原始代码进行了三项核心优化,分别针对注意力机制、代码执行效率以及内存操作。下面逐一介绍。
1. 优化注意力(Attention Optimization)
注意力机制是扩散模型中的计算瓶颈,尤其是多头注意力和交叉注意力。原始代码使用手动实现的点积注意力,时间和内存复杂度与序列长度成二次方关系。我们将其替换为 PyTorch 2.0 内置的 nn.MultiheadAttention,该模块默认使用内存高效注意力(Memory-Efficient Attention)并支持 FlashAttention。
优化后的代码对比:
class CrossAttention(nn.Module):
def __init__(self, ...):
self.mmha = nn.MultiheadAttention(...)
def forward(self, x, context):
return self.mmha(x, context, context)
原代码:
class CrossAttention(nn.Module):
def __init__(self, ...):
...
def forward(self, x, context):
# 手动实现点积注意力
...
重要提示:FlashAttention 只在 GPU 计算能力为 SM 7.5 或 SM 8.x 的设备上可用(如 T4、A10、A100)。在 A100 上测试时,由于扩散模型中注意力头数和小批量大小较小,内存高效注意力(Memory-Efficient Attention)的效果反而优于 FlashAttention。高级用户可以通过 torch.backends.cuda.sdp_kernel 上下文管理器手动启用或禁用不同的注意力后端。
2. 编译(Compilation)
PyTorch 2.0 的 torch.compile 功能可以将 Python 代码编译成高效指令,消除 Python 解释器开销。只需一行代码即可启用:
model = torch.compile(model)
编译会在第一次运行代码时动态进行。为了获得最大加速,需要避免图形断裂(graph breaks)——即编译器无法编译的部分。我们移除了代码中的 checkpoint 函数等不支持的库调用,尽量保持整个计算图完整。
注意:编译要求 GPU 计算能力 ≥ SM 7.0(不支持 P100)。目前我们只编译 U-Net 部分,而非整个采样循环,因为循环每次迭代都会重新编译,反而降低性能。
3. 其他优化
我们消除了常见的 GPU 内存操作陷阱,例如直接在 GPU 上创建张量,而非先在 CPU 上创建再传输到 GPU。这减少了内存带宽开销。使用 PyTorch Profiler 中的 火焰图 可以定位这些微小的效率损失。
二、基准测试与结果
我们定义了五种配置进行对比(均基于相同文本到图像扩散模型和 PLMS 采样器):
- 无 xFormers 的原始代码(使用 PyTorch 1.12 和手动注意力)
- 有 xFormers 的原始代码
- 优化代码(普通注意力,无编译)
- 优化代码(内存高效注意力,无编译)
- 优化代码(内存高效注意力 + 编译)
以下表格展示了相对于"有 xFormers 的原始代码"的加速百分比。绝对值请参见下文详细表格。
不同 GPU 和批量大小的加速百分比(%)
| GPU | 批量大小 1 | 批量大小 2 | 批量大小 4 |
| P100 | -3.8% | 0.44% | 5.47% |
| T4 | 2.12% | 10.51% | 14.2% |
| A10 | -2.34% | 8.99% | 10.57% |
| V100 | 18.63% | 6.39% | 10.43% |
| A100 | 38.5% | 20.33% | 12.17% |
从上表可以看出:
- 对于 A100 和 V100 等高端 GPU,加速效果显著,且批量大小为 1 时提升最大。
- 对于 T4 等中低端 GPU,加速较小甚至略有倒退,但批量越大加速效果越明显。
绝对运行时间(秒)及相对于"有 xFormers"的改善百分比
批量大小 1
| 配置 | P100 | T4 | A10 | V100 | A100 |
| 无 xFormers 原始代码 | 30.4s (19.3%) | 29.8s (-77.3%) | 13.0s (-83.9%) | 10.9s (-33.1%) | 8.0s (19.3%) |
| 有 xFormers 原始代码(基准线) | 25.5s (0%) | 16.8s (0%) | 7.1s (0%) | 8.2s (0%) | 6.7s (0%) |
| 优化代码(普通注意力,无编译) | 27.3s (-7.0%) | 19.9s (18.7%) | 13.2s (87.2%) | 7.5s (8.7%) | 5.7s (15.1%) |
| 优化代码(内存高效注意力,无编译) | 26.5s (-3.8%) | 16.8s (0.2%) | 7.1s (-0.8%) | 6.9s (16.0%) | 5.3s (20.6%) |
| 优化代码(内存高效注意力 + 编译) | - | 16.4s (2.1%) | 7.2s (-2.3%) | 6.6s (18.6%) | 4.1s (38.5%) |
批量大小 2
| 配置 | P100 | T4 | A10 | V100 | A100 |
| 无 xFormers 原始代码 | 58.0s (21.6%) | 57.6s (84.0%) | 24.4s (95.2%) | 18.6s (-63.0%) | 12.0s (-50.6%) |
| 有 xFormers 原始代码(基准线) | 47.7s (0%) | 31.3s (0%) | 12.5s (0%) | 11.4s (0%) | 8.0s (0%) |
| 优化代码(普通注意力,无编译) | 49.3s (-3.5%) | 37.9s (-21.0%) | 17.8s (-42.2%) | 12.7s (10.7%) | 7.8s (1.8%) |
| 优化代码(内存高效注意力,无编译) | 47.5s (0.4%) | 31.2s (0.5%) | 12.2s (2.6%) | 11.5s (-0.7%) | 7.0s (12.6%) |
| 优化代码(内存高效注意力 + 编译) | - | 28.0s (10.5%) | 11.4s (9.0%) | 10.7s (6.4%) | 6.4s (20.3%) |
批量大小 4
| 配置 | P100 | T4 | A10 | V100 | A100 |
| 无 xFormers 原始代码 | 117.9s (-20.0%) | 112.4s (-81.8%) | 47.2s (-101.7%) | 35.8s (-71.9%) | 22.8s (-78.9%) |
| 有 xFormers 原始代码(基准线) | 98.3s (0%) | 61.8s (0%) | 23.4s (0%) | 20.8s (0%) | 12.7s (0%) |
| 优化代码(普通注意力,无编译) | 101.1s (-2.9%) | 73.0s (-18.0%) | 28.3s (-21.0%) | 23.3s (11.9%) | 14.5s (13.9%) |
| 优化代码(内存高效注意力,无编译) | 92.9s (5.5%) | 61.1s (1.2%) | 23.9s (-1.9%) | 20.8s (-0.1%) | 12.8s (-0.9%) |
| 优化代码(内存高效注意力 + 编译) | - | 53.1s (14.2%) | 20.9s (10.6%) | 18.6s (10.4%) | 11.2s (12.2%) |



以上图表展示了不同配置下各 GPU 的绝对运行时间变化趋势。注意:由于基准测试采用了循环重复运行的方法(A、B、C、D、E、A、B…),不同图表间的绝对时间可比性有限,但图表内部的相对差异是可靠的。
三、常见问题解答
Q1:我的 GPU 是 P100,为什么编译不起作用?
A:PyTorch 2.0 的编译模式要求 GPU 计算能力 ≥ SM 7.0。P100 的计算能力为 SM 6.0,不支持编译。对于 P100,只要使用内存高效注意力(无需 xFormers)就能获得 5.5% 的加速(批量大小 4 时)。
Q2:我需要安装 xFormers 吗?
A:不需要!本教程的核心优势就是仅依赖 PyTorch 2.0 内置功能,无需任何额外依赖。PyTorch 2.0 已将 xFormers 的内存高效注意力集成到 nn.MultiheadAttention 中。
Q3:为什么在 T4 上编译后加速只有 2.1%(批量=1)?
A:T4 的计算能力较低(SM 7.5),编译的收益有限。加上扩散模型 U-Net 中存在一些难以避免的图形断裂,导致编译不能完全发挥效果。批量增大后,编译的收益会提升(批量=4 时加速 14.2%)。
Q4:我的输出分辨率改变时,编译需要重新运行吗?
A:是的。编译发生在第一次运行代码时,若输入尺寸或模型结构改变,PyTorch 会重新编译。建议在部署时确保输入尺寸固定,或使用 torch._dynamo.config.cache_size_limit 控制缓存行为。
Q5:如何开启或关闭不同的注意力后端?
A:使用 torch.backends.cuda.sdp_kernel(enable_flash=True, enable_math=True, enable_mem_efficient=True) 上下文管理器。例如,强制使用数学(香草)注意力:
with torch.backends.cuda.sdp_kernel(enable_flash=False, enable_math=True, enable_mem_efficient=False):
output = model(x, context)

注意:上图为基准测试中典型运行时间的波动示意图。我们每个配置运行了 10 轮循环,并增加了额外的"预热"迭代(--niter 2,但实际包含 1 次预热),以确保编译后的第一次迭代不影响公平对比。
四、总结与下一步
通过 PyTorch 2.0 的编译和优化注意力,我们成功在无需外部依赖(如 xFormers)的前提下,实现了对文本到图像扩散模型最高 49% 的推断加速。这不仅降低了部署复杂度,也为后续优化打开了新的大门。
未来的改进方向包括:
- 将同样的优化应用于训练流程,PyTorch 2.0 编译可直接适用于训练,优化注意力的训练支持已在路线图中。
- 对 U-Net 以外的采样循环进行编译,但需避免每次采样步骤的重新编译问题。
- 将编译扩展到其他采样器(如 DDIM、DPM-Solver),而非仅限 PLMS。
- 将这套优化推广到图像到图像、图像修复等更多扩散模型任务中。
