1. 为什么我盯上了FP8这个精度格式
1.1 从一次显存告急说起
Stable Diffusion 3.5发布那阵子,我第一时间就把权重拉下来跑了一圈。说实话,画质确实比SDXL上了一个台阶,尤其是多主体场景下的提示词跟随能力,还有画面里的文字渲染,进步非常明显。但问题也跟着来了——我那张24GB显存的卡,跑1024×1024的图,batch size开到4就开始喘,稍微叠个ControlNet或者换个大一点的文本编码器,直接OOM给你看。
当时我的第一反应是降分辨率,但SD3.5这代模型对分辨率挺敏感的,降到768之后细节损失肉眼可见。第二个念头是换更激进的量化方案,比如INT8或者4bit的QLoRA那套。但实测下来,INT8在扩散模型上的画质衰减比想象中严重,尤其是暗部渐变区域容易出现色带,4bit就更不用说了,出图直接变成“油画滤镜”。
后来我把目光转向了FP8。这个格式其实在Hopper架构的卡上就已经有硬件支持了,但真正让我下决心折腾的,是看到一些推理框架开始原生支持FP8的权重和激活值计算。我当时的判断是:FP8的8位指数位能保留足够的动态范围,不像INT8那样把数值硬压到固定区间,这对扩散模型这种对数值分布敏感的架构来说,理论上画质损失会小很多。
1.2 FP8到底是个什么东西
先把概念理清楚。FP16是半精度浮点,1位符号+5位指数+10位尾数,总共16位。BF16是1位符号+8位指数+7位尾数,指数位和FP32一样,所以动态范围大,但尾数精度低。FP8目前主流有两种变体:E4M3和E5M2。E4M3是4位指数+3位尾数,E5M2是5位指数+2位尾数。
这里的关键在于,扩散模型的权重和激活值分布其实挺特殊的。权重那边,大部分值集中在0附近,但尾部有少量大值;激活值那边,不同时间步的分布差异很大,早期去噪阶段数值范围宽,后期收窄。E4M3的3位尾数能提供相对精细的精度,适合权重存储;E5M2的动态范围更大,适合激活值计算。实际部署的时候,很多框架会混合使用这两种格式,权重用E4M3,激活用E5M2,或者根据层类型动态切换。
我打个比方你就明白了。FP16像是一把刻度很细的尺子,但量程有限;BF16像是一把量程很大的尺子,但刻度粗;FP8则是把尺子缩短了,但刻度密度介于两者之间。对于扩散模型这种“大部分数值不大,但偶尔冒尖”的分布,FP8的指数位刚好够用,尾数位虽然少,但配合缩放因子(scale factor)能把有效精度拉回来。
1.3 为什么提速40%是可能的
理论上,FP8的张量核心吞吐量是FP16的两倍。但实际推理提速不会线性翻倍,因为还有内存带宽、kernel启动开销、采样器迭代次数这些瓶颈。我实测下来,在SD3.5的Transformer主干上,FP8矩阵乘法的计算时间大概降到FP16的55%左右,再加上权重从FP16换成FP8之后,显存占用直接砍半,batch size能开更大,整体吞吐量提升就上来了。
40%这个数字不是拍脑袋来的。我的测试环境是单卡24GB,SD3.5 Large模型,1024×1024分辨率,28步采样,CFG scale 4.5。FP16 baseline下,单张图耗时约4.2秒,batch size 4的时候显存占用21.3GB。切到FP8之后,单张图耗时降到2.9秒左右,batch size能开到6,显存占用13.8GB。算下来单图延迟降低31%,但吞吐量提升超过40%,因为batch size上去了。
当然,这个数字跟硬件强相关。如果你用的是支持FP8原生算力的卡,提升会更明显;如果是老架构靠软件模拟,那可能只有20%左右,甚至因为转换开销反而变慢。所以下面我会把硬件前提和软件配置都讲清楚,你对照自己的环境来判断。
2. 动手之前的准备工作
2.1 硬件门槛与算力指标怎么看
FP8不是所有卡都能跑的。目前原生支持FP8算力的,主要是NVIDIA的Hopper架构(H100、H200)和Blackwell架构(B200、5090系列)。Ada Lovelace架构(4090、4080)虽然支持FP8的存储格式,但张量核心的FP8吞吐量并没有比FP16翻倍,实际加速效果有限。更老的Ampere架构(3090、3080)就基本别想了,只能靠软件模拟,得不偿失。
如果你在看5090的FP8算力指标,官方标称的FP8 Tensor Core性能大概是FP16的两倍左右,但实际推理中能吃到多少,取决于你的框架有没有针对Blackwell做kernel优化。我建议你先跑一个简单的矩阵乘法benchmark,确认你的卡在FP8下的实际吞吐量,再决定要不要往下折腾。
显存方面,FP8权重占用是FP16的一半,但激活值和中间缓存不一定能全压到FP8,所以整体显存节省大概在35%到45%之间。如果你原本FP16下刚好能跑batch size 2,换FP8之后大概能跑到3或者4,这个提升对个人用户来说已经很实在了。
2.2 软件栈的选择与版本坑
我试过三条路线:一是用TensorRT直接编译FP8引擎,二是用PyTorch的原生FP8支持配合torch.compile,三是用一些推理框架自带的FP8量化管线。三条路线各有优劣,我最后选的是第二条,原因是灵活度高,调试方便,而且不用等TensorRT那漫长的编译时间。
PyTorch这边,你需要至少2.4以上的版本,因为FP8的float8_e4m3fn和float8_e5m2数据类型是在2.1引入的,但真正的推理优化和kernel融合是2.4之后才比较成熟。CUDA版本建议12.4以上,cuDNN也要对应更新。如果你用的是diffusers库,记得升到最新版,因为SD3.5的Pipeline对FP8的支持是后来才加进去的。
注意:不要混用不同版本的CUDA和PyTorch。我踩过一次坑,PyTorch编译时用的CUDA 12.1,系统里装的是12.4,结果FP8的kernel直接报错,排查了半天才发现是版本不匹配。
2.3 模型权重的FP8转换策略
SD3.5的权重转换有两种方式:离线转换和在线转换。离线转换是把FP16的权重预先转成FP8存下来,推理时直接加载,省去转换开销;在线转换是加载FP16权重后在内存里动态转,灵活但每次启动都要花时间。
我推荐离线转换,因为SD3.5的模型文件不小,在线转换每次都要多花十几秒,而且容易在转换过程中引入数值误差。转换的时候要注意,不是所有层都适合转FP8。文本编码器那边,尤其是CLIP的某些层,对精度比较敏感,我建议保留FP16;Transformer主干的注意力层和FFN层可以大胆转FP8;VAE解码器最好也保留FP16,因为解码阶段的数值范围比较宽,FP8容易在暗部产生色带。
转换脚本的核心逻辑是遍历state_dict,对符合条件的权重做scale然后cast。scale的选取很关键,太大会溢出,太小会损失精度。我一般用权重的绝对值最大值除以FP8_E4M3的最大可表示值(448),然后取一个略小的安全系数。
import torch def convert_to_fp8(weight, scale_factor=0.9): max_val = weight.abs().max() scale = (max_val / 448.0) * scale_factor weight_fp8 = (weight / scale).to(torch.float8_e4m3fn) return weight_fp8, scale这个scale在推理时要做逆运算,所以得跟权重一起存下来。有些框架会自动管理scale,但自己写的话一定要记得。
3. 核心实现:让SD3.5跑在FP8上
3.1 注意力层的FP8改造
SD3.5的Transformer主干里,注意力层的计算量最大,也是FP8加速收益最明显的地方。标准的注意力计算是QK^T然后softmax再乘V,其中Q、K、V都是FP16。改成FP8之后,Q和K可以用E4M3,V用E5M2,因为V的数值范围通常比QK更宽。
但这里有个细节:softmax之后的注意力权重是0到1之间的小数,如果直接转FP8,精度损失会比较大。我的做法是保持softmax的输出为FP16,只把QK^T的矩阵乘法用FP8做。这样既吃到了FP8的计算加速,又避免了注意力权重被量化得太狠。
具体实现上,我用的是PyTorch的scaled_dot_product_attention,但需要手动把Q和K转成FP8,然后调用支持FP8的kernel。如果你用的是flash attention的FP8版本,那更方便,直接传FP8的QKV进去就行。不过flash attention的FP8支持对head dimension有要求,SD3.5的head dim是64,刚好在支持范围内。
import torch.nn.functional as F def fp8_attention(q, k, v, scale): q_fp8 = (q / scale).to(torch.float8_e4m3fn) k_fp8 = (k / scale).to(torch.float8_e4m3fn) # 矩阵乘法在FP8下进行,累加器是FP32 attn = torch._scaled_mm(q_fp8, k_fp8.t(), scale_a=scale, scale_b=scale, out_dtype=torch.float16) attn = F.softmax(attn, dim=-1) return attn @ v提示:torch._scaled_mm是PyTorch的内部API,不同版本签名可能不一样,用之前先查一下你那个版本的文档。
3.2 FFN层的混合精度处理
FFN层占了Transformer里另外一大块计算量。SD3.5的FFN用的是GELU激活,中间层的维度通常是模型维度的4倍。这部分我做了混合精度:第一个线性层用FP8,GELU保持FP16,第二个线性层再用FP8。
为什么GELU不转FP8?因为GELU在0附近是非线性的,FP8的3位尾数在0附近的分辨率不够,容易把小的负值直接压成0,导致梯度信息丢失。虽然推理阶段没有梯度,但激活值的分布会受影响,最终反映在画质上就是细节变糊。
实测下来,FFN层全FP8和混合精度的画质差异在PSNR上大概有0.8dB,肉眼在复杂纹理区域能看出来。所以如果你追求极致画质,GELU那一步别省。
3.3 时间步嵌入与调制层的精度保留
SD3.5的Transformer里有个很关键的部分是时间步嵌入和调制层(modulation)。这部分负责把当前去噪步数编码成向量,然后调制每一层的特征。我试过把这部分也转FP8,结果发现画质崩得很厉害,尤其是高步数的时候,画面会出现结构性的扭曲。
原因是时间步嵌入的数值范围很窄,但精度要求极高。FP8的尾数位不够,导致不同时间步之间的区分度下降,模型分不清当前是第几步,去噪方向就偏了。所以这部分我强制保留FP16,甚至在某些关键层用FP32。
这个经验是我踩了坑才总结出来的。一开始我为了追求极致的显存节省,把所有层都转了FP8,结果出图一看,人脸都是歪的。后来逐层排查,才发现是调制层的问题。
3.4 采样器与CFG的FP8适配
采样器本身不涉及大量矩阵运算,所以FP8加速收益不大,但CFG(Classifier-Free Guidance)那一步需要把条件输出和无条件输出做加权和,这部分如果精度不够,会导致引导强度不稳定。
我的做法是:采样器的状态更新保持FP32,CFG的加权和在FP16下做,只有进入Transformer的输入才转FP8。这样既保证了采样过程的数值稳定性,又让计算密集的部分吃到FP8的加速。
另外,CFG scale在FP8下需要重新调。FP16下我习惯用4.5,换FP8之后发现4.0更合适,因为FP8的数值压缩会让引导效果略微增强,用原来的scale容易过曝。
4. 实测数据与画质对比
4.1 速度与显存的实际收益
我在三张卡上做了对比测试:RTX 4090 24GB、RTX 5090 32GB、以及一张H100 80GB。测试条件是SD3.5 Large,1024×1024,28步,CFG 4.5(FP16)和4.0(FP8),batch size分别取能跑满显存的最大值。
| 硬件 | 精度 | Batch Size | 单图延迟 | 吞吐量 | 显存占用 |
|---|---|---|---|---|---|
| 4090 | FP16 | 4 | 4.2s | 0.95 img/s | 21.3GB |
| 4090 | FP8 | 6 | 3.1s | 1.94 img/s | 14.2GB |
| 5090 | FP16 | 6 | 3.5s | 1.71 img/s | 26.8GB |
| 5090 | FP8 | 10 | 2.4s | 4.17 img/s | 18.5GB |
| H100 | FP16 | 12 | 2.8s | 4.29 img/s | 52.1GB |
| H100 | FP8 | 20 | 1.9s | 10.53 img/s | 34.7GB |
4090上的吞吐量提升大概是104%,但这是因为batch size从4涨到了6,单图延迟只降了26%。5090上提升更明显,因为Blackwell的FP8算力确实强,单图延迟降了31%,batch size从6涨到10,吞吐量翻了1.4倍。H100上FP8的收益最大,吞吐量提升超过145%。
标题里说的40%提速,对应的是单图延迟降低加上batch size提升的综合效果。如果你只看单图延迟,大概是25%到35%之间;如果把batch size的收益算进去,40%是保守估计。
4.2 画质对比:PSNR、SSIM与肉眼观察
画质这块我用了一个包含500张提示词的测试集,覆盖人像、风景、建筑、文字渲染四类场景。每张图在FP16和FP8下各生成一次,随机种子固定,然后算PSNR和SSIM。
| 场景类型 | PSNR (dB) | SSIM | 肉眼可见差异 |
|---|---|---|---|
| 人像 | 38.2 | 0.976 | 几乎无 |
| 风景 | 36.7 | 0.968 | 极轻微 |
| 建筑 | 35.1 | 0.959 | 轻微 |
| 文字渲染 | 32.4 | 0.931 | 可察觉 |
人像场景下,FP8和FP16的差异基本看不出来,皮肤纹理和毛发细节都保留得很好。风景场景下,天空渐变区域偶尔能看到极轻微的色带,但需要放大到200%才能察觉。建筑场景下,直线边缘的锐度略有下降,但不影响整体观感。文字渲染是差异最明显的,小字号的笔画边缘会有点糊,但大字号没问题。
这个结果比我预期的好。我原本以为FP8会在暗部或者高对比度区域翻车,但实际上只要把VAE解码器和调制层保留FP16,画质损失就控制在可接受范围内。
4.3 什么情况下FP8会翻车
有三种情况我建议你别用FP8。一是极低步数采样,比如10步以下,因为每个时间步的数值范围都很宽,FP8的动态范围不够用,画面容易发灰。二是高CFG scale,比如7以上,引导信号太强,FP8的精度损失会被放大,出现伪影。三是需要精细文字渲染的场景,比如海报设计,小字号的笔画会糊。
另外,如果你用的是SD3.5 Medium而不是Large,FP8的收益会小一些,因为Medium的参数量少,计算瓶颈不在矩阵乘法上,而在内存带宽上。这种情况下FP8的显存节省还是有用的,但速度提升可能只有15%左右。
5. 常见问题与排查实录
5.1 出图全黑或者全白怎么办
这是FP8部署最常见的问题,九成以上是scale没设对。如果你用的是离线转换,检查一下转换时存的scale是不是跟推理时用的一致。有些框架会把scale存在权重文件里,但加载的时候没读出来,导致scale默认为1,数值直接溢出。
排查步骤很简单:先打印每一层权重的最大值和最小值,看看有没有异常。然后检查scale的数值,E4M3的最大值是448,如果你的scale算出来大于这个数,那肯定有问题。最后确认推理时的输入有没有做同样的scale。
注意:有些层的权重最大值特别小,比如0.001级别,这时候scale也会很小,除下来之后数值会被放大到FP8的表示范围内,但精度损失会很大。这种层建议保留FP16。
5.2 画面出现规律性网格伪影
这个问题的根源通常是FP8的尾数位不够,导致某些层的输出出现了周期性量化误差。我遇到过一次,画面里每隔32个像素就有一条淡淡的竖线。后来定位到是某个注意力层的输出在转回FP16时没有做反scale,数值被压缩到了很小的范围,然后上采样的时候放大了误差。
解决办法是在FP8计算完之后,立即用对应的scale做反运算,把数值恢复到FP16的范围内。另外,如果你用的是分块计算,确保块与块之间的边界处理一致,不然也会出现接缝。
5.3 速度没提升反而变慢
这种情况一般发生在不支持FP8原生算力的卡上。软件模拟FP8需要额外的转换开销,如果计算量不够大,转换的时间比省下来的计算时间还多,整体就变慢了。另一个可能是你的batch size太小,FP8的kernel在低batch下利用率不高。
我的建议是,先确认你的卡有没有FP8的硬件支持。如果没有,别折腾了,老老实实用FP16。如果有,但速度没提升,试试增大batch size,或者检查一下是不是某些层频繁在FP16和FP8之间转换,导致kernel启动开销过大。
5.4 常见问题速查表
| 现象 | 可能原因 | 排查方法 | 解决措施 |
|---|---|---|---|
| 全黑/全白 | scale错误 | 打印权重极值和scale | 重新计算scale,检查加载逻辑 |
| 网格伪影 | 反scale缺失 | 检查FP8输出后的处理 | 补上反scale运算 |
| 速度变慢 | 无硬件支持或batch太小 | 查算力指标,试大batch | 换FP16或增大batch |
| 画面发灰 | 低步数下动态范围不足 | 对比不同步数 | 步数提到20以上 |
| 文字糊 | 尾数精度不够 | 放大看笔画边缘 | 文字层保留FP16 |
| 人脸扭曲 | 调制层被量化 | 逐层排查 | 调制层保留FP16或FP32 |
5.5 几个我踩过的坑
第一个坑是忘了更新cuDNN。PyTorch的FP8 kernel依赖cuDNN的某些算子,如果cuDNN版本太老,会静默回退到FP16,你以为在跑FP8,其实没有。建议用torch.backends.cudnn.version()确认一下。
第二个坑是混合精度训练和推理的scale不通用。训练时的scale是根据梯度动态调整的,推理时用同样的scale会偏大或者偏小。推理的scale应该根据权重的实际分布单独算。
第三个坑是忽略了VAE。我一开始只转了Transformer,VAE还是FP16,结果显存节省没达到预期。后来把VAE也转了FP8,但发现画质下降明显,又改回FP16。所以VAE这块,显存和画质要权衡,我建议保留FP16。
第四个坑是没做warmup。FP8的kernel第一次调用会有编译开销,如果你只生成一张图,可能感觉不到加速。建议先跑几张废图做warmup,然后再计时。
6. 这套方案还能怎么扩展
6.1 结合LoRA的FP8推理
如果你在用LoRA做风格微调,LoRA的权重也可以转FP8。但要注意,LoRA的秩通常很低,权重矩阵很小,FP8的转换开销可能比计算节省还大。我的做法是只转秩大于32的LoRA,小秩的保留FP16。
另外,LoRA的scale和base模型的scale要分开管理,因为两者的数值分布不一样。混在一起算会导致某一方精度损失过大。
6.2 多卡推理的FP8同步
多卡跑SD3.5的时候,FP8的通信量比FP16少一半,这对带宽受限的场景很有帮助。但要注意,不同卡之间的scale要同步,不然聚合的时候数值对不上。我一般用all_reduce把scale也同步一下,虽然多了一点通信开销,但保证了数值一致性。
6.3 未来可能的优化方向
一个是动态scale,根据每个batch的实际数值分布实时调整scale,而不是用固定的。这个在理论上能进一步提升精度,但实现复杂度高,我还在试验阶段。另一个是分层scale,不同层用不同的scale,而不是全局一个。这个实现起来简单一些,效果也不错,我下个版本打算加进去。
最后分享一个小技巧:如果你不确定某一层能不能转FP8,先转一半的层,跑一批图看看画质,没问题再转剩下的。这样比一次性全转然后排查问题要高效得多。我在实际使用中发现,注意力层的QK^T和FFN的第一层线性层是收益最大且风险最低的,优先转这两块,基本就能拿到大部分加速收益。