# -*- 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