Qwen-Image-2.1-NVFP4 / runtime /qwen21-nvfp4-conditioning.patch
BennyDaBall's picture
Add tested native workflows, comparisons, encoder patch and validation
2e77d03 verified
Raw History Blame
1.85 kB
diff --git a/comfy/text_encoders/qwen_image21.py b/comfy/text_encoders/qwen_image21.py
index de0d147d..8ee21129 100644
--- a/comfy/text_encoders/qwen_image21.py
+++ b/comfy/text_encoders/qwen_image21.py
@@ -1,7 +1,10 @@
+import json
import numbers
import torch
+import comfy.ops
+import comfy.model_management
import comfy.text_encoders.qwen3vl
from comfy import sd1_clip
@@ -33,9 +36,25 @@ class QwenImage21Qwen3VLClipModel(comfy.text_encoders.qwen3vl.Qwen3VLClipModel):
# last layer without the final RMSNorm: transformers 4.57 hidden_states[-1], which Qwen's results are tuned to (5.x norms it)
self.layer_norm_hidden_state = False
self.image_spans = []
+ self.nvfp4_conditioning = False
+
+ def load_sd(self, sd):
+ self.nvfp4_conditioning = any(json.loads(v.numpy().tobytes()).get("format") == "nvfp4" for k, v in sd.items() if k.endswith(".comfy_quant"))
+ return super().load_sd(sd)
+
+ def forward(self, tokens):
+ device = self.execution_device
+ if device is None:
+ device = self.transformer.get_input_embeddings().weight.device
+ if not self.nvfp4_conditioning or not comfy.model_management.supports_nvfp4_compute(device):
+ return super().forward(tokens)
+ with comfy.ops.use_quantized_matmul(self.transformer, device):
+ return super().forward(tokens)
def process_tokens(self, tokens, device):
embeds, attention_mask, num_tokens, embeds_info = super().process_tokens(tokens, device)
+ if self.nvfp4_conditioning and comfy.model_management.supports_nvfp4_compute(device):
+ embeds = embeds.to(dtype=torch.bfloat16)
self.image_spans = [(e["index"], e["size"]) for e in embeds_info if e["type"] == "image"]
return embeds, attention_mask, num_tokens, embeds_info