Reuse existing method in joint transformer

This commit is contained in:
filipstrand 2025-06-04 22:27:22 +02:00
parent 674c44e30d
commit 55dd6ee3e8
2 changed files with 7 additions and 26 deletions

View File

@ -6,6 +6,7 @@ from mlx import nn
from mflux.community.concept_attention.joint_attention_concept import JointAttentionConcept from mflux.community.concept_attention.joint_attention_concept import JointAttentionConcept
from mflux.models.transformer.ada_layer_norm_zero import AdaLayerNormZero from mflux.models.transformer.ada_layer_norm_zero import AdaLayerNormZero
from mflux.models.transformer.feed_forward import FeedForward from mflux.models.transformer.feed_forward import FeedForward
from mflux.models.transformer.joint_transformer_block import JointTransformerBlock
@dataclass @dataclass
@ -66,7 +67,7 @@ class JointTransformerBlockConcept(nn.Module):
) )
# 3a. Apply norm and feed forward for hidden states # 3a. Apply norm and feed forward for hidden states
hidden_states = JointTransformerBlockConcept._apply_norm_and_feed_forward( hidden_states = JointTransformerBlock.apply_norm_and_feed_forward(
hidden_states=hidden_states, hidden_states=hidden_states,
attn_output=attn_output, attn_output=attn_output,
gate_mlp=gate_mlp, gate_mlp=gate_mlp,
@ -78,7 +79,7 @@ class JointTransformerBlockConcept(nn.Module):
) )
# 3b. Apply norm and feed forward for encoder hidden states # 3b. Apply norm and feed forward for encoder hidden states
encoder_hidden_states = JointTransformerBlockConcept._apply_norm_and_feed_forward( encoder_hidden_states = JointTransformerBlock.apply_norm_and_feed_forward(
hidden_states=encoder_hidden_states, hidden_states=encoder_hidden_states,
attn_output=context_attn_output, attn_output=context_attn_output,
gate_mlp=c_gate_mlp, gate_mlp=c_gate_mlp,
@ -90,7 +91,7 @@ class JointTransformerBlockConcept(nn.Module):
) )
# 3c. Apply norm and feed forward for concept encoder hidden states # 3c. Apply norm and feed forward for concept encoder hidden states
encoder_hidden_states_concept = JointTransformerBlockConcept._apply_norm_and_feed_forward( encoder_hidden_states_concept = JointTransformerBlock.apply_norm_and_feed_forward(
hidden_states=encoder_hidden_states_concept, hidden_states=encoder_hidden_states_concept,
attn_output=context_attn_output_concept, attn_output=context_attn_output_concept,
gate_mlp=c_gate_mlp_concept, gate_mlp=c_gate_mlp_concept,
@ -109,23 +110,3 @@ class JointTransformerBlockConcept(nn.Module):
) )
return encoder_hidden_states, hidden_states, encoder_hidden_states_concept, layer_attention_data return encoder_hidden_states, hidden_states, encoder_hidden_states_concept, layer_attention_data
@staticmethod
def _apply_norm_and_feed_forward(
hidden_states: mx.array,
attn_output: mx.array,
gate_mlp: mx.array,
gate_msa: mx.array,
scale_mlp: mx.array,
shift_mlp: mx.array,
norm_layer: nn.Module,
ff_layer: nn.Module,
) -> mx.array:
attn_output = mx.expand_dims(gate_msa, axis=1) * attn_output
hidden_states = hidden_states + attn_output
norm_hidden_states = norm_layer(hidden_states)
norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None]
ff_output = ff_layer(norm_hidden_states)
ff_output = mx.expand_dims(gate_mlp, axis=1) * ff_output
hidden_states = hidden_states + ff_output
return hidden_states

View File

@ -45,7 +45,7 @@ class JointTransformerBlock(nn.Module):
) )
# 3a. Apply norm and feed forward for hidden states # 3a. Apply norm and feed forward for hidden states
hidden_states = JointTransformerBlock._apply_norm_and_feed_forward( hidden_states = JointTransformerBlock.apply_norm_and_feed_forward(
hidden_states=hidden_states, hidden_states=hidden_states,
attn_output=attn_output, attn_output=attn_output,
gate_mlp=gate_mlp, gate_mlp=gate_mlp,
@ -57,7 +57,7 @@ class JointTransformerBlock(nn.Module):
) )
# 3b. Apply norm and feed forward for encoder hidden states # 3b. Apply norm and feed forward for encoder hidden states
encoder_hidden_states = JointTransformerBlock._apply_norm_and_feed_forward( encoder_hidden_states = JointTransformerBlock.apply_norm_and_feed_forward(
hidden_states=encoder_hidden_states, hidden_states=encoder_hidden_states,
attn_output=context_attn_output, attn_output=context_attn_output,
gate_mlp=c_gate_mlp, gate_mlp=c_gate_mlp,
@ -71,7 +71,7 @@ class JointTransformerBlock(nn.Module):
return encoder_hidden_states, hidden_states return encoder_hidden_states, hidden_states
@staticmethod @staticmethod
def _apply_norm_and_feed_forward( def apply_norm_and_feed_forward(
hidden_states: mx.array, hidden_states: mx.array,
attn_output: mx.array, attn_output: mx.array,
gate_mlp: mx.array, gate_mlp: mx.array,