Modelos de Difusão
Os modelos de difusão para geração de imagens e vídeos têm ganhado popularidade, entregando mídia visual super-realista. No entanto, sua adoção é frequentemente limitada pelos requisitos de memória e processamento. A quantização é essencial para o serviço eficiente desses modelos. Neste artigo, demonstramos acelerações de inferência reprodutíveis de até 1,26x com MXFP8 e 1,68x com NVFP4 com diffusers e torchao no Flux.1-Dev, QwenImage e LTX-2 modelos no NVIDIA B200. Também delineamos como usamos quantização seletiva, CUDA Graphs e LPIPS como medida para iterar sobre a precisão e o desempenho óptimo desses modelos. O código para reproduzir os experimentos neste artigo está disponível.
Os formatos de microescala MXFP8 e NVFP4 são suportados nativamente pela arquitetura Blackwell da NVIDIA (por exemplo, GPUs B200). Ao contrário da quantização padrão, que dimensiona um tensor inteiro, a microescala agrupa elementos em blocos pequenos (por exemplo, 16 ou 32 valores) que compartilham um fator de escala de alta precisão. Isso permite bit-depths significativamente menores, preservando a faixa dinâmica e a precisão. Para saber mais sobre isso, você pode consultar este artigo. O NVFP4 requer uma capacidade CUDA de pelo menos 10,0. Portanto, certifique-se de que você tem uma GPU que atenda a esse requisito. Os benchmarks apresentados neste documento foram realizados em uma máquina B200 (B200 DGX). Para o ambiente virtual, você pode usar conda.
No momento da escrita, as versões noturnas eram 2.12.0.dev20260315+cu130, 0.17.0.dev20260316+cu130 e 2026.3.15+cu130 para PyTorch, TorchAO e MSLK, respectivamente. Alguns modelos requerem que os usuários sejam autenticados na plataforma Hugging Face Hub. Portanto, certifique-se de executar hf auth login antes de executar os exemplos, se não tiver feito isso anteriormente. Usar a configuração de quantização NVFP4 do TorchAO é direto com sua integração nativa em Diffusers. O snippet de código acima quantiza cada camada torch.nn.Linear do modelo. Para este artigo, sempre usamos compilação regional com fullgraph=True, pois isso reduz significativamente o tempo de compilação e produz resultados quase tão bons quanto a compilação do modelo completo.
Saiba mais sobre compilação regional a partir daqui. O snippet de código abaixo mostra como configurar a inferência MXFP8 e NVFP4 com TorchAO. Os seguintes parâmetros de inferência foram usados durante a benchmarking FLUX.1-dev. Primeiro, apresentamos a latência e o pico de consumo de memória em diferentes configurações e benchmarks, com acelerações de até 1,26x com MXFP8 e até 1,59x com NVFP4. Observe que esses resultados usam quantização seletiva, na qual excluímos certas camadas da quantização. Discutimos mais sobre quantização seletiva mais adiante neste artigo. NVIDIA B200, quantização seletiva, torch.compile com compilação regional; batch_size=1 usa torch.compile(…, mode=’reduce-overhead’). O modo de quantização “Nenhum” significa sem quantização.
As imagens MXFP8 e NVFP4 geradas para um prompt de teste estão próximas da linha de base bfloat16. Para uma avaliação de precisão mais aprofundada, calculamos a pontuação LPIPS média entre as imagens bfloat16 (linha de base) e as imagens MXFP8|NVFP4 (experimento), média sobre os prompts no conjunto de dados Drawbench. NVIDIA B200, quantização seletiva, torch.compile com compilação regional. Uma pontuação LPIPS de zero significa “imagens idênticas” e pontuações LPIPS mais baixas correspondem a maior semelhança perceptual. O código que usamos para calcular a pontuação LPIPS média está aqui. Por favor, consulte a seção LPIPS mais adiante neste artigo para mais detalhes sobre as avaliações de precisão com LPIPS.
Para LTX-2, habilitamos a divisão em blocos no VAE para manter os requisitos de memória administráveis. Os seguintes parâmetros de inferência foram usados para obter os resultados. NVIDIA B200, quantização seletiva, torch.compile com compilação regional. O modo de quantização “Nenhum” significa sem quantização. Confira este link para uma comparação dos resultados de vídeo em um prompt de teste. Calcular pontuações de avaliação sobre um conjunto de dados de prompts (como fizemos para Flux-1.dev) é deixado para um estudo futuro. Os seguintes parâmetros de inferência foram usados para obter os resultados. NVIDIA B200, quantização seletiva, torch.compile com compilação regional, batch_size=1 usa torch.compile(…, mode=’reduce-overhead’). O modo de quantização “Nenhum” significa sem quantização.
As imagens MXFP8 e NVFP4 geradas para um prompt de teste estão próximas da linha de base bfloat16, com NVFP4 mostrando diferenças ligeiramente maiores em comparação com MXFP8. Na tabela a seguir, relatamos as pontuações LPIPS semelhantes às do Flux.1-Dev. Observe que, em nossos experimentos, encontramos que QwenImage é mais sensível à quantização do que Flux.1-Dev, como evidenciado pela pontuação LPIPS média MXFP8 de 0,34 para QwenImage (em comparação com uma pontuação LPIPS média de 0,11 para MXP8 no Flux-1.Dev). Reduzir a pontuação LPIPS média para QwenImage ainda mais por meio de quantização seletiva mais agressiva ou algoritmos numéricos mais avançados (GPTQ, QAT, etc.) é deixado para um estudo futuro.
Nesta seção, compartilhamos como usamos quantização seletiva, CUDA Graphs e LPIPS para iterar sobre