Skip to content

Commit 1cac51b

Browse files
committed
add Turing (20~ Series) GPU compatibility for Z Image Turbo
1 parent 936fe40 commit 1cac51b

4 files changed

Lines changed: 24 additions & 9 deletions

File tree

examples/v1/z-image-turbo.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -2,17 +2,18 @@
22
from diffusers.pipelines.z_image.pipeline_z_image import ZImagePipeline
33

44
from nunchaku import NunchakuZImageTransformer2DModel
5-
from nunchaku.utils import get_precision
5+
from nunchaku.utils import get_precision, is_turing
66

77
if __name__ == "__main__":
88
precision = get_precision() # auto-detect your precision is 'int4' or 'fp4' based on your GPU
99
rank = 128 # Use 32 for faster sampling; 256 (INT4 only) for best quality
10+
dtype = torch.float16 if is_turing() else torch.bfloat16 # Use float16 when Turing (20- series) GPU is used.
1011
transformer = NunchakuZImageTransformer2DModel.from_pretrained(
11-
f"nunchaku-tech/nunchaku-z-image-turbo/svdq-{precision}_r{rank}-z-image-turbo.safetensors"
12+
f"nunchaku-tech/nunchaku-z-image-turbo/svdq-{precision}_r{rank}-z-image-turbo.safetensors", torch_dtype=dtype
1213
)
1314

1415
pipe = ZImagePipeline.from_pretrained(
15-
"Tongyi-MAI/Z-Image-Turbo", transformer=transformer, torch_dtype=torch.bfloat16, low_cpu_mem_usage=False
16+
"Tongyi-MAI/Z-Image-Turbo", transformer=transformer, torch_dtype=dtype, low_cpu_mem_usage=False
1617
).to("cuda")
1718

1819
prompt = "a young military male cooking in the kitchen for therapy"
@@ -26,4 +27,4 @@
2627
generator=torch.Generator().manual_seed(12345),
2728
).images[0]
2829

29-
image.save(f"z-image-turbo-{precision}_r{rank}.png")
30+
image.save(f"z-image-turbo-{precision}_r{rank}_{str(dtype)}.png")

nunchaku/models/linear.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -149,11 +149,12 @@ def from_linear(cls, linear: nn.Linear, **kwargs):
149149
SVDQW4A4Linear
150150
"""
151151
in_features = kwargs.pop("in_features", linear.in_features)
152+
torch_dtype = kwargs.pop("torch_dtype", linear.weight.dtype)
152153
return cls(
153154
in_features=in_features,
154155
out_features=linear.out_features,
155156
bias=linear.bias is not None,
156-
torch_dtype=linear.weight.dtype,
157+
torch_dtype=torch_dtype,
157158
device=linear.weight.device,
158159
**kwargs,
159160
)

nunchaku/models/transformers/transformer_zimage.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,7 @@
2121
from ..attention_processors.zimage import NunchakuZSingleStreamAttnProcessor
2222
from ..linear import SVDQW4A4Linear
2323
from ..utils import fuse_linears
24-
from .utils import NunchakuModelLoaderMixin, patch_scale_key
24+
from .utils import NunchakuModelLoaderMixin, convert_fp16, patch_scale_key
2525

2626

2727
class NunchakuZImageAttention(NunchakuBaseAttention):
@@ -259,10 +259,12 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike[str],
259259

260260
print(f"quantization_config: {quantization_config}, rank={rank}, skip_refiners={skip_refiners}")
261261

262-
transformer._patch_model(skip_refiners=skip_refiners, precision=precision, rank=rank)
262+
transformer._patch_model(skip_refiners=skip_refiners, precision=precision, rank=rank, **kwargs)
263263
transformer = transformer.to_empty(device=device)
264264

265265
patch_scale_key(transformer, model_state_dict)
266+
if torch_dtype == torch.float16:
267+
convert_fp16(transformer, model_state_dict)
266268

267269
transformer.load_state_dict(model_state_dict)
268270

nunchaku/models/transformers/utils.py

Lines changed: 13 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -164,10 +164,21 @@ def patch_scale_key(transformer_from_config: nn.Module, state_dict_from_checkpoi
164164
if k not in state_dict_from_checkpoint:
165165
assert ".wcscales" in k
166166
state_dict_from_checkpoint[k] = torch.ones_like(state_dict[k])
167-
else:
168-
assert state_dict[k].dtype == state_dict_from_checkpoint[k].dtype
169167

170168
for n, m in transformer_from_config.named_modules():
171169
if isinstance(m, SVDQW4A4Linear):
172170
if m.wtscale is not None:
173171
m.wtscale = state_dict_from_checkpoint.pop(f"{n}.wtscale", 1.0)
172+
173+
174+
def convert_fp16(transformer_from_config: nn.Module, state_dict_from_checkpoint: dict):
175+
state_dict = transformer_from_config.state_dict()
176+
for k in state_dict.keys():
177+
if state_dict[k].dtype != state_dict_from_checkpoint[k].dtype:
178+
assert (
179+
state_dict[k].dtype == torch.float16 and state_dict_from_checkpoint[k].dtype == torch.bfloat16
180+
), f"Unexpected dtype difference for key: {k}, model dtype: {state_dict[k].dtype}, \
181+
checkpoint dtype: {state_dict_from_checkpoint[k].dtype}"
182+
state_dict_from_checkpoint[k] = torch.nan_to_num(
183+
state_dict_from_checkpoint[k].to(torch.float16), nan=0.0, posinf=65504, neginf=-65504
184+
)

0 commit comments

Comments
 (0)