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.
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.
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écnica | Qué resuelve |
|---|---|
dit_quant_scheme: nvfp4 | El 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_attn3 | SageAttention, 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.0 | Este 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: block | Té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
- 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.
- 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.
- 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.
- 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.