import torch import torch.nn.functional as F from typing import Any, Dict, Optional, Tuple from torch import nn import einops from diffusers.models.transformers.transformer_flux import FluxTransformerBlock from diffusers.models.attention import Attention from diffusers.models.embeddings import apply_rotary_emb def _trim_rope(rope, target_len: int): """ Ensure rotary embedding length matches the attention sequence length. Works for Diffusers' rope tuples (cos, sin) or a single tensor. Rope format: [seq_len, embed_dim], so dimension 0 is the sequence. Args: rope: Either a tuple (cos, sin) or a single tensor, or None target_len: Target sequence length Returns: Trimmed rope in the same format as input (takes last target_len positions) """ if rope is None: return None # Rope is a (cos, sin) tuple (Diffusers format) # Each tensor is [seq_len, embed_dim] if isinstance(rope, tuple): cos, sin = rope if cos.shape[0] != target_len: # Trim dimension 0 (sequence length) cos = cos[-target_len:, ...] sin = sin[-target_len:, ...] return (cos, sin) # Rope is a single tensor if rope.shape[0] != target_len: rope = rope[-target_len:, ...] return rope class FluxConceptAttentionProcessor: """ Custom attention processor for FLUX that implements concept attention. Exactly matches original FLUX attention pattern while adding concept observation stream. """ def __init__(self): if not hasattr(F, "scaled_dot_product_attention"): raise ImportError("FluxConceptAttentionProcessor requires PyTorch 2.0") def __call__( self, attn: Attention, hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor, attention_mask: Optional[torch.FloatTensor] = None, image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, concept_hidden_states: Optional[torch.Tensor] = None, concept_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, q_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, kv_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, **kwargs ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: # Concept-specific arguments are now explicit parameters (required for inspect.signature) batch_size, _, _ = encoder_hidden_states.shape # *** MAIN GENERATION STREAM *** - Exact FLUX Pattern # 1. `sample` projections (image hidden states) query = attn.to_q(hidden_states) key = attn.to_k(hidden_states) value = attn.to_v(hidden_states) inner_dim = key.shape[-1] head_dim = inner_dim // attn.heads image_query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) image_key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) image_value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) if attn.norm_q is not None: image_query = attn.norm_q(image_query) if attn.norm_k is not None: image_key = attn.norm_k(image_key) # 2. `context` projections (text encoder hidden states) - FLUX specific encoder_hidden_states_query_proj = attn.add_q_proj(encoder_hidden_states) encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states) encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states) encoder_hidden_states_query_proj = encoder_hidden_states_query_proj.view( batch_size, -1, attn.heads, head_dim ).transpose(1, 2) encoder_hidden_states_key_proj = encoder_hidden_states_key_proj.view( batch_size, -1, attn.heads, head_dim ).transpose(1, 2) encoder_hidden_states_value_proj = encoder_hidden_states_value_proj.view( batch_size, -1, attn.heads, head_dim ).transpose(1, 2) if attn.norm_added_q is not None: encoder_hidden_states_query_proj = attn.norm_added_q(encoder_hidden_states_query_proj) if attn.norm_added_k is not None: encoder_hidden_states_key_proj = attn.norm_added_k(encoder_hidden_states_key_proj) # 3. FLUX concatenation order: [encoder, image] query = torch.cat([encoder_hidden_states_query_proj, image_query], dim=2) key = torch.cat([encoder_hidden_states_key_proj, image_key], dim=2) value = torch.cat([encoder_hidden_states_value_proj, image_value], dim=2) # 4. Apply rotary embeddings to FULL concatenated tensors (FLUX way) # Use explicit q/kv ropes; fallback to image_rotary_emb rope_q = q_rotary_emb if q_rotary_emb is not None else image_rotary_emb rope_kv = kv_rotary_emb if kv_rotary_emb is not None else image_rotary_emb if rope_q is not None: # query shape after transpose: [B, heads, seq, head_dim] q_len = query.shape[2] # sequence dimension is at index 2 rope_q = _trim_rope(rope_q, q_len) query = apply_rotary_emb(query, rope_q, sequence_dim=2) if rope_kv is not None: k_len = key.shape[2] # sequence dimension is at index 2 rope_kv = _trim_rope(rope_kv, k_len) key = apply_rotary_emb(key, rope_kv, sequence_dim=2) # 5. Main text+image attention (FLUX standard) hidden_states = F.scaled_dot_product_attention( query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False ) hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) hidden_states = hidden_states.to(query.dtype) # *** CONCEPT STREAM *** - Exact same logic as text-image (following reference) concept_hidden_states_output = None concept_attention_maps = None if concept_hidden_states is not None: # Use TEXT projections for concepts (like reference: txt_attn.qkv) concept_query = attn.add_q_proj(concept_hidden_states) # Same as text! concept_key = attn.add_k_proj(concept_hidden_states) # Same as text! concept_value = attn.add_v_proj(concept_hidden_states) # Same as text! concept_query = concept_query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) concept_key = concept_key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) concept_value = concept_value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) if attn.norm_added_q is not None: concept_query = attn.norm_added_q(concept_query) # Text normalization if attn.norm_added_k is not None: concept_key = attn.norm_added_k(concept_key) # Text normalization concept_image_q = torch.cat([concept_query, image_query], dim=2) concept_image_k = torch.cat([concept_key, image_key], dim=2) concept_image_v = torch.cat([concept_value, image_value], dim=2) # Apply concept rotary embeddings to FULL concatenated tensor (like reference) if concept_rotary_emb is not None: # concept_image_q shape after transpose: [B, heads, seq, head_dim] cq_len = concept_image_q.shape[2] # sequence dimension is at index 2 ck_len = concept_image_k.shape[2] rope_cq = _trim_rope(concept_rotary_emb, cq_len) rope_ck = _trim_rope(concept_rotary_emb, ck_len) concept_image_q = apply_rotary_emb(concept_image_q, rope_cq, sequence_dim=2) concept_image_k = apply_rotary_emb(concept_image_k, rope_ck, sequence_dim=2) # Do the joint attention operation (like reference) concept_image_attn = F.scaled_dot_product_attention( concept_image_q, concept_image_k, concept_image_v, dropout_p=0.0, is_causal=False ) # Separate the concept attention (like reference: concept_attn = concept_image_attn[:, :, :concepts.shape[1]]) concept_attn = concept_image_attn[:, :, :concept_hidden_states.size(1)] concept_hidden_states_output = concept_attn.transpose(1, 2).reshape( batch_size, -1, attn.heads * head_dim ) # Compute attention maps from concept and image queries (before attention) # Save vectors for postprocessing (like reference implementation) concept_attention_maps = { 'concept_vectors': concept_query, # (batch, heads, concepts, dim) 'image_vectors': image_query # (batch, heads, patches, dim) } # 6. FLUX output processing encoder_hidden_states_out, hidden_states_out = ( hidden_states[:, : encoder_hidden_states.shape[1]], hidden_states[:, encoder_hidden_states.shape[1] :], ) # linear proj hidden_states_out = attn.to_out[0](hidden_states_out) # dropout hidden_states_out = attn.to_out[1](hidden_states_out) # FLUX specific: separate output projection for encoder encoder_hidden_states_out = attn.to_add_out(encoder_hidden_states_out) # Process concept outputs with same projections if concept_hidden_states_output is not None: concept_hidden_states_output = attn.to_out[0](concept_hidden_states_output) concept_hidden_states_output = attn.to_out[1](concept_hidden_states_output) # Return vectors for postprocessing (like reference implementation) concept_attention_maps = { 'concept_vectors': concept_hidden_states_output, # Final processed concept features 'image_vectors': hidden_states_out # Final processed image features } return hidden_states_out, encoder_hidden_states_out, concept_hidden_states_output, concept_attention_maps class FluxTransformerBlockWithConceptAttention(FluxTransformerBlock): """ Simplified FLUX transformer block with concept attention. Uses the elegant CogVideoX approach with custom attention processor. """ def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) self.attn.processor = FluxConceptAttentionProcessor() def forward( self, hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor, concept_hidden_states: Optional[torch.Tensor], temb: torch.Tensor, concept_temb: Optional[torch.Tensor], image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, concept_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, joint_attention_kwargs: Optional[Dict[str, Any]] = None, concept_attention_kwargs: Optional[Dict[str, Any]] = None, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1( hidden_states, emb=temb ) norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context( encoder_hidden_states, emb=temb ) # Concept normalization (use same temb as main stream if concept_temb is None) norm_concept_hidden_states = None concept_gate_msa = concept_shift_mlp = concept_scale_mlp = concept_gate_mlp = None concept_gate_ff = None if concept_hidden_states is not None: effective_concept_temb = concept_temb if concept_temb is not None else temb norm_concept_hidden_states, concept_gate_msa, concept_shift_mlp, concept_scale_mlp, concept_gate_mlp = self.norm1_context( concept_hidden_states, emb=effective_concept_temb ) joint_attention_kwargs = joint_attention_kwargs or {} # Pass concept-specific args through joint_attention_kwargs # (they're not accepted as direct args by FluxAttention.forward) if norm_concept_hidden_states is not None: joint_attention_kwargs['concept_hidden_states'] = norm_concept_hidden_states joint_attention_kwargs['concept_rotary_emb'] = concept_rotary_emb # Attention with concept attention (using our custom processor) attention_outputs = self.attn( hidden_states=norm_hidden_states, encoder_hidden_states=norm_encoder_hidden_states, image_rotary_emb=image_rotary_emb, **joint_attention_kwargs, ) if len(attention_outputs) == 4: attn_output, context_attn_output, concept_attn_output, concept_attention_maps = attention_outputs ip_attn_output = None elif len(attention_outputs) == 5: attn_output, context_attn_output, concept_attn_output, concept_attention_maps, ip_attn_output = attention_outputs else: # Fallback for when no concept attention attn_output, context_attn_output = attention_outputs[:2] concept_attn_output = None concept_attention_maps = None ip_attn_output = attention_outputs[2] if len(attention_outputs) > 2 else None ################## Process Concept Features FIRST (like CogVideoX) ################## if concept_attn_output is not None and concept_hidden_states is not None: # Apply concept attention gate and residual concept_attn_output = concept_gate_msa.unsqueeze(1) * concept_attn_output concept_hidden_states = concept_hidden_states + concept_attn_output # Concept feedforward processing (norm2_context is regular LayerNorm, not adaptive) norm_concept_hidden_states = self.norm2_context(concept_hidden_states) norm_concept_hidden_states = norm_concept_hidden_states * (1 + concept_scale_mlp[:, None]) + concept_shift_mlp[:, None] concept_ff_output = self.ff_context(norm_concept_hidden_states) concept_hidden_states = concept_hidden_states + concept_gate_mlp.unsqueeze(1) * concept_ff_output if concept_hidden_states.dtype == torch.float16: concept_hidden_states = concept_hidden_states.clip(-65504, 65504) ################## Now Process Main Generation Stream ################## # Standard FLUX processing for image features attn_output = gate_msa.unsqueeze(1) * attn_output hidden_states = hidden_states + attn_output norm_hidden_states = self.norm2(hidden_states) norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] ff_output = self.ff(norm_hidden_states) ff_output = gate_mlp.unsqueeze(1) * ff_output hidden_states = hidden_states + ff_output if ip_attn_output is not None: hidden_states = hidden_states + ip_attn_output # Standard FLUX processing for text features if context_attn_output is not None: context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output encoder_hidden_states = encoder_hidden_states + context_attn_output norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] context_ff_output = self.ff_context(norm_encoder_hidden_states) encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output if encoder_hidden_states.dtype == torch.float16: encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504) return encoder_hidden_states, hidden_states, concept_hidden_states, concept_attention_maps