Curso / Aplicado a tus proyectos / NVFP4 en difusión
🔧 basado en tu código real

Aplicado a tus proyectos · Flux2 + LightX2V

NVFP4 en difusión: la misma cuantización, un motor completamente distinto

Ya viste NVFP4 en un decoder autoregresivo (Gemma-4). Acá es el mismo formato de 4 bits, integrado a mano en dos arquitecturas de difusión — una que genera una imagen, otra que genera video — con problemas que nunca aparecen en un LLM de texto.

Flux2: diffusion transformer texto→imagen (BFL) Wan 2.1: image-to-video, destilado a 4 pasos Hardware: RTX 5090 (Blackwell, sm_120)

0. Misma técnica, arquitectura completamente distinta

En la página de Gemma-4, cuantizar a NVFP4 fue casi un one-liner: un flag (--qformat nvfp4_mse) que le pasás a una herramienta de NVIDIA, y vLLM se encarga del resto en runtime (--quantization modelopt_fp4). Acá no hay vLLM ni ModelOpt de por medio — es torch puro, invocando el kernel de matmul cuantizado de Blackwell directamente. Ver esta versión "a mano" es lo que te muestra qué es lo que vLLM te estaba escondiendo cuando usás el flag fácil.

1. Un diffusion transformer, en una frase

Un LLM decoder-only genera una vez: calcula la probabilidad del siguiente token y listo, un forward pass por token. Un modelo de difusión genera iterando: arranca de ruido puro y, en varias pasadas (típicamente 20-50, a veces mucho menos — sección 4), predice cuánto ruido hay que restarle a la imagen actual para acercarla un poco más al resultado final. Flux2 usa bloques transformer para esa predicción — de ahí "diffusion transformer" — pero no hay nada autoregresivo ni causal en juego: no genera token por token, refina la imagen entera en cada paso.

Para lo que importa en esta página, el detalle que sí es igual a un LLM: por dentro, la mayoría del cómputo son las mismas operaciones — proyecciones lineales grandes (attn.to_q/k/v, capas de feed-forward) — y son exactamente esas matrices las que se cuantizan a NVFP4.

2. NVFP4 sin vLLM: torch._scaled_mm a mano

blackwell_utils.py define BlackwellLinear, un reemplazo a medida para nn.Linear que hace exactamente lo que ModelOpt+vLLM hacen por vos en la página anterior, pero visible línea por línea:

# Pesos: (out, in // 2) uint8 — dos valores FP4 empaquetados por byte
self.register_buffer("qweight", torch.empty((out_features, in_features // 2), dtype=torch.uint8))
# Escala: un float8_e4m3fn cada 16 valores — mismo esquema de bloque que Gemma-4
self.register_buffer("weight_scale", torch.empty((out_features, in_features // 16), dtype=torch.float8_e4m3fn))

En el forward, la activación se cuantiza dinámicamente a FP4 (a diferencia de los pesos, que ya vienen cuantizados desde el checkpoint), y el matmul cuantizado corre directo en hardware:

res = torch._scaled_mm(
    qact.view(torch.float4_e2m1fn_x2),
    self.qweight.view(torch.float4_e2m1fn_x2).t(),
    scale_a=act_scale,
    scale_b=self.weight_scale,
    out_dtype=torch.bfloat16,
)

float4_e2m1fn_x2 es exactamente el E2M1 de la página de Gemma-4 (1 signo, 2 exponente, 1 mantisa), con el "x2" del nombre del dtype indicando que cada elemento de PyTorch empaqueta dos valores FP4 en un byte — la razón por la que el tensor de pesos declara la mitad de columnas (in_features // 2). torch._scaled_mm es la función que sabe desempaquetar eso y multiplicar en el tensor core de Blackwell sin pasar por una conversión intermedia a BF16.

🔧 Un problema que no existe en un LLM de texto

El checkpoint NVFP4 de Flux2 lo publica Black Forest Labs con su propia convención de nombres (double_blocks.N.img_attn.qkv, heredada de Flux 1), pero el modelo que corre en memoria usa la convención de diffusers (transformer_blocks.N.attn.to_q). Ningún framework resuelve ese mapeo por vos. _nvfp4_key_for_module() es, literalmente, una tabla de traducción escrita a mano entre los dos esquemas de nombres — incluyendo el caso de to_q/to_k/to_v separados en diffusers mapeando a una única matriz qkv fusionada del lado de BFL, con un índice de slice para saber qué tercio de esa matriz corresponde a cada uno.

3. ¿Cómo fine-tuneás algo que ya está en 4 bits?

Con un modelo en BF16, aplicar un LoRA es casi trivial: sumás la matriz de bajo rango entrenada (A·B) directamente a los pesos originales, o la dejás separada y la sumás en runtime. Con pesos empaquetados en NVFP4 eso deja de ser posible — no hay forma de "sumarle" una corrección continua a un valor que ya está discretizado a uno de 16 niveles por bloque.

benchmark_lora_nunchaku.py lo deja explícito en el código: el método merge() de la capa LoRA-sobre-NVFP4 directamente levanta NotImplementedError("Merging into NVFP4 relies on fusion at runtime."). La única salida es fusionar el cómputo del LoRA dentro del kernel en cada forward — nunca modificar los pesos cuantizados en sí. nunchaku (la librería de la que viene svdq_quantize_w4a4_act_fuse_lora_cuda) hace justamente eso: cuantiza la activación y aplica la corrección de LoRA en la misma pasada CUDA.

Cuantizar no es solo comprimir el modelo — es cerrarle una puerta a cualquier técnica que asuma que podés sumarle un número real a un peso.

4. Video: lo mismo, con una dimensión más

Tu configuración de producción para Wan 2.1 (image-to-video) en LightX2V agrega, sobre la misma base de NVFP4, un conjunto de técnicas que no tienen equivalente en un LLM de texto:

TécnicaQué resuelve
dit_quant_scheme: nvfp4El transformer de difusión (DiT) va a 4 bits. El text encoder (T5) y CLIP se quedan en BF16 — mismo patrón que embeddings/lm_head en Gemma-4: las capas que traducen entre modalidades se salvan de la cuantización.
self_attn/cross_attn: sage_attn3SageAttention, el kernel de atención optimizado para Blackwell — el mismo tipo de rol que Flash Attention o PagedAttention cumplen para LLMs de texto.
infer_steps: 4, guidance_scale: 1.0Este checkpoint está destilado: entrenado para necesitar solo 4 pasos de denoising en vez de 20-50. La destilación también elimina la necesidad de classifier-free guidance (correr el modelo dos veces por paso, con y sin condicionamiento, para reforzar la adherencia al prompt) — guidance_scale=1.0 apaga esa segunda pasada, la mitad del cómputo por paso.
cpu_offload, offload_granularity: blockTécnica de memoria distinta a la cuantización: los bloques del transformer que no están calculando activamente se mandan a RAM del sistema y vuelven a la GPU cuando les toca — la forma de hacer entrar un modelo de video en 24 GB sin tocar su precisión.
RIFE (interpolación de frames)Un paso de post-proceso separado del modelo de difusión: en vez de generar más frames con el DiT (caro), un modelo chico de interpolación estima los frames intermedios — 16 fps generados se convierten en 60 fps de salida.

🟢 El principio que se repite en las tres páginas de NVFP4

Gemma-4 excluye embeddings/lm_head/proyecciones MTP de la cuantización. Flux2 mapea a mano el checkpoint cuantizado a la arquitectura de diffusers. Wan 2.1 deja el text encoder y CLIP en BF16. Ningún caso cuantiza el modelo entero de forma uniforme — en los tres, el trabajo real está en decidir qué no cuantizar, y verificar que las piezas que sí se cuantizaron sigan hablando el mismo idioma que las que no.

5. Resumen

  1. NVFP4 es un formato de pesos — no le importa si el modelo es un decoder autoregresivo o un diffusion transformer; el mismo E2M1 en bloques de 16 aplica a los dos.
  2. Integrarlo sin un framework que lo resuelva por vos expone lo que ModelOpt+vLLM esconden: empaquetado de bits, matmul cuantizado explícito, y el problema de mapear convenciones de nombres entre checkpoints de distinto origen.
  3. Fine-tunear con LoRA sobre pesos en 4 bits no se puede hacer por fusión de pesos — se fusiona el cómputo en el kernel, en cada forward.
  4. Video agrega memoria (offload por bloque), velocidad de sampling (destilación, sin CFG) y un truco de post-proceso (interpolación de frames) que no tienen contraparte en texto.