DiT 推理优化

DiT

DiT,在扩散模型领域一般指的是 Diffusion Transformer

  • 传统扩散模型(比如 Stable Diffusion)里的 U-Net 主干是 卷积网络 (CNN)
  • DiT 则是把主干替换成了 Transformer 架构,用自注意力机制来处理扩散过程中的特征。

为什么要用 Transformer?

  • 全局感受野:CNN 善于处理局部特征,而 Transformer 可以直接建模图像中任意两个位置的关系,更利于捕捉长程依赖。
  • 可扩展性:Transformer 在大规模训练(比如 GPT-类模型)里已经证明了很强的 scaling law,扩散模型也想利用这个优势。
  • 统一架构:有研究希望文本、图像、视频都用类似的 Transformer 主干,更方便多模态统一。

FBCache

原版满血Flux推理很慢,调研了下目前比较流行的denoising caching算法TeacacheWeaveSpeed ,主要原理

  • 计算首块残差:在每个推理步骤(timestep)中,正常计算第一个 Transformer 模块的输出。记录这个输出(或其与输入的残差)。
  • 比较残差变化:将当前步骤计算得到的首块残差与上一个步骤缓存的首块残差进行比较。
    • 如果两个步骤间的首块残差差异足够小(低于阈值),则认为当前步骤的输入与上一步骤足够相似,可以直接重用上一个步骤计算得到的最终输出(或最终残差),并跳过所有后续 Transformer 模块(从第二个模块开始)的计算。
    • 如果差异较大(高于阈值),则认为输入变化显著,需要完整执行所有 Transformer 模块的计算,得到新的最终输出,并更新缓存中的首块残差和最终输出(或最终残差),供下一个步骤使用。

直接用FLux底模 + diffusers 看,对比有无使用denoising caching的情况,发现确实比没有应用前有明显加速

参考

  • https://github.com/chengzeyi/Comfy-WaveSpeed

  • https://github.com/welltop-cn/ComfyUI-TeaCache

在测试调参时候肉眼可见有差距就淘汰了,业务代码不能直接使用diffusers进行推理,所以参考纯Flux的推理实现,主要是修改FLux的forward方法。实现后让业务对比纯flux推理结果(脱离后处理)。相同pipeline,prompt, seed等参数一样情况下,对比是否使用WaveSpeed生成图存在的差异。

ComfyUI 插件原理:根据模型名,通过unittest.mock.patch.object替代模型的forward_orig方法。

Diffusers

  • 对于FBCache + 适配最新版的diffusers PR fix FluxSingleTransformerBlock missing encoder_hidden_states when use Flux-Kontext

  • 使用compile,加速transformer

    1
    2
    3
    4
    def _compile(self):
        if hasattr(self.pipe, "transformer") and self.pipe.transformer is not None:
            # 仅编译模型中较小且频繁重复的模块(通常是转换层)来缩短冷启动延迟,并支持在后续每次执行时重用已编译的构件,将编译时间缩短
            self.pipe.transformer.compile_repeated_blocks(fullgraph=True, dynamic=True)
    

nunchaku

扩散模型 (Diffusion Models)生成图像效果特别好,但模型越来越大,需要更多的显存,推理速度也变慢。量化(quantization),就是把模型的权重和激活从 16/32 位浮点数压缩到更低的精度(比如 4 位),以减少内存和计算量。

权重 (weights)激活 (activations) 对量化非常敏感,传统方法(比如 LLM 常用的 smoothing,把异常值均摊到权重/激活里)不够用了,会导致生成质量下降。

SVDQuant:不用把所有东西都硬塞进 4 位,而是给“异常值”开一个低秩、高精度的“旁路(branch)”去专门处理

  • 把激活里的异常值转移到权重里。

  • 对权重里这些异常值,用 SVD(奇异值分解) 搞一个低秩分支(low-rank branch),保持高精度。

  • 这样就“吸收”了异常值,使主干部分更适合 4 位量化。

如果直接让低秩分支单独跑,会增加额外的内存访问(需要额外搬运数据),结果推理反而变慢,抵消了量化带来的加速。

所以使用Nunchaku 推理引擎,把低秩分支和低比特分支的运算融合起来

  • 避免重复的数据搬运。
  • 保持速度优势。
  • 还能直接兼容 LoRA,不用再重新量化

量化工具 deepcompressor

兼具压缩显存,快速推理 cpu offload + set attention + fbcache

1
2
3
4
5
6
7
8
9
10
11
12
13
from nunchaku import NunchakuFluxTransformer2dModel
from nunchaku.utils import get_precision
from nunchaku.caching.diffusers_adapters import apply_cache_on_pipe

transformer = NunchakuFluxTransformer2dModel.from_pretrained(
  f"{self.model_dir}/nunchaku-flux.1-kontext-dev/svdq-{get_precision()}_r32-flux.1-kontext-dev.safetensors",
  offload=True
)
transformer.set_attention_impl("nunchaku-fp16")
pipe = FluxKontextPipeline.from_pretrained(konext_model_dir, transformer=transformer,
                                           torch_dtype=torch.bfloat16)
apply_cache_on_pipe(pipe, residual_diff_threshold=0.06)
pipe.enable_model_cpu_offload()
  • 替换FluxSingleTransformerBlock的proj_out(block 最后一层线性层)为ConcatLinear(大Linear,输入向量会被 切分 (split) 成多个子块,然后每个子块单独过一个 nn.Linear,最后再把结果加起来)
  • 替换Unet