Qwen-Image-2.1-NVFP4 / text_encoder /modeling_nunchaku_qwen3vl.py
joseplcam's picture
Inline kernel setup: Diffusers downloads only each component's own module
c10a8f8 verified
Raw History Blame Contribute Delete
3.49 kB
"""Qwen3-VL text encoder with SVDQuant NVFP4 (Nunchaku Lite) linears.
Loaded by Diffusers as a custom pipeline component (`trust_remote_code=True`).
`config.json` holds the compact Nunchaku Lite config under `nunchaku_lite`; the
listed linears are swapped for Diffusers' own `SVDQW4A4Linear` before the
quantized state dict is loaded, so inference uses the same NVFP4 kernels as the
transformer. Everything else (vision tower, embeddings, the BF16 layers) loads
as in the stock model.
"""
import os
from pathlib import Path
import torch
from accelerate import init_empty_weights
from huggingface_hub import snapshot_download
# --- Kernel setup. Kept identical in transformer/modeling_nunchaku_qwenimage21.py: Diffusers downloads
# only each component's own module file, and whichever component loads first must run it. ---
# Diffusers fetches the Nunchaku Lite kernels from `rootonchair/nunchaku-lite-kernels`, which is
# no longer downloadable. Loading this repo with trust_remote_code=True already runs this code,
# so point that kernel name at the rebuilt copy before Diffusers imports its Nunchaku utilities.
KERNELS = "rootonchair/nunchaku-lite-kernels"
if KERNELS not in os.environ.get("LOCAL_KERNELS", ""):
local = f"{KERNELS}={snapshot_download('joseplcam/nunchaku-lite-kernels')}"
os.environ["LOCAL_KERNELS"] = ":".join(filter(None, [os.environ.get("LOCAL_KERNELS"), local]))
if "DIFFUSERS_TRUST_REMOTE_KERNELS" not in os.environ:
# Diffusers read this variable into a constant at import time, so set both.
import diffusers.utils.constants
os.environ["DIFFUSERS_TRUST_REMOTE_KERNELS"] = "true"
diffusers.utils.constants.DIFFUSERS_TRUST_REMOTE_KERNELS = True
# --- end of kernel setup ---
# isort: off -- must stay below the kernel setup, which prepares the kernels this import loads
from diffusers.quantizers.nunchaku.utils import replace_with_nunchaku_linear # noqa: E402
from safetensors.torch import load_file # noqa: E402
from transformers import Qwen3VLForConditionalGeneration # noqa: E402
# isort: on
class NunchakuQwen3VLForConditionalGeneration(Qwen3VLForConditionalGeneration):
@classmethod
def from_pretrained(cls, pretrained_model_name_or_path, *args, subfolder: str = "", **kwargs):
path = Path(pretrained_model_name_or_path) / subfolder
if not path.is_dir():
path = Path(snapshot_download(str(pretrained_model_name_or_path), allow_patterns=[f"{subfolder}/*"])) / subfolder
dtype = kwargs.get("dtype") or kwargs.get("torch_dtype") or torch.bfloat16
if dtype == "auto":
dtype = torch.bfloat16
config = cls.config_class.from_pretrained(path)
with init_empty_weights(): # parameters on meta, buffers (e.g. rotary tables) materialized
model = cls(config)
replace_with_nunchaku_linear(model, config.nunchaku_lite, dtype)
# strict=False only because non-persistent buffers are absent; both checks below stay strict
result = model.load_state_dict(load_file(path / "model.safetensors"), strict=False, assign=True)
if result.unexpected_keys:
raise ValueError(f"Checkpoint has {len(result.unexpected_keys)} unexpected tensors, e.g. {result.unexpected_keys[:3]}")
missing = [name for name, p in model.named_parameters() if p.device.type == "meta"]
if missing:
raise ValueError(f"Checkpoint is missing {len(missing)} parameters, e.g. {missing[:3]}")
return model.eval()