Qwen21_Text_Encoder_Heretic / extras /QwenImage21NVFP4Backend.py
catplusplus's picture
Upload folder using huggingface_hub
7b00440 verified
Raw History Blame Contribute Delete
7.18 kB
# -*- coding: utf-8 -*-
"""Nunchaku NVFP4 Resident Backend for Qwen-Image-2.1.
Loads the forged SVDQuant NVFP4 r32 Qwen-Image-2.1 transformer directly into VRAM,
enabling blazingly fast inference with zero layerwise PCIe streaming bottlenecks.
"""
import os
import sys
import torch
from diffusers import QwenImage21Pipeline
# Ensure local packages are on path
ROOT_DIR = "/auto/home/amano/olegk/Nikola"
for p in [f"{ROOT_DIR}/packages/nunchaku", f"{ROOT_DIR}/packages/deepcompressor", ROOT_DIR]:
if p not in sys.path:
sys.path.insert(0, p)
from nunchaku.models.transformers.transformer_qwenimage21 import NunchakuQwenImage21Transformer2DModel
class QwenImage21NVFP4Backend:
def __init__(
self,
model_id="/home/olegk/Nikola/models/Qwen/Qwen-Image-2.1",
optimized_model_path="/home/olegk/Nikola/models/nunchaku-qwen-image-2.1/best_quality_fp4.safetensors",
gpu_id=0,
enable_tiling=True,
dynamic_scale_k=0.0,
stream_text_encoder=True,
text_encoder_path=None,
):
self.model_id = model_id
self.optimized_model_path = optimized_model_path
self.gpu_id = gpu_id
self.enable_tiling = enable_tiling
self.dynamic_scale_k = dynamic_scale_k
self.stream_text_encoder = stream_text_encoder
self.text_encoder_path = text_encoder_path
self.pipeline = None
def load(self):
print(f"Loading QwenImage21NVFP4Backend...")
print(f" • Base Pipeline: {self.model_id}")
print(f" • Quantized DiT: {self.optimized_model_path}")
print(f" • Target GPU: cuda:{self.gpu_id}")
print(f" • Stream Text Encoder: {self.stream_text_encoder}")
if self.text_encoder_path:
print(f" • Custom Text Encoder: {self.text_encoder_path}")
device = f"cuda:{self.gpu_id}"
# 1. Load base pipeline without heavy transformer
# Pass transformer=None or dummy to save loading 13GB BF16 transformer
print("Loading peripheral pipeline components (Text Encoder, Tokenizer, VAE, Scheduler)...")
if self.text_encoder_path:
from transformers import Qwen3VLForConditionalGeneration, Qwen3VLProcessor
te_dir = (
os.path.join(self.text_encoder_path, "text_encoder")
if os.path.isdir(os.path.join(self.text_encoder_path, "text_encoder"))
else self.text_encoder_path
)
proc_dir = (
os.path.join(self.text_encoder_path, "processor")
if os.path.isdir(os.path.join(self.text_encoder_path, "processor"))
else self.text_encoder_path
)
if not os.path.exists(os.path.join(proc_dir, "tokenizer.json")):
proc_dir = os.path.join(self.model_id, "processor")
print(f" • Loading custom Qwen3-VL text encoder from: {te_dir}")
custom_te = Qwen3VLForConditionalGeneration.from_pretrained(
te_dir,
torch_dtype=torch.bfloat16,
low_cpu_mem_usage=True,
)
print(f" • Loading processor from: {proc_dir}")
custom_proc = Qwen3VLProcessor.from_pretrained(proc_dir)
pipeline = QwenImage21Pipeline.from_pretrained(
self.model_id,
transformer=None,
text_encoder=custom_te,
processor=custom_proc,
torch_dtype=torch.bfloat16,
)
else:
pipeline = QwenImage21Pipeline.from_pretrained(
self.model_id,
transformer=None,
torch_dtype=torch.bfloat16,
)
# 2. Load forged Nunchaku NVFP4 transformer directly into resident VRAM
print(f"Loading resident NVFP4 DiT into {device}...")
quantized_transformer = NunchakuQwenImage21Transformer2DModel.from_pretrained(
self.optimized_model_path,
device=device,
torch_dtype=torch.bfloat16,
)
pipeline.transformer = quantized_transformer
# 3. Place VAE directly on target device with memory-safe tiling
print(f"Placing VAE on {device}...")
pipeline.vae = pipeline.vae.to(device)
if self.enable_tiling:
try:
pipeline.vae.enable_tiling()
print("VAE tiling enabled successfully.")
except Exception as e:
print(f"Note: VAE tiling could not be enabled ({e})")
# 4. Text Encoder Configuration: Layerwise PCIe Streaming or CPU Fallback
if self.stream_text_encoder:
print(f"Configuring PCIe Layerwise Weight Streaming for Qwen3-VL on {device}...")
from stream_encoder import attach_qwen3vl_streamer
self.streamer = attach_qwen3vl_streamer(pipeline, device=device)
orig_encode_prompt = pipeline.encode_prompt
def streamed_safe_encode_prompt(*args, **kwargs):
kwargs.pop("device", None)
embeds = orig_encode_prompt(*args, device=torch.device(device), **kwargs)
target_dev = pipeline.transformer.device
return tuple(x.to(target_dev) if isinstance(x, torch.Tensor) else x for x in embeds)
pipeline.encode_prompt = streamed_safe_encode_prompt
else:
print(f"Configuring hybrid CPU prompt encoding (DiT & VAE 100% resident on {device})...")
orig_encode_prompt = pipeline.encode_prompt
def safe_encode_prompt(*args, **kwargs):
kwargs.pop("device", None)
te_device = pipeline.text_encoder.device
embeds = orig_encode_prompt(*args, device=te_device, **kwargs)
target_dev = pipeline.transformer.device
return tuple(x.to(target_dev) if isinstance(x, torch.Tensor) else x for x in embeds)
pipeline.encode_prompt = safe_encode_prompt
# 5. Automatically attach optimal dynamic latent variance damping s(t)
if self.dynamic_scale_k > 0.0:
k = self.dynamic_scale_k
print(f"Attaching automatic late-stage variance damping (k={k:.3f}, +1.20 dB fidelity boost)...")
orig_call = pipeline.__call__
def scaled_call(*args, **kwargs):
user_cb = kwargs.pop("callback_on_step_end", None)
def combined_cb(pipe_obj, step_idx, timestep, callback_kwargs):
t = float(timestep.item() if isinstance(timestep, torch.Tensor) else timestep)
if t < 600.0:
scale = 1.0 - k * ((600.0 - t) / 600.0)
callback_kwargs["latents"] = callback_kwargs["latents"] * scale
if user_cb is not None:
return user_cb(pipe_obj, step_idx, timestep, callback_kwargs)
return callback_kwargs
kwargs["callback_on_step_end"] = combined_cb
return orig_call(*args, **kwargs)
pipeline.__call__ = scaled_call
self.pipeline = pipeline
return self.pipeline, self.pipeline