Skip to content

Commit 0caf3ff

Browse files
committed
Add First Block Cache support for Flux 2
1 parent ac56fa2 commit 0caf3ff

4 files changed

Lines changed: 39 additions & 25 deletions

File tree

src/diffusers/hooks/_helpers.py

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -175,6 +175,7 @@ def _register_transformer_blocks_metadata():
175175
from ..models.transformers.transformer_bria import BriaTransformerBlock
176176
from ..models.transformers.transformer_cogview4 import CogView4TransformerBlock
177177
from ..models.transformers.transformer_flux import FluxSingleTransformerBlock, FluxTransformerBlock
178+
from ..models.transformers.transformer_flux2 import Flux2SingleTransformerBlock, Flux2TransformerBlock
178179
from ..models.transformers.transformer_hunyuan_video import (
179180
HunyuanVideoSingleTransformerBlock,
180181
HunyuanVideoTokenReplaceSingleTransformerBlock,
@@ -246,6 +247,22 @@ def _register_transformer_blocks_metadata():
246247
),
247248
)
248249

250+
# Flux2
251+
TransformerBlockRegistry.register(
252+
model_class=Flux2TransformerBlock,
253+
metadata=TransformerBlockMetadata(
254+
return_hidden_states_index=1,
255+
return_encoder_hidden_states_index=0,
256+
),
257+
)
258+
TransformerBlockRegistry.register(
259+
model_class=Flux2SingleTransformerBlock,
260+
metadata=TransformerBlockMetadata(
261+
return_hidden_states_index=1,
262+
return_encoder_hidden_states_index=0,
263+
),
264+
)
265+
249266
# HunyuanVideo
250267
TransformerBlockRegistry.register(
251268
model_class=HunyuanVideoTransformerBlock,

src/diffusers/models/transformers/transformer_flux2.py

Lines changed: 12 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -837,18 +837,13 @@ def __init__(
837837
def forward(
838838
self,
839839
hidden_states: torch.Tensor,
840-
encoder_hidden_states: torch.Tensor | None,
840+
encoder_hidden_states: torch.Tensor,
841841
temb_mod: torch.Tensor,
842842
image_rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None,
843843
joint_attention_kwargs: dict[str, Any] | None = None,
844-
split_hidden_states: bool = False,
845-
text_seq_len: int | None = None,
846844
) -> tuple[torch.Tensor, torch.Tensor]:
847-
# If encoder_hidden_states is None, hidden_states is assumed to have encoder_hidden_states already
848-
# concatenated
849-
if encoder_hidden_states is not None:
850-
text_seq_len = encoder_hidden_states.shape[1]
851-
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
845+
text_seq_len = encoder_hidden_states.shape[1]
846+
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
852847

853848
mod_shift, mod_scale, mod_gate = Flux2Modulation.split(temb_mod, 1)[0]
854849

@@ -866,11 +861,8 @@ def forward(
866861
if hidden_states.dtype == torch.float16:
867862
hidden_states = hidden_states.clip(-65504, 65504)
868863

869-
if split_hidden_states:
870-
encoder_hidden_states, hidden_states = hidden_states[:, :text_seq_len], hidden_states[:, text_seq_len:]
871-
return encoder_hidden_states, hidden_states
872-
else:
873-
return hidden_states
864+
encoder_hidden_states, hidden_states = hidden_states[:, :text_seq_len], hidden_states[:, text_seq_len:]
865+
return encoder_hidden_states, hidden_states
874866

875867

876868
class Flux2TransformerBlock(nn.Module):
@@ -1374,12 +1366,9 @@ def forward(
13741366
joint_attention_kwargs=kv_attn_kwargs,
13751367
)
13761368

1377-
# Concatenate text and image streams for single-block inference
1378-
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
1379-
13801369
# Blend single block modulation for extract mode: [txt_mod, ref_mod, img_mod]
13811370
if kv_cache_mode == "extract" and num_ref_tokens > 0:
1382-
total_single_len = hidden_states.shape[1]
1371+
total_single_len = num_txt_tokens + hidden_states.shape[1]
13831372
single_stream_mod = _blend_single_block_mods(
13841373
single_stream_mod, ref_single_mod, num_txt_tokens, num_ref_tokens, total_single_len
13851374
)
@@ -1396,28 +1385,26 @@ def forward(
13961385
kv_attn_kwargs_single["kv_cache"] = kv_cache.get_single(index_block)
13971386

13981387
if torch.is_grad_enabled() and self.gradient_checkpointing:
1399-
hidden_states = self._gradient_checkpointing_func(
1388+
encoder_hidden_states, hidden_states = self._gradient_checkpointing_func(
14001389
block,
14011390
hidden_states,
1402-
None,
1391+
encoder_hidden_states,
14031392
single_stream_mod,
14041393
concat_rotary_emb,
14051394
kv_attn_kwargs_single,
14061395
)
14071396
else:
1408-
hidden_states = block(
1397+
encoder_hidden_states, hidden_states = block(
14091398
hidden_states=hidden_states,
1410-
encoder_hidden_states=None,
1399+
encoder_hidden_states=encoder_hidden_states,
14111400
temb_mod=single_stream_mod,
14121401
image_rotary_emb=concat_rotary_emb,
14131402
joint_attention_kwargs=kv_attn_kwargs_single,
14141403
)
14151404

1416-
# Remove text tokens (and ref tokens in extract mode) from concatenated stream
1405+
# Remove ref tokens (extract mode only) from the image stream
14171406
if kv_cache_mode == "extract" and num_ref_tokens > 0:
1418-
hidden_states = hidden_states[:, num_txt_tokens + num_ref_tokens :, ...]
1419-
else:
1420-
hidden_states = hidden_states[:, num_txt_tokens:, ...]
1407+
hidden_states = hidden_states[:, num_ref_tokens:, ...]
14211408

14221409
# 7. Output layers
14231410
hidden_states = self.norm_out(hidden_states, temb)

tests/models/transformers/test_models_transformer_flux2.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,7 @@
3434
BaseModelTesterConfig,
3535
BitsAndBytesTesterMixin,
3636
ContextParallelTesterMixin,
37+
FirstBlockCacheTesterMixin,
3738
GGUFCompileTesterMixin,
3839
GGUFTesterMixin,
3940
LoraHotSwappingForModelTesterMixin,
@@ -198,6 +199,10 @@ def test_tensor_parallel_neuron_inference(self):
198199
)
199200

200201

202+
class TestFlux2TransformerFBCCache(Flux2TransformerTesterConfig, FirstBlockCacheTesterMixin):
203+
"""FirstBlockCache tests for Flux2 Transformer."""
204+
205+
201206
class TestFlux2TransformerLoRA(Flux2TransformerTesterConfig, LoraTesterMixin):
202207
"""LoRA adapter tests for Flux2 Transformer."""
203208

tests/pipelines/flux2/test_pipeline_flux2_klein.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323
)
2424
from ..testing_utils import (
2525
BasePipelineTesterConfig,
26+
FirstBlockCacheTesterMixin,
2627
MemoryTesterMixin,
2728
PipelineTesterMixin,
2829
check_qkv_fused_layers_exist,
@@ -198,6 +199,10 @@ class TestFlux2KleinPipelineMemory(Flux2KleinPipelineTesterConfig, MemoryTesterM
198199
"""Memory optimization tests (CPU offload, group offload, layerwise casting) for the Flux2 Klein pipeline."""
199200

200201

202+
class TestFlux2KleinPipelineFirstBlockCache(Flux2KleinPipelineTesterConfig, FirstBlockCacheTesterMixin):
203+
"""First Block Cache tests for the Flux2 Klein pipeline."""
204+
205+
201206
@require_torch_neuron
202207
class TestFlux2KleinPipelineIntegration:
203208
ckpt_id = "black-forest-labs/FLUX.2-klein-4B"

0 commit comments

Comments
 (0)