1+ import functools
2+ import unittest
3+
14from diffusers import DiffusionPipeline , FluxTransformer2DModel
25from torch import nn
36
7+ from nunchaku .caching .utils import cache_context , create_cache_context
8+ from nunchaku .models .IP_adapter .utils import undo_all_mods_on_transformer
9+
410from ...IP_adapter import utils
511
612
713def apply_IPA_on_transformer (transformer : FluxTransformer2DModel , * , ip_adapter_scale : float = 1.0 , repo_id : str ):
8-
914 IPA_transformer_blocks = nn .ModuleList (
1015 [
1116 utils .IPA_TransformerBlocks (
@@ -16,19 +21,52 @@ def apply_IPA_on_transformer(transformer: FluxTransformer2DModel, *, ip_adapter_
1621 )
1722 ]
1823 )
24+ if getattr (transformer , "_is_cached" , False ):
25+ IPA_transformer_blocks [0 ].update_residual_diff_threshold (
26+ use_double_fb_cache = transformer .use_double_fb_cache ,
27+ residual_diff_threshold_multi = transformer .residual_diff_threshold_multi ,
28+ residual_diff_threshold_single = transformer .residual_diff_threshold_single ,
29+ )
30+ undo_all_mods_on_transformer (transformer )
31+ if not hasattr (transformer , "_original_forward" ):
32+ transformer ._original_forward = transformer .forward
33+ if not hasattr (transformer , "_original_blocks" ):
34+ transformer ._original_blocks = transformer .transformer_blocks
35+
1936 dummy_single_transformer_blocks = nn .ModuleList ()
2037
2138 IPA_transformer_blocks [0 ].load_ip_adapter_weights_per_layer (repo_id = repo_id )
2239
2340 transformer .transformer_blocks = IPA_transformer_blocks
2441 transformer .single_transformer_blocks = dummy_single_transformer_blocks
42+ original_forward = transformer .forward
2543
44+ @functools .wraps (original_forward )
45+ def new_forward (self , * args , ** kwargs ):
46+ with (
47+ unittest .mock .patch .object (self , "transformer_blocks" , IPA_transformer_blocks ),
48+ unittest .mock .patch .object (self , "single_transformer_blocks" , dummy_single_transformer_blocks ),
49+ ):
50+ return original_forward (* args , ** kwargs )
51+
52+ transformer .forward = new_forward .__get__ (transformer )
2653 transformer ._is_IPA = True
2754
2855 return transformer
2956
3057
3158def apply_IPA_on_pipe (pipe : DiffusionPipeline , * , shallow_patch : bool = False , ** kwargs ):
59+ if getattr (pipe , "_is_cached" , False ):
60+ original_call = pipe .__class__ .__call__
61+
62+ @functools .wraps (original_call )
63+ def new_call (self , * args , ** kwargs ):
64+ with cache_context (create_cache_context ()):
65+ return original_call (self , * args , ** kwargs )
66+
67+ pipe .__class__ .__call__ = new_call
68+ pipe .__class__ ._is_cached = True
69+
3270 if not shallow_patch :
3371 apply_IPA_on_transformer (pipe .transformer , ** kwargs )
3472
0 commit comments