@@ -101,10 +101,12 @@ def __call__(
101101 k = cache .update_and_fetch (k )
102102 if k .shape [2 ] <= self .index_topk :
103103 return None
104+ scores = q @ k .swapaxes (- 1 , - 2 )
105+ scores = mx .maximum (scores , 0 )
104106 weights = self .weights_proj (x ) * (self .n_heads ** - 0.5 )
105107 weights = (weights * self .softmax_scale ).swapaxes (- 1 , - 2 )[..., None ]
106- q_scaled = q * weights
107- scores = ( q * weights ) @ k . swapaxes ( - 1 , - 2 )
108+ scores = scores * weights
109+ scores = scores . sum ( axis = 1 )
108110 if mask is not None :
109111 scores = mx .where (mask , scores , - float ("inf" ))
110112 return mx .argpartition (scores , kth = - self .index_topk , axis = - 1 )[
@@ -219,24 +221,17 @@ def __call__(
219221 queries = mx .concatenate ([q_nope , q_pe ], axis = - 1 )
220222 topk_indices = self .indexer (x , qr , mask , cache = cache [1 ])
221223 if topk_indices is not None :
222- repeats = self .num_heads // self .config .index_n_heads
223- if L == 1 :
224- topk_indices = mx .repeat (topk_indices , repeats , axis = 1 ).squeeze (- 2 )[
225- ..., None
226- ]
227- keys = mx .take_along_axis (keys , topk_indices , axis = - 2 )
228- values = mx .take_along_axis (values , topk_indices , axis = - 2 )
229- else :
230- topk_mask = mx .zeros (
231- (B , self .config .index_n_heads , * mask .shape [- 2 :]), mx .bool_
232- )
233- topk_mask = mx .put_along_axis (
234- topk_mask , topk_indices , mx .array (True ), axis = - 1
235- )
236- mask = mask & topk_mask
237- mask = mx .repeat (mask , repeats , axis = 1 )
224+ k_seq = keys .shape [2 ]
225+ sparse_mask = mx .zeros ((B , L , k_seq ), dtype = mx .bool_ )
226+ sparse_mask = mx .put_along_axis (
227+ sparse_mask , topk_indices , mx .array (True ), axis = - 1
228+ )
229+ sparse_mask = sparse_mask [:, None , :, :]
230+ if mask is not None :
231+ sparse_mask = sparse_mask & mask
232+ mask = sparse_mask
238233 output = scaled_dot_product_attention (
239- queries , keys , values , cache = cache , scale = self .scale , mask = mask
234+ queries , keys , values , cache = cache [ 0 ] , scale = self .scale , mask = mask
240235 )
241236 output = output .transpose (0 , 2 , 1 , 3 ).reshape (B , L , - 1 )
242237 return self .o_proj (output )
0 commit comments