Skip to content
Merged
1 change: 1 addition & 0 deletions docs/source/index.rst
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@ Check out `DeepCompressor <github_deepcompressor_>`_ for the quantization librar
usage/fbcache.rst
usage/pulid.rst
usage/ip_adapter.rst
usage/zimage.rst

.. toctree::
:maxdepth: 1
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,3 +6,4 @@ nunchaku.models.attention_processors

nunchaku.models.attention_processors.flux
nunchaku.models.attention_processors.qwenimage
nunchaku.models.attention_processors.zimage
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
nunchaku.models.attention_processors.zimage
===========================================

.. automodule:: nunchaku.models.attention_processors.zimage
:members:
:undoc-members:
:show-inheritance:
1 change: 1 addition & 0 deletions docs/source/python_api/nunchaku.models.transformers.rst
Original file line number Diff line number Diff line change
Expand Up @@ -7,5 +7,6 @@ nunchaku.models.transformers
nunchaku.models.transformers.transformer_flux
nunchaku.models.transformers.transformer_flux_v2
nunchaku.models.transformers.transformer_qwenimage
nunchaku.models.transformers.transformer_zimage
nunchaku.models.transformers.transformer_sana
nunchaku.models.transformers.utils
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
nunchaku.models.transformers.transformer\_zimage
================================================

.. automodule:: nunchaku.models.transformers.transformer_zimage
:members:
:undoc-members:
:show-inheritance:
16 changes: 16 additions & 0 deletions docs/source/usage/zimage.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1,16 @@
Z-Image
=======

The following is the example of running Nunchaku version of Z-Image text-to-image pipeline.

.. tabs::

.. tab:: Z-Image-Turbo

.. literalinclude:: ../../../examples/v1/z-image-turbo.py
:language: python
:caption: Running Z-Image-Turbo (`examples/v1/z-image-turbo.py <https://github.com/nunchaku-tech/nunchaku/blob/main/examples/v1/z-image-turbo.py>`__)
:linenos:


For more details, see :class:`~nunchaku.models.transformers.transformer_zimage.NunchakuZImageTransformer2DModel`.
29 changes: 29 additions & 0 deletions examples/v1/z-image-turbo.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
import torch
from diffusers.pipelines.z_image.pipeline_z_image import ZImagePipeline

from nunchaku import NunchakuZImageTransformer2DModel
from nunchaku.utils import get_precision

if __name__ == "__main__":
precision = get_precision() # auto-detect your precision is 'int4' or 'fp4' based on your GPU
rank = 128 # Use 32 for faster sampling; 256 (INT4 only) for best quality
transformer = NunchakuZImageTransformer2DModel.from_pretrained(
f"nunchaku-tech/nunchaku-z-image-turbo/svdq-{precision}_r{rank}-z-image-turbo.safetensors"
)

pipe = ZImagePipeline.from_pretrained(
"Tongyi-MAI/Z-Image-Turbo", transformer=transformer, torch_dtype=torch.bfloat16, low_cpu_mem_usage=False
).to("cuda")

prompt = "a young military male cooking in the kitchen for therapy"

image = pipe(
prompt=prompt,
height=1024,
width=1024,
num_inference_steps=8, # This actually results in 8 DiT forwards

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This should probably be set to num_inference_steps=9 to match the official example: https://github.com/Tongyi-MAI/Z-Image?tab=readme-ov-file#-quick-start

guidance_scale=0.0, # Guidance should be 0 for the Turbo models
generator=torch.Generator().manual_seed(12345),
).images[0]

image.save(f"z-image-turbo-{precision}_r{rank}.png")
2 changes: 2 additions & 0 deletions nunchaku/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
NunchakuQwenImageTransformer2DModel,
NunchakuSanaTransformer2DModel,
NunchakuT5EncoderModel,
NunchakuZImageTransformer2DModel,
)

__all__ = [
Expand All @@ -12,4 +13,5 @@
"NunchakuT5EncoderModel",
"NunchakuFluxTransformer2DModelV2",
"NunchakuQwenImageTransformer2DModel",
"NunchakuZImageTransformer2DModel",
]
23 changes: 20 additions & 3 deletions nunchaku/merge_safetensors.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,7 @@


def merge_safetensors(
pretrained_model_name_or_path: str | os.PathLike[str], **kwargs
pretrained_model_name_or_path: str | os.PathLike[str], model_class: str, **kwargs
) -> tuple[dict[str, torch.Tensor], dict[str, str]]:
"""
Merge split safetensors model files into a single state dict and metadata.
Expand All @@ -47,6 +47,8 @@ def merge_safetensors(
----------
pretrained_model_name_or_path : str or os.PathLike
Path to the model directory or HuggingFace repo.
model_class : str
Specify model class. E.g. NunchakuFluxTransformer2dModel or NunchakuZImageTransformer2DModel
**kwargs
Additional keyword arguments for subfolder, comfy_config_path, and HuggingFace download options.

Expand Down Expand Up @@ -108,6 +110,9 @@ def merge_safetensors(
state_dict.update(transformer_block_sd)

rank = next((v.shape[1] for k, v in transformer_block_sd.items() if ".lora_down" in k), 32)
if "ZImage" in model_class:
rank = next((v.shape[1] for k, v in transformer_block_sd.items() if ".proj_down" in k), 32)
skip_refiners = not any(("refiner" in k and "attention.to_qkv" in k) for k in transformer_block_sd.keys())

precision = "int4"
for v in state_dict.values():
Expand All @@ -134,10 +139,12 @@ def merge_safetensors(
},
"rank": rank,
}
if "ZImage" in model_class:
quantization_config["skip_refiners"] = skip_refiners
return state_dict, {
"config": Path(config_path).read_text(),
"comfy_config": Path(comfy_config_path).read_text(),
"model_class": "NunchakuFluxTransformer2dModel",
"model_class": model_class,
"quantization_config": json.dumps(quantization_config),
}

Expand All @@ -151,10 +158,20 @@ def merge_safetensors(
required=True,
help="Path to model directory. It can also be a huggingface repo.",
)
parser.add_argument(
"-m",
"--model-class",
type=str,
required=True,
help="Specify model class. E.g. NunchakuFluxTransformer2dModel or NunchakuZImageTransformer2DModel",
)
parser.add_argument("-o", "--output-path", type=Path, required=True, help="Path to output path")
args = parser.parse_args()
state_dict, metadata = merge_safetensors(args.input_path)
state_dict, metadata = merge_safetensors(args.input_path, args.model_class)
output_path = Path(args.output_path)
print(f" --input-path: {args.input_path}")
print(f" --model-class: {args.model_class}")
print(f" --output-path: {args.output_path}")
dirpath = output_path.parent
dirpath.mkdir(parents=True, exist_ok=True)
save_file(state_dict, output_path, metadata=metadata)
2 changes: 2 additions & 0 deletions nunchaku/models/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
NunchakuFluxTransformer2DModelV2,
NunchakuQwenImageTransformer2DModel,
NunchakuSanaTransformer2DModel,
NunchakuZImageTransformer2DModel,
)

__all__ = [
Expand All @@ -12,4 +13,5 @@
"NunchakuT5EncoderModel",
"NunchakuFluxTransformer2DModelV2",
"NunchakuQwenImageTransformer2DModel",
"NunchakuZImageTransformer2DModel",
]
88 changes: 88 additions & 0 deletions nunchaku/models/attention_processors/zimage.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,88 @@
"""
Attention processor implementations for :class:`~nunchaku.models.transformers.transformer_zimage.NunchakuZImageAttention`.
"""

from typing import Optional

import torch
from diffusers.models.attention_dispatch import dispatch_attention_fn
from diffusers.models.transformers.transformer_z_image import ZSingleStreamAttnProcessor


class NunchakuZSingleStreamAttnProcessor(ZSingleStreamAttnProcessor):
"""
Nunchaku attention processor for Z-Image-Turbo.
Adapted from diffusers.models.transformers.transformer_z_image.ZSingleStreamAttnProcessor.

"""

def __init__(self):
super().__init__()

# Adapted from diffusers.models.transformers.transformer_z_image.ZSingleStreamAttnProcessor#__call__
def __call__(
self,
attn,
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
freqs_cis: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""
Forward pass of the attention module. Adapted from diffusers.models.transformers.transformer_z_image.ZSingleStreamAttnProcessor#__call__.
"""

qkv = attn.to_qkv(hidden_states)
query, key, value = qkv.chunk(3, dim=-1)

query = query.unflatten(-1, (attn.heads, -1))
key = key.unflatten(-1, (attn.heads, -1))
value = value.unflatten(-1, (attn.heads, -1))

# Apply Norms
if attn.norm_q is not None:
query = attn.norm_q(query)
if attn.norm_k is not None:
key = attn.norm_k(key)

# Apply RoPE
def apply_rotary_emb(x_in: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
with torch.amp.autocast("cuda", enabled=False):
x = torch.view_as_complex(x_in.float().reshape(*x_in.shape[:-1], -1, 2))
freqs_cis = freqs_cis.unsqueeze(2)
x_out = torch.view_as_real(x * freqs_cis).flatten(3)
return x_out.type_as(x_in) # todo

if freqs_cis is not None:
query = apply_rotary_emb(query, freqs_cis)
key = apply_rotary_emb(key, freqs_cis)

# Cast to correct dtype
dtype = query.dtype
query, key = query.to(dtype), key.to(dtype)

# From [batch, seq_len] to [batch, 1, 1, seq_len] -> broadcast to [batch, heads, seq_len, seq_len]
if attention_mask is not None and attention_mask.ndim == 2:
attention_mask = attention_mask[:, None, None, :]

# Compute joint attention
hidden_states = dispatch_attention_fn(
query,
key,
value,
attn_mask=attention_mask,
dropout_p=0.0,
is_causal=False,
backend=self._attention_backend,
parallel_config=self._parallel_config,
)

# Reshape back
hidden_states = hidden_states.flatten(2, 3)
hidden_states = hidden_states.to(dtype)

output = attn.to_out[0](hidden_states)
if len(attn.to_out) > 1: # dropout
output = attn.to_out[1](output)

return output
2 changes: 2 additions & 0 deletions nunchaku/models/transformers/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,10 +2,12 @@
from .transformer_flux_v2 import NunchakuFluxTransformer2DModelV2
from .transformer_qwenimage import NunchakuQwenImageTransformer2DModel
from .transformer_sana import NunchakuSanaTransformer2DModel
from .transformer_zimage import NunchakuZImageTransformer2DModel

__all__ = [
"NunchakuFluxTransformer2dModel",
"NunchakuSanaTransformer2DModel",
"NunchakuFluxTransformer2DModelV2",
"NunchakuQwenImageTransformer2DModel",
"NunchakuZImageTransformer2DModel",
]
17 changes: 2 additions & 15 deletions nunchaku/models/transformers/transformer_flux_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,7 @@
from ..linear import SVDQW4A4Linear
from ..normalization import NunchakuAdaLayerNormZero, NunchakuAdaLayerNormZeroSingle
from ..utils import fuse_linears
from .utils import NunchakuModelLoaderMixin
from .utils import NunchakuModelLoaderMixin, patch_scale_key


class NunchakuFluxAttention(NunchakuBaseAttention):
Expand Down Expand Up @@ -421,20 +421,7 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike[str],
transformer = transformer.to_empty(device=device)
converted_state_dict = convert_flux_state_dict(model_state_dict)

state_dict = transformer.state_dict()

for k in state_dict.keys():
if k not in converted_state_dict:
assert ".wcscales" in k
converted_state_dict[k] = torch.ones_like(state_dict[k])
else:
assert state_dict[k].dtype == converted_state_dict[k].dtype

# Load the wtscale from the converted state dict.
for n, m in transformer.named_modules():
if isinstance(m, SVDQW4A4Linear):
if m.wtscale is not None:
m.wtscale = converted_state_dict.pop(f"{n}.wtscale", 1.0)
patch_scale_key(transformer, converted_state_dict)

transformer.load_state_dict(converted_state_dict)

Expand Down
17 changes: 3 additions & 14 deletions nunchaku/models/transformers/transformer_qwenimage.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,7 @@
from ..attention_processors.qwenimage import NunchakuQwenImageNaiveFA2Processor
from ..linear import AWQW4A16Linear, SVDQW4A4Linear
from ..utils import CPUOffloadManager, fuse_linears
from .utils import NunchakuModelLoaderMixin
from .utils import NunchakuModelLoaderMixin, patch_scale_key

logger = diffusers_logging.get_logger(__name__)

Expand Down Expand Up @@ -405,19 +405,8 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike[str],
theta=10000, axes_dim=list(config.get("axes_dims_rope", [16, 56, 56])), scale_rope=True
)

state_dict = transformer.state_dict()
for k in state_dict.keys():
if k not in model_state_dict:
assert ".wcscales" in k
model_state_dict[k] = torch.ones_like(state_dict[k])
else:
assert state_dict[k].dtype == model_state_dict[k].dtype

# load the wtscale from the state dict, as it is a float on CPU
for n, m in transformer.named_modules():
if isinstance(m, SVDQW4A4Linear):
if m.wtscale is not None:
m.wtscale = model_state_dict.pop(f"{n}.wtscale", 1.0)
patch_scale_key(transformer, model_state_dict)

transformer.load_state_dict(model_state_dict)
transformer.set_offload(offload)

Expand Down
Loading