Qwen-Image-Layered-MRP-MLX/src/mflux/models/transformer/joint_attention.py

85 lines
2.8 KiB
Python

import mlx.core as mx
from mlx import nn
from mflux.models.transformer.common.attention_utils import AttentionUtils
class JointAttention(nn.Module):
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.to_out = [nn.Linear(3072, 3072)]
self.add_q_proj = nn.Linear(3072, 3072)
self.add_k_proj = nn.Linear(3072, 3072)
self.add_v_proj = nn.Linear(3072, 3072)
self.to_add_out = nn.Linear(3072, 3072)
self.norm_q = nn.RMSNorm(128)
self.norm_k = nn.RMSNorm(128)
self.norm_added_q = nn.RMSNorm(128)
self.norm_added_k = nn.RMSNorm(128)
def __call__(
self,
hidden_states: mx.array,
encoder_hidden_states: mx.array,
image_rotary_emb: mx.array,
) -> (mx.array, mx.array):
# 1a. Compute Q,K,V for hidden_states
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,
)
# 1b. Compute Q,K,V for encoder_hidden_states
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,
)
# 1c. Concatenate results
query = mx.concatenate([enc_query, query], axis=2)
key = mx.concatenate([enc_key, key], axis=2)
value = mx.concatenate([enc_value, value], axis=2)
# 1d. Apply rope to Q,K
query, key = AttentionUtils.apply_rope(xq=query, xk=key, freqs_cis=image_rotary_emb)
# 2. Compute attention
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,
)
# 3. Separate the results
encoder_hidden_states, hidden_states = (
hidden_states[:, : encoder_hidden_states.shape[1]],
hidden_states[:, encoder_hidden_states.shape[1] :],
)
# 4. Project the output
hidden_states = self.to_out[0](hidden_states)
encoder_hidden_states = self.to_add_out(encoder_hidden_states)
return hidden_states, encoder_hidden_states