Skip to content

Commit e2169b0

Browse files
committed
comfyUI bug fix
1 parent 91c18fd commit e2169b0

2 files changed

Lines changed: 113 additions & 3 deletions

File tree

nunchaku/caching/diffusers_adapters/flux.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -47,14 +47,14 @@ def new_forward(self, *args, **kwargs):
4747
unittest.mock.patch.object(self, "transformer_blocks", cached_transformer_blocks),
4848
unittest.mock.patch.object(self, "single_transformer_blocks", dummy_single_transformer_blocks),
4949
):
50+
transformer._is_cached = True
51+
transformer.cached_transformer_blocks = cached_transformer_blocks
52+
transformer.single_transformer_blocks = dummy_single_transformer_blocks
5053
return original_forward(*args, **kwargs)
5154
else:
5255
return original_forward(*args, **kwargs)
5356

5457
transformer.forward = new_forward.__get__(transformer)
55-
transformer.cached_transformer_blocks = cached_transformer_blocks
56-
transformer.single_transformer_blocks = dummy_single_transformer_blocks
57-
transformer._is_cached = True
5858

5959
return transformer
6060

Lines changed: 110 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,110 @@
1+
import torch
2+
from diffusers import (
3+
AutoencoderKL,
4+
FlowMatchEulerDiscreteScheduler,
5+
FluxControlNetModel,
6+
FluxControlNetPipeline,
7+
FluxPipeline,
8+
)
9+
from diffusers.models import FluxMultiControlNetModel
10+
from diffusers.utils import load_image
11+
from transformers import CLIPTextModel, CLIPTokenizer, T5EncoderModel, T5TokenizerFast
12+
13+
from nunchaku import NunchakuT5EncoderModel
14+
from nunchaku.caching.diffusers_adapters import apply_cache_on_pipe
15+
from nunchaku.models.transformers.transformer_flux import NunchakuFluxTransformer2dModel
16+
from nunchaku.utils import get_precision
17+
18+
bfl_repo = "black-forest-labs/FLUX.1-dev"
19+
dtype = torch.bfloat16 # or torch.float16, or torch.float32
20+
device = "cuda" # or "cpu" if you want to run on CPU
21+
22+
scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(bfl_repo, subfolder="scheduler", torch_dtype=dtype)
23+
text_encoder = CLIPTextModel.from_pretrained(bfl_repo, subfolder="text_encoder", torch_dtype=dtype)
24+
text_encoder_2 = T5EncoderModel.from_pretrained(bfl_repo, subfolder="text_encoder_2", torch_dtype=dtype)
25+
tokenizer = CLIPTokenizer.from_pretrained(
26+
bfl_repo, subfolder="tokenizer", torch_dtype=dtype, clean_up_tokenization_spaces=True
27+
)
28+
tokenizer_2 = T5TokenizerFast.from_pretrained(
29+
bfl_repo, subfolder="tokenizer_2", torch_dtype=dtype, clean_up_tokenization_spaces=True
30+
)
31+
vae = AutoencoderKL.from_pretrained(bfl_repo, subfolder="vae", torch_dtype=dtype)
32+
precision = get_precision()
33+
transformer = NunchakuFluxTransformer2dModel.from_pretrained(
34+
f"mit-han-lab/nunchaku-flux.1-dev/svdq-{precision}_r32-flux.1-dev.safetensors",
35+
# offload=True
36+
)
37+
transformer.set_attention_impl("nunchaku-fp16")
38+
39+
# qencoder
40+
text_encoder_2 = NunchakuT5EncoderModel.from_pretrained("mit-han-lab/nunchaku-t5/awq-int4-flux.1-t5xxl.safetensors")
41+
controlnet_union = FluxControlNetModel.from_pretrained(
42+
"Shakker-Labs/FLUX.1-dev-ControlNet-Union-Pro-2.0", torch_dtype=torch.bfloat16
43+
)
44+
controlnet = FluxMultiControlNetModel([controlnet_union]) # we always recommend loading via FluxMultiControlNetModel
45+
46+
47+
params = {
48+
"scheduler": scheduler,
49+
"vae": vae,
50+
"tokenizer": tokenizer,
51+
"tokenizer_2": tokenizer_2,
52+
"text_encoder": text_encoder,
53+
"text_encoder_2": text_encoder_2,
54+
"transformer": transformer,
55+
}
56+
# pipe
57+
pipe = FluxPipeline(**params).to(device, dtype=dtype)
58+
pipe_cn = FluxControlNetPipeline(**params, controlnet=controlnet).to(device, dtype)
59+
60+
# offload
61+
pipe.enable_sequential_cpu_offload(device=device)
62+
pipe_cn.enable_sequential_cpu_offload(device=device)
63+
64+
# cache
65+
apply_cache_on_pipe(
66+
pipe_cn,
67+
use_double_fb_cache=True,
68+
residual_diff_threshold_multi=0.09,
69+
residual_diff_threshold_single=0.12,
70+
)
71+
apply_cache_on_pipe(
72+
pipe_cn,
73+
use_double_fb_cache=False,
74+
residual_diff_threshold_multi=0,
75+
residual_diff_threshold_single=0,
76+
)
77+
78+
params = {
79+
"prompt": "A bohemian-style female travel blogger with sun-kissed skin and messy beach waves.",
80+
"height": 1152,
81+
"width": 768,
82+
"num_inference_steps": 30,
83+
"guidance_scale": 3.5,
84+
}
85+
86+
# pipe
87+
txt2img_res = pipe(
88+
**params,
89+
).images[0]
90+
txt2img_res.save("flux.1-dev-txt2img.jpg")
91+
92+
# cache
93+
apply_cache_on_pipe(
94+
pipe_cn,
95+
use_double_fb_cache=True,
96+
residual_diff_threshold_multi=0.09,
97+
residual_diff_threshold_single=0.12,
98+
)
99+
100+
# pipe_cn
101+
control_iamge = load_image(
102+
"https://huggingface.co/Shakker-Labs/FLUX.1-dev-ControlNet-Union-Pro/resolve/main/assets/openpose.jpg"
103+
)
104+
params["control_image"] = [control_iamge]
105+
params["controlnet_conditioning_scale"] = [0.9]
106+
params["control_guidance_end"] = [0.65]
107+
cn_res = pipe_cn(
108+
**params,
109+
).images[0]
110+
cn_res.save("flux.1-dev-cn-txt2img.jpg")

0 commit comments

Comments
 (0)