From febae44b799f89917ee56062977a7062f384f3e0 Mon Sep 17 00:00:00 2001 From: BuildTools Date: Fri, 31 Jul 2026 20:52:42 -0600 Subject: [PATCH 1/5] fix hooks --- src/diffusers/hooks/_helpers.py | 24 ++++++++++++++++++++++++ 1 file changed, 24 insertions(+) diff --git a/src/diffusers/hooks/_helpers.py b/src/diffusers/hooks/_helpers.py index 9cbe5bc8108f..77d982fd6681 100644 --- a/src/diffusers/hooks/_helpers.py +++ b/src/diffusers/hooks/_helpers.py @@ -109,6 +109,7 @@ def _register_attention_processors_metadata(): from ..models.attention_processor import AttnProcessor2_0 from ..models.transformers.transformer_cogview4 import CogView4AttnProcessor from ..models.transformers.transformer_flux import FluxAttnProcessor + from ..models.transformers.transformer_flux2 import Flux2AttnProcessor from ..models.transformers.transformer_hunyuanimage import HunyuanImageAttnProcessor from ..models.transformers.transformer_qwenimage import QwenDoubleStreamAttnProcessor2_0 from ..models.transformers.transformer_wan import WanAttnProcessor2_0 @@ -144,6 +145,12 @@ def _register_attention_processors_metadata(): metadata=AttentionProcessorMetadata(skip_processor_output_fn=_skip_proc_output_fn_Attention_FluxAttnProcessor), ) + # Flux2AttnProcessor + AttentionProcessorRegistry.register( + model_class=Flux2AttnProcessor, + metadata=AttentionProcessorMetadata(skip_processor_output_fn=_skip_proc_output_fn_Attention_FluxAttnProcessor), + ) + # QwenDoubleStreamAttnProcessor2 AttentionProcessorRegistry.register( model_class=QwenDoubleStreamAttnProcessor2_0, @@ -175,6 +182,7 @@ def _register_transformer_blocks_metadata(): from ..models.transformers.transformer_bria import BriaTransformerBlock from ..models.transformers.transformer_cogview4 import CogView4TransformerBlock from ..models.transformers.transformer_flux import FluxSingleTransformerBlock, FluxTransformerBlock + from ..models.transformers.transformer_flux2 import Flux2SingleTransformerBlock, Flux2TransformerBlock from ..models.transformers.transformer_hunyuan_video import ( HunyuanVideoSingleTransformerBlock, HunyuanVideoTokenReplaceSingleTransformerBlock, @@ -246,6 +254,22 @@ def _register_transformer_blocks_metadata(): ), ) + # Flux2 + TransformerBlockRegistry.register( + model_class=Flux2TransformerBlock, + metadata=TransformerBlockMetadata( + return_hidden_states_index=1, + return_encoder_hidden_states_index=0, + ), + ) + TransformerBlockRegistry.register( + model_class=Flux2SingleTransformerBlock, + metadata=TransformerBlockMetadata( + return_hidden_states_index=1, + return_encoder_hidden_states_index=0, + ), + ) + # HunyuanVideo TransformerBlockRegistry.register( model_class=HunyuanVideoTransformerBlock, From 17ea853eddf21db878940fec8d1c5b67717b61b8 Mon Sep 17 00:00:00 2001 From: BuildTools Date: Fri, 31 Jul 2026 21:38:54 -0600 Subject: [PATCH 2/5] add tail overrides --- src/diffusers/hooks/_helpers.py | 2 +- src/diffusers/hooks/mag_cache.py | 4 +++- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/src/diffusers/hooks/_helpers.py b/src/diffusers/hooks/_helpers.py index 77d982fd6681..6207a96b176d 100644 --- a/src/diffusers/hooks/_helpers.py +++ b/src/diffusers/hooks/_helpers.py @@ -13,7 +13,7 @@ # limitations under the License. import inspect -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import Any, Callable, Type diff --git a/src/diffusers/hooks/mag_cache.py b/src/diffusers/hooks/mag_cache.py index e5f0aaebc01a..72ee8280a147 100644 --- a/src/diffusers/hooks/mag_cache.py +++ b/src/diffusers/hooks/mag_cache.py @@ -346,8 +346,10 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs): diff = in_hidden.shape[1] - out_hidden.shape[1] if diff == 0: residual = out_hidden - in_hidden + elif diff > 0: + residual = out_hidden - in_hidden[:, diff:] # Fallback to matching tail else: - residual = out_hidden - in_hidden # Fallback to matching tail + residual = out_hidden[:, -diff:] - in_hidden # Fallback to matching tail else: # Fallback for completely mismatched shapes residual = out_hidden From 290efa3cc38a288cb8bc36e80e569b33a24cbd97 Mon Sep 17 00:00:00 2001 From: BuildTools Date: Sat, 1 Aug 2026 00:21:06 -0600 Subject: [PATCH 3/5] fix --- src/diffusers/hooks/_helpers.py | 4 ++-- src/diffusers/hooks/mag_cache.py | 13 +++++++++++++ 2 files changed, 15 insertions(+), 2 deletions(-) diff --git a/src/diffusers/hooks/_helpers.py b/src/diffusers/hooks/_helpers.py index 6207a96b176d..8a3ae18de7da 100644 --- a/src/diffusers/hooks/_helpers.py +++ b/src/diffusers/hooks/_helpers.py @@ -265,8 +265,8 @@ def _register_transformer_blocks_metadata(): TransformerBlockRegistry.register( model_class=Flux2SingleTransformerBlock, metadata=TransformerBlockMetadata( - return_hidden_states_index=1, - return_encoder_hidden_states_index=0, + return_hidden_states_index=0, + return_encoder_hidden_states_index=None, ), ) diff --git a/src/diffusers/hooks/mag_cache.py b/src/diffusers/hooks/mag_cache.py index 72ee8280a147..ac700f12bae6 100644 --- a/src/diffusers/hooks/mag_cache.py +++ b/src/diffusers/hooks/mag_cache.py @@ -328,6 +328,7 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs): output = self.fn_ref.original_forward(*args, **kwargs) if self.is_tail: + fuse = False # Calculate residual for next steps if isinstance(output, tuple): out_hidden = output[self._metadata.return_hidden_states_index] @@ -344,12 +345,17 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs): residual = out_hidden - in_hidden elif out_hidden.ndim == 3 and in_hidden.ndim == 3 and out_hidden.shape[2] == in_hidden.shape[2]: diff = in_hidden.shape[1] - out_hidden.shape[1] + print(diff) if diff == 0: residual = out_hidden - in_hidden elif diff > 0: + print("falling back in", diff) residual = out_hidden - in_hidden[:, diff:] # Fallback to matching tail + fuse = diff else: + print("falling back out", diff) residual = out_hidden[:, -diff:] - in_hidden # Fallback to matching tail + fuse = diff else: # Fallback for completely mismatched shapes residual = out_hidden @@ -359,6 +365,13 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs): state.previous_residual = residual self._advance_step(state) + print(fuse) + # if fuse: + # if fuse > 0: + # text_tokens = in_hidden[:, :fuse] + # return torch.cat([text_tokens, output], dim=1) + # else: + # return out_hidden return output From f8034378851771e374843509e1a76ab20ac0f62c Mon Sep 17 00:00:00 2001 From: BuildTools Date: Sat, 1 Aug 2026 00:23:32 -0600 Subject: [PATCH 4/5] remove debug --- src/diffusers/hooks/mag_cache.py | 13 ------------- 1 file changed, 13 deletions(-) diff --git a/src/diffusers/hooks/mag_cache.py b/src/diffusers/hooks/mag_cache.py index ac700f12bae6..72ee8280a147 100644 --- a/src/diffusers/hooks/mag_cache.py +++ b/src/diffusers/hooks/mag_cache.py @@ -328,7 +328,6 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs): output = self.fn_ref.original_forward(*args, **kwargs) if self.is_tail: - fuse = False # Calculate residual for next steps if isinstance(output, tuple): out_hidden = output[self._metadata.return_hidden_states_index] @@ -345,17 +344,12 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs): residual = out_hidden - in_hidden elif out_hidden.ndim == 3 and in_hidden.ndim == 3 and out_hidden.shape[2] == in_hidden.shape[2]: diff = in_hidden.shape[1] - out_hidden.shape[1] - print(diff) if diff == 0: residual = out_hidden - in_hidden elif diff > 0: - print("falling back in", diff) residual = out_hidden - in_hidden[:, diff:] # Fallback to matching tail - fuse = diff else: - print("falling back out", diff) residual = out_hidden[:, -diff:] - in_hidden # Fallback to matching tail - fuse = diff else: # Fallback for completely mismatched shapes residual = out_hidden @@ -365,13 +359,6 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs): state.previous_residual = residual self._advance_step(state) - print(fuse) - # if fuse: - # if fuse > 0: - # text_tokens = in_hidden[:, :fuse] - # return torch.cat([text_tokens, output], dim=1) - # else: - # return out_hidden return output From 3ca8570c92c89699cc178b654cd6abf8827b1841 Mon Sep 17 00:00:00 2001 From: BuildTools Date: Sat, 1 Aug 2026 01:35:09 -0600 Subject: [PATCH 5/5] make style, quality --- src/diffusers/hooks/_helpers.py | 2 +- src/diffusers/hooks/mag_cache.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/src/diffusers/hooks/_helpers.py b/src/diffusers/hooks/_helpers.py index 8a3ae18de7da..4a55bb12fe47 100644 --- a/src/diffusers/hooks/_helpers.py +++ b/src/diffusers/hooks/_helpers.py @@ -13,7 +13,7 @@ # limitations under the License. import inspect -from dataclasses import dataclass, field +from dataclasses import dataclass from typing import Any, Callable, Type diff --git a/src/diffusers/hooks/mag_cache.py b/src/diffusers/hooks/mag_cache.py index 72ee8280a147..c90a850a7c87 100644 --- a/src/diffusers/hooks/mag_cache.py +++ b/src/diffusers/hooks/mag_cache.py @@ -347,7 +347,7 @@ def new_forward(self, module: torch.nn.Module, *args, **kwargs): if diff == 0: residual = out_hidden - in_hidden elif diff > 0: - residual = out_hidden - in_hidden[:, diff:] # Fallback to matching tail + residual = out_hidden - in_hidden[:, diff:] # Fallback to matching tail else: residual = out_hidden[:, -diff:] - in_hidden # Fallback to matching tail else: