diff --git a/src/diffusers/hooks/_helpers.py b/src/diffusers/hooks/_helpers.py index 9cbe5bc8108f..4a55bb12fe47 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=0, + return_encoder_hidden_states_index=None, + ), + ) + # HunyuanVideo TransformerBlockRegistry.register( model_class=HunyuanVideoTransformerBlock, diff --git a/src/diffusers/hooks/mag_cache.py b/src/diffusers/hooks/mag_cache.py index e5f0aaebc01a..c90a850a7c87 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