File size: 3,490 Bytes
7cbf234
 
 
 
 
 
 
 
 
 
c10a8f8
7cbf234
 
 
 
 
 
c10a8f8
 
 
 
 
 
 
 
 
 
 
 
7cbf234
c10a8f8
 
 
 
 
7cbf234
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
"""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()