diff --git a/src/mflux/models/text_encoder/clip_encoder/clip_sdpa_attention.py b/src/mflux/models/text_encoder/clip_encoder/clip_sdpa_attention.py index e56a0b2..36937ba 100644 --- a/src/mflux/models/text_encoder/clip_encoder/clip_sdpa_attention.py +++ b/src/mflux/models/text_encoder/clip_encoder/clip_sdpa_attention.py @@ -1,5 +1,6 @@ import mlx.core as mx from mlx import nn +from mlx.core.fast import scaled_dot_product_attention class CLIPSdpaAttention(nn.Module): @@ -23,22 +24,15 @@ class CLIPSdpaAttention(nn.Module): key = CLIPSdpaAttention.reshape_and_transpose(key, self.batch_size, self.num_heads, self.head_dimension) value = CLIPSdpaAttention.reshape_and_transpose(value, self.batch_size, self.num_heads, self.head_dimension) - hidden_states = CLIPSdpaAttention.masked_attention(query, key, value, causal_attention_mask) + scale = 1 / mx.sqrt(query.shape[-1]) + hidden_states = scaled_dot_product_attention(query, key, value, scale=scale, mask=causal_attention_mask) + hidden_states = mx.transpose(hidden_states, (0, 2, 1, 3)) hidden_states = mx.reshape(hidden_states, (self.batch_size, -1, self.num_heads * self.head_dimension)) hidden_states = self.out_proj(hidden_states) return hidden_states - @staticmethod - def masked_attention(query, key, value, mask): - scale = 1 / mx.sqrt(query.shape[-1]) - scores = (query * scale) @ key.transpose(0, 1, 3, 2) - scores = scores + mask - attn = mx.softmax(scores, axis=-1) - hidden_states = attn @ value - return hidden_states - @staticmethod def reshape_and_transpose(x, batch_size, num_heads, head_dim): return mx.transpose(mx.reshape(x, (batch_size, -1, num_heads, head_dim)), (0, 2, 1, 3)) diff --git a/src/mflux/models/transformer/joint_attention.py b/src/mflux/models/transformer/joint_attention.py index 2dd991a..eda3074 100644 --- a/src/mflux/models/transformer/joint_attention.py +++ b/src/mflux/models/transformer/joint_attention.py @@ -1,5 +1,6 @@ import mlx.core as mx from mlx import nn +from mlx.core.fast import scaled_dot_product_attention class JointAttention(nn.Module): @@ -65,7 +66,9 @@ class JointAttention(nn.Module): query, key = JointAttention.apply_rope(query, key, image_rotary_emb) - hidden_states = JointAttention.attention(query, key, value) + scale = 1 / mx.sqrt(query.shape[-1]) + hidden_states = scaled_dot_product_attention(query, key, value, scale=scale) + hidden_states = mx.transpose(hidden_states, (0, 2, 1, 3)) hidden_states = mx.reshape( hidden_states, @@ -81,14 +84,6 @@ class JointAttention(nn.Module): return hidden_states, encoder_hidden_states - @staticmethod - def attention(query, key, value): - scale = 1 / mx.sqrt(query.shape[-1]) - scores = (query * scale) @ key.transpose(0, 1, 3, 2) - attn = mx.softmax(scores, axis=-1) - hidden_states = attn @ value - return hidden_states - @staticmethod def apply_rope(xq: mx.array, xk: mx.array, freqs_cis: mx.array): xq_ = xq.astype(mx.float32).reshape(*xq.shape[:-1], -1, 1, 2) diff --git a/src/mflux/models/transformer/single_block_attention.py b/src/mflux/models/transformer/single_block_attention.py index 8b9aec4..ea0a1a2 100644 --- a/src/mflux/models/transformer/single_block_attention.py +++ b/src/mflux/models/transformer/single_block_attention.py @@ -1,5 +1,6 @@ import mlx.core as mx from mlx import nn +from mlx.core.fast import scaled_dot_product_attention class SingleBlockAttention(nn.Module): @@ -29,7 +30,9 @@ class SingleBlockAttention(nn.Module): query, key = SingleBlockAttention.apply_rope(query, key, image_rotary_emb) - hidden_states = SingleBlockAttention.attention(query, key, value) + scale = 1 / mx.sqrt(query.shape[-1]) + hidden_states = scaled_dot_product_attention(query, key, value, scale=scale) + hidden_states = mx.transpose(hidden_states, (0, 2, 1, 3)) hidden_states = mx.reshape( hidden_states, @@ -38,14 +41,6 @@ class SingleBlockAttention(nn.Module): return hidden_states - @staticmethod - def attention(query, key, value): - scale = 1 / mx.sqrt(query.shape[-1]) - scores = (query * scale) @ key.transpose(0, 1, 3, 2) - attn = mx.softmax(scores, axis=-1) - hidden_states = attn @ value - return hidden_states - @staticmethod def apply_rope(xq: mx.array, xk: mx.array, freqs_cis: mx.array): xq_ = xq.astype(mx.float32).reshape(*xq.shape[:-1], -1, 1, 2) diff --git a/src/mflux/models/vae/common/attention.py b/src/mflux/models/vae/common/attention.py index a23059c..1ad1568 100644 --- a/src/mflux/models/vae/common/attention.py +++ b/src/mflux/models/vae/common/attention.py @@ -1,5 +1,6 @@ import mlx.core as mx from mlx import nn +from mlx.core.fast import scaled_dot_product_attention from mflux.config.config import Config @@ -20,14 +21,18 @@ class Attention(nn.Module): y = self.group_norm(input_array.astype(mx.float32)).astype(Config.precision) - queries = self.to_q(y).reshape(B, H * W, C) - keys = self.to_k(y).reshape(B, H * W, C) - values = self.to_v(y).reshape(B, H * W, C) + queries = self.to_q(y).reshape(B, H * W, 1, C) + keys = self.to_k(y).reshape(B, H * W, 1, C) + values = self.to_v(y).reshape(B, H * W, 1, C) + + queries = mx.transpose(queries, (0, 2, 1, 3)) + keys = mx.transpose(keys, (0, 2, 1, 3)) + values = mx.transpose(values, (0, 2, 1, 3)) scale = 1 / mx.sqrt(queries.shape[-1]) - scores = (queries * scale) @ keys.transpose(0, 2, 1) - attn = mx.softmax(scores, axis=-1) - y = (attn @ values).reshape(B, H, W, C) + y = scaled_dot_product_attention(queries, keys, values, scale=scale) + + y = mx.transpose(y, (0, 2, 1, 3)).reshape(B, H, W, C) y = self.to_out[0](y) output_tensor = input_array + y diff --git a/tests/resources/reference_controlnet_dev.png b/tests/resources/reference_controlnet_dev.png index e62f8b4..a27827a 100644 Binary files a/tests/resources/reference_controlnet_dev.png and b/tests/resources/reference_controlnet_dev.png differ diff --git a/tests/resources/reference_controlnet_dev_lora.png b/tests/resources/reference_controlnet_dev_lora.png index d6c2ffb..3a0d94c 100644 Binary files a/tests/resources/reference_controlnet_dev_lora.png and b/tests/resources/reference_controlnet_dev_lora.png differ diff --git a/tests/resources/reference_controlnet_schnell.png b/tests/resources/reference_controlnet_schnell.png index 0ba9dd3..217392c 100644 Binary files a/tests/resources/reference_controlnet_schnell.png and b/tests/resources/reference_controlnet_schnell.png differ diff --git a/tests/resources/reference_dev.png b/tests/resources/reference_dev.png index 5b5713a..dcd63e6 100644 Binary files a/tests/resources/reference_dev.png and b/tests/resources/reference_dev.png differ diff --git a/tests/resources/reference_dev_image_to_image_result.png b/tests/resources/reference_dev_image_to_image_result.png index 5adfe09..799d0f4 100644 Binary files a/tests/resources/reference_dev_image_to_image_result.png and b/tests/resources/reference_dev_image_to_image_result.png differ diff --git a/tests/resources/reference_dev_lora.png b/tests/resources/reference_dev_lora.png index f52cb8f..cd22b5e 100644 Binary files a/tests/resources/reference_dev_lora.png and b/tests/resources/reference_dev_lora.png differ diff --git a/tests/resources/reference_dev_lora_multiple.png b/tests/resources/reference_dev_lora_multiple.png index 9229ae5..cf24dde 100644 Binary files a/tests/resources/reference_dev_lora_multiple.png and b/tests/resources/reference_dev_lora_multiple.png differ diff --git a/tests/resources/reference_schnell.png b/tests/resources/reference_schnell.png index 2d61b2c..699ae6f 100644 Binary files a/tests/resources/reference_schnell.png and b/tests/resources/reference_schnell.png differ