Skip to content

Commit d08e6a7

Browse files
committed
fix:txt2_img pipe and controlnet pipe share components
1 parent e302ec3 commit d08e6a7

2 files changed

Lines changed: 15 additions & 5 deletions

File tree

nunchaku/caching/diffusers_adapters/flux.py

Lines changed: 10 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -41,13 +41,19 @@ def apply_cache_on_transformer(
4141

4242
@functools.wraps(original_forward)
4343
def new_forward(self, *args, **kwargs):
44-
with (
45-
unittest.mock.patch.object(self, "transformer_blocks", cached_transformer_blocks),
46-
unittest.mock.patch.object(self, "single_transformer_blocks", dummy_single_transformer_blocks),
47-
):
44+
cache_context = utils.get_current_cache_context()
45+
if cache_context is not None:
46+
with (
47+
unittest.mock.patch.object(self, "transformer_blocks", cached_transformer_blocks),
48+
unittest.mock.patch.object(self, "single_transformer_blocks", dummy_single_transformer_blocks),
49+
):
50+
return original_forward(*args, **kwargs)
51+
else:
4852
return original_forward(*args, **kwargs)
4953

5054
transformer.forward = new_forward.__get__(transformer)
55+
transformer.cached_transformer_blocks = cached_transformer_blocks
56+
transformer.single_transformer_blocks = dummy_single_transformer_blocks
5157
transformer._is_cached = True
5258

5359
return transformer

nunchaku/caching/diffusers_adapters/sana.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,11 @@ def apply_cache_on_transformer(transformer: SanaTransformer2DModel, *, residual_
2323

2424
@functools.wraps(original_forward)
2525
def new_forward(self, *args, **kwargs):
26-
with unittest.mock.patch.object(self, "transformer_blocks", cached_transformer_blocks):
26+
cache_context = utils.get_current_cache_context()
27+
if cache_context is not None:
28+
with unittest.mock.patch.object(self, "transformer_blocks", cached_transformer_blocks):
29+
return original_forward(*args, **kwargs)
30+
else:
2731
return original_forward(*args, **kwargs)
2832

2933
transformer.forward = new_forward.__get__(transformer)

0 commit comments

Comments
 (0)