Merge pull request #116 from filipstrand/optimizations

Optimizations and small refactoring of attention blocks
This commit is contained in:
Filip Strand 2025-01-14 21:18:04 +01:00 committed by GitHub
commit 3b0f5a4516
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
14 changed files with 135 additions and 112 deletions

View File

@ -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))

View File

@ -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)

View File

@ -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)

View File

@ -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,
)

View File

@ -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

Binary file not shown.

Before

Width:  |  Height:  |  Size: 416 KiB

After

Width:  |  Height:  |  Size: 416 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 484 KiB

After

Width:  |  Height:  |  Size: 484 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 273 KiB

After

Width:  |  Height:  |  Size: 273 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 367 KiB

After

Width:  |  Height:  |  Size: 367 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 350 KiB

After

Width:  |  Height:  |  Size: 350 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 374 KiB

After

Width:  |  Height:  |  Size: 374 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 445 KiB

After

Width:  |  Height:  |  Size: 445 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 392 KiB

After

Width:  |  Height:  |  Size: 392 KiB