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/common/__init__.py b/src/mflux/models/transformer/common/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/mflux/models/transformer/common/attention_utils.py b/src/mflux/models/transformer/common/attention_utils.py new file mode 100644 index 0000000..424ae48 --- /dev/null +++ b/src/mflux/models/transformer/common/attention_utils.py @@ -0,0 +1,59 @@ +import mlx.core as mx +from mlx import nn +from mlx.core.fast import scaled_dot_product_attention + + +class AttentionUtils: + @staticmethod + def process_qkv( + hidden_states: mx.array, + to_q: nn.Linear, + to_k: nn.Linear, + to_v: nn.Linear, + norm_q: nn.RMSNorm, + norm_k: nn.RMSNorm, + num_heads: int, + head_dim: int, + ): + query = to_q(hidden_states) + key = to_k(hidden_states) + value = to_v(hidden_states) + + # Reshape and transpose + query = mx.transpose(mx.reshape(query, (1, -1, num_heads, head_dim)), (0, 2, 1, 3)) + key = mx.transpose(mx.reshape(key, (1, -1, num_heads, head_dim)), (0, 2, 1, 3)) + value = mx.transpose(mx.reshape(value, (1, -1, num_heads, head_dim)), (0, 2, 1, 3)) + + # Apply normalization + query = norm_q(query) + key = norm_k(key) + + return query, key, value + + @staticmethod + def compute_attention( + query: mx.array, + key: mx.array, + value: mx.array, + batch_size: int, + num_heads: int, + head_dim: int + ) -> mx.array: # fmt: off + 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, + (batch_size, -1, num_heads * head_dim), + ) + + 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) + xk_ = xk.astype(mx.float32).reshape(*xk.shape[:-1], -1, 1, 2) + xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1] + xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1] + return xq_out.reshape(*xq.shape).astype(mx.float32), xk_out.reshape(*xk.shape).astype(mx.float32) diff --git a/src/mflux/models/transformer/joint_attention.py b/src/mflux/models/transformer/joint_attention.py index 2dd991a..c114ae4 100644 --- a/src/mflux/models/transformer/joint_attention.py +++ b/src/mflux/models/transformer/joint_attention.py @@ -1,14 +1,16 @@ import mlx.core as mx from mlx import nn +from mflux.models.transformer.common.attention_utils import AttentionUtils + class JointAttention(nn.Module): - head_dimension = 128 - batch_size = 1 - num_heads = 24 - def __init__(self): super().__init__() + self.head_dimension = 128 + self.batch_size = 1 + self.num_heads = 24 + self.to_q = nn.Linear(3072, 3072) self.to_k = nn.Linear(3072, 3072) self.to_v = nn.Linear(3072, 3072) @@ -28,49 +30,42 @@ class JointAttention(nn.Module): encoder_hidden_states: mx.array, image_rotary_emb: mx.array, ) -> (mx.array, mx.array): - query = self.to_q(hidden_states) - key = self.to_k(hidden_states) - value = self.to_v(hidden_states) - - query = mx.transpose(mx.reshape(query, (1, -1, 24, 128)), (0, 2, 1, 3)) - key = mx.transpose(mx.reshape(key, (1, -1, 24, 128)), (0, 2, 1, 3)) - value = mx.transpose(mx.reshape(value, (1, -1, 24, 128)), (0, 2, 1, 3)) - - query = self.norm_q(query) - key = self.norm_k(key) - - encoder_hidden_states_query_proj = self.add_q_proj(encoder_hidden_states) - encoder_hidden_states_key_proj = self.add_k_proj(encoder_hidden_states) - encoder_hidden_states_value_proj = self.add_v_proj(encoder_hidden_states) - - encoder_hidden_states_query_proj = mx.transpose( - mx.reshape(encoder_hidden_states_query_proj, (1, -1, 24, 128)), - (0, 2, 1, 3), + query, key, value = AttentionUtils.process_qkv( + hidden_states=hidden_states, + to_q=self.to_q, + to_k=self.to_k, + to_v=self.to_v, + norm_q=self.norm_q, + norm_k=self.norm_k, + num_heads=self.num_heads, + head_dim=self.head_dimension, ) - encoder_hidden_states_key_proj = mx.transpose( - mx.reshape(encoder_hidden_states_key_proj, (1, -1, 24, 128)), - (0, 2, 1, 3), - ) - encoder_hidden_states_value_proj = mx.transpose( - mx.reshape(encoder_hidden_states_value_proj, (1, -1, 24, 128)), - (0, 2, 1, 3), + enc_query, enc_key, enc_value = AttentionUtils.process_qkv( + hidden_states=encoder_hidden_states, + to_q=self.add_q_proj, + to_k=self.add_k_proj, + to_v=self.add_v_proj, + norm_q=self.norm_added_q, + norm_k=self.norm_added_k, + num_heads=self.num_heads, + head_dim=self.head_dimension, ) - encoder_hidden_states_query_proj = self.norm_added_q(encoder_hidden_states_query_proj) - encoder_hidden_states_key_proj = self.norm_added_k(encoder_hidden_states_key_proj) + query = mx.concatenate([enc_query, query], axis=2) + key = mx.concatenate([enc_key, key], axis=2) + value = mx.concatenate([enc_value, value], axis=2) - query = mx.concatenate([encoder_hidden_states_query_proj, query], axis=2) - key = mx.concatenate([encoder_hidden_states_key_proj, key], axis=2) - value = mx.concatenate([encoder_hidden_states_value_proj, value], axis=2) + query, key = AttentionUtils.apply_rope(xq=query, xk=key, freqs_cis=image_rotary_emb) - query, key = JointAttention.apply_rope(query, key, image_rotary_emb) - - hidden_states = JointAttention.attention(query, key, value) - 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 = AttentionUtils.compute_attention( + query=query, + key=key, + value=value, + batch_size=self.batch_size, + num_heads=self.num_heads, + head_dim=self.head_dimension, ) + encoder_hidden_states, hidden_states = ( hidden_states[:, : encoder_hidden_states.shape[1]], hidden_states[:, encoder_hidden_states.shape[1] :], @@ -80,19 +75,3 @@ class JointAttention(nn.Module): encoder_hidden_states = self.to_add_out(encoder_hidden_states) 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) - xk_ = xk.astype(mx.float32).reshape(*xk.shape[:-1], -1, 1, 2) - xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1] - xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1] - return xq_out.reshape(*xq.shape).astype(mx.float32), xk_out.reshape(*xk.shape).astype(mx.float32) diff --git a/src/mflux/models/transformer/single_block_attention.py b/src/mflux/models/transformer/single_block_attention.py index 8b9aec4..88526b1 100644 --- a/src/mflux/models/transformer/single_block_attention.py +++ b/src/mflux/models/transformer/single_block_attention.py @@ -1,55 +1,41 @@ import mlx.core as mx from mlx import nn +from mflux.models.transformer.common.attention_utils import AttentionUtils + class SingleBlockAttention(nn.Module): - head_dimension = 128 - batch_size = 1 - num_heads = 24 - def __init__(self): super().__init__() + self.head_dimension = 128 + self.batch_size = 1 + self.num_heads = 24 + self.to_q = nn.Linear(3072, 3072) self.to_k = nn.Linear(3072, 3072) self.to_v = nn.Linear(3072, 3072) self.norm_q = nn.RMSNorm(128) self.norm_k = nn.RMSNorm(128) - def __call__(self, hidden_states: mx.array, image_rotary_emb: mx.array) -> (mx.array, mx.array): - query = self.to_q(hidden_states) - key = self.to_k(hidden_states) - value = self.to_v(hidden_states) - - query = mx.transpose(mx.reshape(query, (1, -1, 24, 128)), (0, 2, 1, 3)) - key = mx.transpose(mx.reshape(key, (1, -1, 24, 128)), (0, 2, 1, 3)) - value = mx.transpose(mx.reshape(value, (1, -1, 24, 128)), (0, 2, 1, 3)) - - query = self.norm_q(query) - key = self.norm_k(key) - - query, key = SingleBlockAttention.apply_rope(query, key, image_rotary_emb) - - hidden_states = SingleBlockAttention.attention(query, key, value) - 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), + def __call__(self, hidden_states: mx.array, image_rotary_emb: mx.array) -> mx.array: + query, key, value = AttentionUtils.process_qkv( + hidden_states=hidden_states, + to_q=self.to_q, + to_k=self.to_k, + to_v=self.to_v, + norm_q=self.norm_q, + norm_k=self.norm_k, + num_heads=self.num_heads, + head_dim=self.head_dimension, ) - return hidden_states + query, key = AttentionUtils.apply_rope(xq=query, xk=key, freqs_cis=image_rotary_emb) - @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) - xk_ = xk.astype(mx.float32).reshape(*xk.shape[:-1], -1, 1, 2) - xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1] - xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1] - return xq_out.reshape(*xq.shape).astype(mx.float32), xk_out.reshape(*xk.shape).astype(mx.float32) + return AttentionUtils.compute_attention( + query=query, + key=key, + value=value, + batch_size=self.batch_size, + num_heads=self.num_heads, + head_dim=self.head_dimension, + ) 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