Skip to content
31 changes: 31 additions & 0 deletions examples/v1/z-image-turbo.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,31 @@
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
transformer = NunchakuZImageTransformer2DModel.from_pretrained(
f"/PATH/TO/svdq-{precision}_r128-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=9, # This actually results in 8 DiT forwards
guidance_scale=0.0, # Guidance should be 0 for the Turbo models
generator=torch.Generator("cuda").manual_seed(12345),
).images[0]

image.save(f"tmp_imgs/z-image-turbo-{precision}.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",
]
21 changes: 18 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 Down Expand Up @@ -108,6 +108,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 +137,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 +156,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",
]
76 changes: 76 additions & 0 deletions nunchaku/models/attention_processors/zimage.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
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):

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

# Apapted 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:

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