mirror of
https://github.com/lllyasviel/stable-diffusion-webui-forge.git
synced 2026-07-21 21:01:24 +08:00
Nunchaku Z-Image
This commit is contained in:
parent
2ff53f45c8
commit
a44f84830c
@ -260,9 +260,14 @@ def load_huggingface_component(guess, component_name, lib_name, cls_name, repo_p
|
||||
|
||||
model_loader = lambda c: QwenImageTransformer2DModel(**c)
|
||||
elif cls_name in ("Lumina2Transformer2DModel", "ZImageTransformer2DModel"):
|
||||
from backend.nn.lumina import NextDiT
|
||||
if guess.nunchaku:
|
||||
from backend.nn.svdq import NunchakuZImageModel
|
||||
|
||||
model_loader = lambda c: NextDiT(**c)
|
||||
model_loader = lambda c: NunchakuZImageModel(**c)
|
||||
else:
|
||||
from backend.nn.lumina import NextDiT
|
||||
|
||||
model_loader = lambda c: NextDiT(**c)
|
||||
|
||||
unet_config = guess.unet_config.copy()
|
||||
state_dict_parameters = memory_management.state_dict_parameters(state_dict)
|
||||
@ -598,7 +603,7 @@ def split_state_dict(sd, additional_state_dicts: list = None):
|
||||
sd = preprocess_state_dict(sd)
|
||||
guess = huggingface_guess.guess(sd)
|
||||
|
||||
if "Qwen" in guess.huggingface_repo and getattr(guess, "nunchaku", False):
|
||||
if getattr(guess, "nunchaku", False) and ("Z-Image" in guess.huggingface_repo or "Qwen" in guess.huggingface_repo):
|
||||
import json
|
||||
|
||||
from nunchaku.utils import get_precision_from_quantization_config
|
||||
|
||||
@ -8,6 +8,7 @@ import types
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from nunchaku import NunchakuFluxTransformer2dModel, NunchakuT5EncoderModel
|
||||
from nunchaku.caching.diffusers_adapters.flux import apply_cache_on_transformer
|
||||
@ -42,7 +43,7 @@ class NunchakuModelMixin(nn.Module):
|
||||
return super().to(*args, **kwargs)
|
||||
|
||||
|
||||
# ========== Flux ========== #
|
||||
# region: Flux
|
||||
|
||||
|
||||
class SVDQFluxTransformer2DModel(nn.Module):
|
||||
@ -225,7 +226,8 @@ class SVDQT5(torch.nn.Module):
|
||||
self.logit_scale = torch.nn.Parameter(torch.tensor(4.6055))
|
||||
|
||||
|
||||
# ========== Qwen ========== #
|
||||
# region: Qwen
|
||||
|
||||
|
||||
from backend.memory_management import xformers_enabled
|
||||
|
||||
@ -773,3 +775,141 @@ class NunchakuQwenImageTransformer2DModel(NunchakuModelMixin, QwenImageTransform
|
||||
m.wtscale = sd.pop(f"{n}.wtscale", 1.0)
|
||||
|
||||
return super().load_state_dict(sd, *args, **kwargs)
|
||||
|
||||
|
||||
# region: Z-Image
|
||||
|
||||
from backend.nn.lumina import JointAttention, NextDiT, clamp_fp16
|
||||
|
||||
|
||||
def fuse_to_svdquant_linear(linear1: nn.Linear, linear2: nn.Linear, **kwargs) -> SVDQW4A4Linear:
|
||||
assert linear1.in_features == linear2.in_features
|
||||
assert linear1.bias is None and linear2.bias is None
|
||||
return SVDQW4A4Linear(
|
||||
linear1.in_features,
|
||||
linear1.out_features + linear2.out_features,
|
||||
bias=False,
|
||||
torch_dtype=linear1.weight.dtype,
|
||||
device=linear1.weight.device,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
class NunchakuZImageAttention(JointAttention):
|
||||
|
||||
def __init__(self, orig_attn: JointAttention, **kwargs):
|
||||
nn.Module.__init__(self)
|
||||
self.n_kv_heads = orig_attn.n_kv_heads
|
||||
self.n_local_heads = orig_attn.n_local_heads
|
||||
self.n_local_kv_heads = orig_attn.n_local_kv_heads
|
||||
self.n_rep = orig_attn.n_rep
|
||||
self.head_dim = orig_attn.head_dim
|
||||
|
||||
self.qkv = SVDQW4A4Linear.from_linear(orig_attn.qkv, **kwargs)
|
||||
self.out = SVDQW4A4Linear.from_linear(orig_attn.out, **kwargs)
|
||||
|
||||
self.q_norm = orig_attn.q_norm
|
||||
self.k_norm = orig_attn.k_norm
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
x_mask: torch.Tensor,
|
||||
freqs_cis: torch.Tensor,
|
||||
transformer_options={},
|
||||
) -> torch.Tensor:
|
||||
return super().forward(x, x_mask, freqs_cis, transformer_options)
|
||||
|
||||
|
||||
class NunchakuZImageFeedForward(nn.Module):
|
||||
|
||||
def __init__(self, orig_ff: FeedForward, **kwargs):
|
||||
super().__init__()
|
||||
self.w13 = fuse_to_svdquant_linear(orig_ff.w1, orig_ff.w3, **kwargs)
|
||||
self.w2 = SVDQW4A4Linear.from_linear(orig_ff.w2, **kwargs)
|
||||
|
||||
def _forward_silu_gating(self, x1, x3):
|
||||
return clamp_fp16(F.silu(x1) * x3)
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
x = self.w13(x)
|
||||
x3, x1 = x.chunk(2, dim=-1)
|
||||
return self.w2(self._forward_silu_gating(x1, x3))
|
||||
|
||||
|
||||
def patch_z_image_state_dict(state_dict: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
|
||||
patched_state_dict = {}
|
||||
quant_sub_keys = ["wscales", "wcscales", "wtscale", "smooth_factor_orig", "smooth_factor", "proj_down", "proj_up"]
|
||||
|
||||
for key, value in state_dict.items():
|
||||
if "attention.to_qkv" in key:
|
||||
patched_state_dict[key.replace("to_qkv", "qkv")] = value
|
||||
elif "attention.to_q" in key:
|
||||
q_weight = state_dict[key]
|
||||
k_weight = state_dict[key.replace("to_q", "to_k")]
|
||||
v_weight = state_dict[key.replace("to_q", "to_v")]
|
||||
patched_state_dict[key.replace("to_q", "qkv")] = torch.cat([q_weight, k_weight, v_weight], dim=0)
|
||||
elif "attention.to_k" in key or "attention.to_v" in key:
|
||||
continue
|
||||
elif "attention.to_out" in key:
|
||||
patched_state_dict[key.replace("to_out.0", "out")] = value
|
||||
elif "feed_forward.net.0.proj.qweight" in key:
|
||||
patched_state_dict[key.replace("net.0.proj", "w13")] = value
|
||||
for subkey in quant_sub_keys:
|
||||
quant_param_key = key.replace("qweight", subkey)
|
||||
if quant_param_key in state_dict:
|
||||
patched_state_dict[quant_param_key.replace("net.0.proj", "w13")] = state_dict[quant_param_key]
|
||||
elif any("feed_forward.net.0.proj." + subkey in key for subkey in quant_sub_keys):
|
||||
continue
|
||||
elif "feed_forward.net.2.qweight" in key:
|
||||
patched_state_dict[key.replace("net.2", "w2")] = value
|
||||
for subkey in quant_sub_keys:
|
||||
quant_param_key = key.replace("qweight", subkey)
|
||||
if quant_param_key in state_dict:
|
||||
patched_state_dict[quant_param_key.replace("net.2", "w2")] = state_dict[quant_param_key]
|
||||
elif any("feed_forward.net.2." + subkey in key for subkey in quant_sub_keys):
|
||||
continue
|
||||
elif "feed_forward.net.0.proj.weight" in key:
|
||||
w3, w1 = torch.chunk(value, chunks=2, dim=0)
|
||||
w2 = state_dict[key.replace("0.proj", "2")]
|
||||
patched_state_dict[key.replace("net.0.proj", "w1")] = w1
|
||||
patched_state_dict[key.replace("net.0.proj", "w2")] = w2
|
||||
patched_state_dict[key.replace("net.0.proj", "w3")] = w3
|
||||
elif "feed_forward.net.2.weight" in key:
|
||||
continue
|
||||
elif "attention.norm_q" in key:
|
||||
patched_state_dict[key.replace("norm_q", "q_norm")] = value
|
||||
elif "attention.norm_k" in key:
|
||||
patched_state_dict[key.replace("norm_k", "k_norm")] = value
|
||||
elif "all_final_layer.2-1" in key:
|
||||
patched_state_dict[key.replace("all_final_layer.2-1", "final_layer")] = value
|
||||
elif "all_x_embedder.2-1" in key:
|
||||
patched_state_dict[key.replace("all_x_embedder.2-1", "x_embedder")] = value
|
||||
else:
|
||||
patched_state_dict[key] = value
|
||||
|
||||
return patched_state_dict
|
||||
|
||||
|
||||
class NunchakuZImageModel(NextDiT):
|
||||
"""https://github.com/nunchaku-tech/ComfyUI-nunchaku/blob/dev/models/zimage.py#L101"""
|
||||
|
||||
def __init__(self, filename: str, precision: str, rank: int, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
_kwargs = {"precision": precision, "rank": rank}
|
||||
|
||||
def _patch_transformer_block(block_list: list[torch.nn.Module]):
|
||||
for _, block in enumerate(block_list):
|
||||
block.attention = NunchakuZImageAttention(block.attention, **_kwargs)
|
||||
block.feed_forward = NunchakuZImageFeedForward(block.feed_forward, **_kwargs)
|
||||
|
||||
_patch_transformer_block(self.layers)
|
||||
_patch_transformer_block(self.noise_refiner)
|
||||
_patch_transformer_block(self.context_refiner)
|
||||
|
||||
self.norm_final = None
|
||||
|
||||
def load_state_dict(self, sd, *args, **kwargs):
|
||||
sd = patch_z_image_state_dict(sd)
|
||||
return super().load_state_dict(sd, *args, **kwargs)
|
||||
|
||||
@ -129,7 +129,7 @@ class ForgeOperations:
|
||||
self.in_features = in_features
|
||||
self.out_features = out_features
|
||||
self.dummy = torch.nn.Parameter(torch.empty(1, device=current_device, dtype=current_dtype))
|
||||
self.weight = None
|
||||
self.weight = torch.empty([0], device=current_device, dtype=current_dtype) # SVDQW4A4Linear.from_linear
|
||||
self.scale_weight = None
|
||||
self.bias = None
|
||||
self.parameters_manual_cast = current_manual_cast_enabled
|
||||
@ -137,6 +137,7 @@ class ForgeOperations:
|
||||
def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs):
|
||||
if hasattr(self, "dummy"):
|
||||
if prefix + "weight" in state_dict:
|
||||
del self.weight
|
||||
self.weight = torch.nn.Parameter(state_dict[prefix + "weight"].to(self.dummy))
|
||||
if prefix + "scale_weight" in state_dict:
|
||||
self.scale_weight = torch.nn.Parameter(state_dict[prefix + "scale_weight"])
|
||||
|
||||
@ -57,6 +57,7 @@ def detect_unet_config(state_dict: dict, key_prefix: str):
|
||||
dit_config["rope_theta"] = 10000.0
|
||||
dit_config["ffn_dim_multiplier"] = 4.0
|
||||
elif dit_config["dim"] == 3840: # Z-Image
|
||||
dit_config["nunchaku"] = "{}layers.0.attention.to_out.0.qweight".format(key_prefix) in state_dict_keys
|
||||
dit_config["n_heads"] = 30
|
||||
dit_config["n_kv_heads"] = 30
|
||||
dit_config["axes_dims"] = [32, 48, 48]
|
||||
|
||||
@ -61,6 +61,7 @@ class BASE:
|
||||
|
||||
def __init__(self, unet_config):
|
||||
self.unet_config = unet_config.copy()
|
||||
self.nunchaku: bool = self.unet_config.pop("nunchaku", False)
|
||||
self.sampling_settings = self.sampling_settings.copy()
|
||||
self.latent_format = self.latent_format()
|
||||
for x in self.unet_extra_config:
|
||||
@ -246,10 +247,6 @@ class Flux(BASE):
|
||||
|
||||
unet_target = "transformer"
|
||||
|
||||
def __init__(self, unet_config):
|
||||
super().__init__(unet_config)
|
||||
self.nunchaku: bool = self.unet_config.pop("nunchaku", False)
|
||||
|
||||
def model_type(self, state_dict):
|
||||
return ModelType.FLUX
|
||||
|
||||
@ -442,10 +439,6 @@ class QwenImage(BASE):
|
||||
|
||||
unet_target = "transformer"
|
||||
|
||||
def __init__(self, unet_config):
|
||||
super().__init__(unet_config)
|
||||
self.nunchaku: bool = self.unet_config.pop("nunchaku", False)
|
||||
|
||||
def model_type(self, state_dict):
|
||||
return ModelType.FLOW
|
||||
|
||||
|
||||
Loading…
Reference in New Issue
Block a user