Specify precision for causal attention mask
This commit is contained in:
parent
3b0f5a4516
commit
16a0225b41
@ -2,6 +2,8 @@ import mlx.core as mx
|
|||||||
from mlx import nn
|
from mlx import nn
|
||||||
from mlx.core.fast import scaled_dot_product_attention
|
from mlx.core.fast import scaled_dot_product_attention
|
||||||
|
|
||||||
|
from mflux import Config
|
||||||
|
|
||||||
|
|
||||||
class CLIPSdpaAttention(nn.Module):
|
class CLIPSdpaAttention(nn.Module):
|
||||||
head_dimension = 64
|
head_dimension = 64
|
||||||
@ -16,6 +18,8 @@ class CLIPSdpaAttention(nn.Module):
|
|||||||
self.out_proj = nn.Linear(input_dims=768, output_dims=768)
|
self.out_proj = nn.Linear(input_dims=768, output_dims=768)
|
||||||
|
|
||||||
def __call__(self, hidden_states: mx.array, causal_attention_mask: mx.array) -> mx.array:
|
def __call__(self, hidden_states: mx.array, causal_attention_mask: mx.array) -> mx.array:
|
||||||
|
causal_attention_mask = causal_attention_mask.astype(Config.precision)
|
||||||
|
|
||||||
query = self.q_proj(hidden_states)
|
query = self.q_proj(hidden_states)
|
||||||
key = self.k_proj(hidden_states)
|
key = self.k_proj(hidden_states)
|
||||||
value = self.v_proj(hidden_states)
|
value = self.v_proj(hidden_states)
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user