ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

FP8加速Stable Diffusion 3.5:显存减半,吞吐量提升40%的实战指南

FP8加速Stable Diffusion 3.5:显存减半,吞吐量提升40%的实战指南 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%这个数字不是拍脑袋来的。我的测试环境是单卡24GBSD3.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的某些层对精度比较敏感我建议保留FP16Transformer主干的注意力层和FFN层可以大胆转FP8VAE解码器最好也保留FP16因为解码阶段的数值范围比较宽FP8容易在暗部产生色带。转换脚本的核心逻辑是遍历state_dict对符合条件的权重做scale然后cast。scale的选取很关键太大会溢出太小会损失精度。我一般用权重的绝对值最大值除以FP8_E4M3的最大可表示值448然后取一个略小的安全系数。import torch def convert_to_fp8(weight, scale_factor0.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可以用E4M3V用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_ascale, scale_bscale, out_dtypetorch.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倍。这部分我做了混合精度第一个线性层用FP8GELU保持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加速收益不大但CFGClassifier-Free Guidance那一步需要把条件输出和无条件输出做加权和这部分如果精度不够会导致引导强度不稳定。我的做法是采样器的状态更新保持FP32CFG的加权和在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 Large1024×102428步CFG 4.5FP16和4.0FP8batch size分别取能跑满显存的最大值。硬件精度Batch Size单图延迟吞吐量显存占用4090FP1644.2s0.95 img/s21.3GB4090FP863.1s1.94 img/s14.2GB5090FP1663.5s1.71 img/s26.8GB5090FP8102.4s4.17 img/s18.5GBH100FP16122.8s4.29 img/s52.1GBH100FP8201.9s10.53 img/s34.7GB4090上的吞吐量提升大概是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.20.976几乎无风景36.70.968极轻微建筑35.10.959轻微文字渲染32.40.931可察觉人像场景下FP8和FP16的差异基本看不出来皮肤纹理和毛发细节都保留得很好。风景场景下天空渐变区域偶尔能看到极轻微的色带但需要放大到200%才能察觉。建筑场景下直线边缘的锐度略有下降但不影响整体观感。文字渲染是差异最明显的小字号的笔画边缘会有点糊但大字号没问题。这个结果比我预期的好。我原本以为FP8会在暗部或者高对比度区域翻车但实际上只要把VAE解码器和调制层保留FP16画质损失就控制在可接受范围内。4.3 什么情况下FP8会翻车有三种情况我建议你别用FP8。一是极低步数采样比如10步以下因为每个时间步的数值范围都很宽FP8的动态范围不够用画面容易发灰。二是高CFG scale比如7以上引导信号太强FP8的精度损失会被放大出现伪影。三是需要精细文字渲染的场景比如海报设计小字号的笔画会糊。另外如果你用的是SD3.5 Medium而不是LargeFP8的收益会小一些因为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或FP325.5 几个我踩过的坑第一个坑是忘了更新cuDNN。PyTorch的FP8 kernel依赖cuDNN的某些算子如果cuDNN版本太老会静默回退到FP16你以为在跑FP8其实没有。建议用torch.backends.cudnn.version()确认一下。第二个坑是混合精度训练和推理的scale不通用。训练时的scale是根据梯度动态调整的推理时用同样的scale会偏大或者偏小。推理的scale应该根据权重的实际分布单独算。第三个坑是忽略了VAE。我一开始只转了TransformerVAE还是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的第一层线性层是收益最大且风险最低的优先转这两块基本就能拿到大部分加速收益。
返回列表