
    GJjT                         d dl Z d dlmZ d dlmZ  G d dej                        Z G d de      Z G d de      Z	 G d d	ej                        Z
 G d
 de
      Zy)    Nc                        e Zd Z	 ddedef fdZ	 	 	 ddej                  dej                  dej                  dej                  dz  d	ej                  dz  d
ej                  fdZ xZS )MultiHeadAttentionn_headn_featc                 T   t         |           || _        ||z  | _        | j                  dz  | _        t        j                  |||      | _        t        j                  |||      | _        t        j                  |||      | _	        t        j                  |||      | _
        y )Ng      ࿩bias)super__init__r   head_dimscalennLinearlinear_qlinear_klinear_v
linear_out)selfr   r   r	   	__class__s       `/Users/ahmed/devFolder/claude-voice/.venv/lib/python3.12/site-packages/parakeet_mlx/attention.pyr   zMultiHeadAttention.__init__   s     	&(]]D(
		&&t<		&&t<		&&t<))FF>    Nqkvpos_embmaskreturnc                 (   | j                  |      | j                  |      | j                  |      }}}|j                  \  }}}	|j                  \  }	}
}	|j	                  ||| j
                  | j                        j                  dddd      }|j	                  ||
| j
                  | j                        j                  dddd      }|j	                  ||
| j
                  | j                        j                  dddd      }|r|j                  ||      \  }}t        j                  j                  |||| j                  |      }|j                  dddd      j	                  ||| j                  | j
                  z        }| j                  |      S )Nr            r   r   )r   r   r   shapereshaper   r   	transposeupdate_and_fetch_kvmxfastscaled_dot_product_attentionr   r   )r   r   r   r   r   r   cachebatchq_seq_k_seqos               r   __call__zMultiHeadAttention.__call__   sS    --"DMM!$4dmmA6Fa1''uagg5!IIeUDKK?II!QPQSTUIIeUDKK?II!QPQSTUIIeUDKK?II!QPQSTU,,Q2DAqGG00Aq

QU0VKK1a#++E5$--$++:UVq!!r   )TNNN)	__name__
__module____qualname__intr   r'   arrayr0   __classcell__r   s   @r   r   r      s    
 	?? ?, $( $"88" 88" 88	"
 D" hho" 
"r   r   c                   T    e Zd Z	 	 	 ddedededej                  dz  dej                  dz  f
 fdZdej                  d	ej                  fd
Z	 	 	 ddej                  dej                  dej                  dej                  dz  dej                  dz  d	ej                  fdZ	 xZ
S )RelPositionMultiHeadAttentionNr   r   r	   
pos_bias_u
pos_bias_vc                    t         |   |||       t        j                  ||d      | _        |1t        j                  | j                  | j                  f      | _	        n|| _	        |1t        j                  | j                  | j                  f      | _
        n|| _
        | j                  | _        | j                  | _        y )N)r   r   r	   Fr   )r
   r   r   r   
linear_posr'   zerosr   r   _pos_bias_u_init_pos_bias_v_initr;   r<   )r   r   r   r	   r;   r<   r   s         r   r   z&RelPositionMultiHeadAttention.__init__5   s     	 	 	
 ))FF?$&HHdkk4==-I$JD!$.D!$&HHdkk4==-I$JD!$.D!////r   xr   c                     |j                   \  }}}}dg|j                  dz
  z  dgz   }t        j                  ||      }|j	                  |||dz   |      }|d d d d dd d d f   }|j	                  ||||      }|S )Nr   r   r    )r    r   )r#   ndimr'   padr$   )r   rB   BHTqpos_lenpaddings          r   	rel_shiftz'RelPositionMultiHeadAttention.rel_shiftR   s    GG1b'(affqj)VH4FF1gIIaGaK,aABkNIIaB(r   r   r   r   r   r   c                    |t        d      | j                  |      | j                  |      | j                  |      }}}| j	                  |      }|j
                  \  }}	}
|j
                  \  }
}}
|j
                  \  }}}
|dk(  r,|dkD  r't        j                  ||||j
                  d   f      }n||k7  rt        d| d| d      |j                  ||	| j                  | j                        }|| j                  z   j                  dddd	      }|| j                  z   j                  dddd	      }|j                  ||| j                  | j                        j                  dddd	      }|j                  ||| j                  | j                        j                  dddd	      }|j                  ||| j                  | j                        j                  dddd	      }||j                  ||      \  }}t        j                  ||j!                  d
d            }| j#                  |      }|d d d d d d d |j
                  d
   f   | j$                  z  }|*t        j&                  |d      }t        j(                   ||<   t        j*                  j-                  |||| j$                  |      }|j                  dddd	      j                  ||	d      }| j/                  |      S )Npos_emb is necessary!r    pos_emb batch (") must be 1 or match query batch ()r   r   r!   r"   )
ValueErrorr   r   r   r>   r#   r'   broadcast_tor$   r   r   r;   r%   r<   r&   matmulswapaxesrL   r   expand_dimsinfr(   r)   r   )r   r   r   r   r   r   r*   pr+   r,   r-   r.   p_batchrJ   q_uq_v	matrix_bdr/   s                     r   r0   z&RelPositionMultiHeadAttention.__call__]   s    ?455--"DMM!$4dmmA6Fa1OOG$''uagg5!gg!a<EAIE7AGGBK#@AA!'*LUGSTU  IIeUDKK?4??"--aAq94??"--aAq9IIeUDKK?II!QPQSTUIIeUDKK?II!QPQSTUIIeWdkk4==AKKAqRSUVW,,Q2DAqIIc1::b"#56	NN9-	aA}}45

B	>>$*D!vvgIdOGG00ATZZi 1 
 KK1a#++E5"=q!!r   )TNNr1   )r2   r3   r4   r5   boolr'   r6   r   rL   r0   r7   r8   s   @r   r:   r:   4   s    
 &*&*00 0 	0
 HHtO0 HHtO0:	288 	 	  $( $3"883" 883" 88	3"
 D3" hho3" 
3"r   r:   c                       e Zd Z	 	 	 	 ddedededej                  dz  dej                  dz  deeef   f fdZ	 	 	 dd	ej                  d
ej                  dej                  dej                  dz  dej                  dz  dej                  fdZ	d	ej                  d
ej                  dedej                  fdZ
dej                  dej                  dedej                  fdZ xZS )"RelPositionMultiHeadLocalAttentionNr   r   r	   r;   r<   context_sizec                 l    t         |   |||||       || _        t        |      dk  rt	        d      y )Nr   z@Context size for RelPositionMultiHeadLocalAttention must be > 0.)r
   r   rb   minrT   )r   r   r   r	   r;   r<   rb   r   s          r   r   z+RelPositionMultiHeadLocalAttention.__init__   sD     	z:F(|!R  "r   r   r   r   r   r   r   c           	      
   |t        d      |2t        j                  |j                  d d t        j                        }| j                  |      | j                  |      | j                  |      }}}| j                  |      }|j                  \  }}	}
|j                  \  }
}}
|j                  \  }}}
|dk(  r,|dkD  r't        j                  ||||j                  d   f      }n||k7  rt        d| d| d      |j                  ||	| j                  | j                        j                  d	ddd
      }|j                  ||| j                  | j                        j                  d	ddd
      }|j                  ||| j                  | j                        j                  d	ddd
      }|j                  ||| j                  | j                        j                  d	ddd
      }||j                  ||      \  }}t        | j                         }d|z  |j                  d   d|z  z  z
  d|z  z  }t        j"                  |ddd	|fdf      }t        j"                  |ddd	|fdf      }t        j"                  |ddd	|fdf      }t        j"                  |dd	|ffd      }|t        j$                  | j&                  d      z   }|t        j$                  | j(                  d      z   }| j+                  |||      }t        j,                  ||j/                  dd            }|d d d d d d d | j                   d	   f   |d d d d d d d | j                   d	   f   z   |d d d d d d d | j                   d	   f<   |d d d d d d | j                   d   dz    d f   |d d d d d d | j                   d	   d f   z   |d d d d d d | j                   d   dz    d f<   t        j0                   |d d d d d d d || j                   d	   z
  f<   t        j0                   |d d d d d d || j                   d   z   dz   d f<   || j2                  z  }t        j$                  t        j$                  |d      d      }t        j4                  |t        j0                   d      j7                  |j8                        }t        j:                  |      }| j+                  |||      }||z   }t        j<                  |d      }t        j4                  |d	|      }| j?                  |||      }|j                  |d| j                  | j                  z        d d d |	f   }| jA                  |      S )NrN   r   dtyper    rO   rP   rQ   rR   r   r!   rD   T)constant_valuesrS   g        )!rT   r'   r?   r#   bool_r   r   r   r>   rU   r$   r   r   r%   r&   maxrb   rF   rX   r;   r<   	matmul_qkrV   rW   rY   r   whereastyperg   	ones_likesoftmax	matmul_pvr   )r   r   r   r   r   r   r*   rZ   r+   r,   r-   r.   r[   rJ   wpad_lenr\   r]   	matrix_acr^   scores
float_maskonesd_maskattnouts                             r   r0   z+RelPositionMultiHeadLocalAttention.__call__   s    ?455<88QWWRa[:D--"DMM!$4dmmA6Fa1OOG$''uagg5!gg!a<EAIE7AGGBK#@AA!'*LUGSTU  IIeUDKK?II!QPQSTUIIeUDKK?II!QPQSTUIIeUDKK?II!QPQSTUIIeWdkk4==AKKAqRSUVW,,Q2DAq !!"q51771:Q//AE:FF1vv7|V<=FF1vv7|V<=FF1vv7|V<=vvdVa\2DI"..!44"..!44NN31-	IIc1::b"#56	 aA5!2!21!55561a!74#4#4Q#7!7789 	!Q1T..q1112
 aA!2!21!5!9:<<=1a!2!21!5!7789 	!QT..q1A56889 =?FF7	!Q7a$"3"3A"66778@Bw	!QA 1 1! 44q8;;<TZZ'~~bnnT15r:XXdRVVGS188I
||J'j!4&zz&"%xxa&nnT1a(kk%T[[4==%@A!VeV)Ls##r   rq   c                    d}|j                   \  }}}}|j                   \  }	}	}
}	|||d|z  dz   f}||z  |z  }d|z  dz   }d}t        j                  j                  dddgdg|      }t	        d|      }t	        d|      }|d	k\  rt        |d
      }t        |d	      }nT|dk\  rt        |d      }t        |d      }n6|dk\  rt        |d      }t        |d      }nt        |d      }t        |d      }|dkD  rd}n$|dkD  rd}n|dkD  rd}n|d
kD  rd}nt	        |d      }t	        |d      }t	        |d      } |||gd|j                  fd|fd|fg|||f||df|g|j                  g      }|d   S )Na  
        // D, W are provided as constant
        uint B = q_shape[0];
        uint H = q_shape[1];
        uint S_q = q_shape[2];
        uint S_k = k_shape[2];
        uint K_rel = 2 * W + 1;

        uint target_idx = thread_position_in_grid.x;
        uint k_rel_idx = thread_position_in_grid.y;

        if (target_idx >= B * H * S_q) return;

        uint s_q_idx = target_idx % S_q;
        uint remaining_idx = target_idx / S_q;
        uint h_idx = remaining_idx % H;
        uint b_idx = remaining_idx / H;
        uint k_offset = k_rel_idx;

        uint stick_q_k_idx = S_k - S_q + s_q_idx;
        // stick to right (assuming S_k >= S_q)

        int s_k_idx_signed = int(stick_q_k_idx) + int(k_offset) - int(W);
        bool is_out_of_bounds = (s_k_idx_signed < 0) || (s_k_idx_signed >= S_k);

        T result;

        if (!is_out_of_bounds) {
            uint s_k_idx = uint(s_k_idx_signed);

            // q[b, h, s_q, d]
            uint Q_D_stride = D;
            uint Q_S_stride = S_q * Q_D_stride;
            uint Q_H_stride = H * Q_S_stride;
            // k[b, h, s_k, d]
            uint K_D_stride = D;
            uint K_S_stride = S_k * K_D_stride;
            uint K_H_stride = H * K_S_stride;

            uint q_base_offset =
                b_idx * Q_H_stride + h_idx * Q_S_stride + s_q_idx * Q_D_stride;
            uint k_base_offset =
                b_idx * K_H_stride + h_idx * K_S_stride + s_k_idx * K_D_stride;

            const device T* q_vec_ptr = q + q_base_offset;
            const device T* k_vec_ptr = k + k_base_offset;

            result = T(0.0);
            uint d_idx = 0;

            // hand unrolling
            for (; d_idx + 16 <= D; d_idx += 16) {
                T q_vals[16], k_vals[16];

                for (uint i = 0; i < 16; ++i) {
                    q_vals[i] = q_vec_ptr[d_idx + i];
                    k_vals[i] = k_vec_ptr[d_idx + i];
                }

                result +=
                    q_vals[0] * k_vals[0] + q_vals[1] * k_vals[1] +
                    q_vals[2] * k_vals[2] + q_vals[3] * k_vals[3] +
                    q_vals[4] * k_vals[4] + q_vals[5] * k_vals[5] +
                    q_vals[6] * k_vals[6] + q_vals[7] * k_vals[7] +
                    q_vals[8] * k_vals[8] + q_vals[9] * k_vals[9] +
                    q_vals[10] * k_vals[10] + q_vals[11] * k_vals[11] +
                    q_vals[12] * k_vals[12] + q_vals[13] * k_vals[13] +
                    q_vals[14] * k_vals[14] + q_vals[15] * k_vals[15];
            }

            for (; d_idx + 8 <= D; d_idx += 8) {
                result +=
                    q_vec_ptr[d_idx] * k_vec_ptr[d_idx] +
                    q_vec_ptr[d_idx + 1] * k_vec_ptr[d_idx + 1] +
                    q_vec_ptr[d_idx + 2] * k_vec_ptr[d_idx + 2] +
                    q_vec_ptr[d_idx + 3] * k_vec_ptr[d_idx + 3] +
                    q_vec_ptr[d_idx + 4] * k_vec_ptr[d_idx + 4] +
                    q_vec_ptr[d_idx + 5] * k_vec_ptr[d_idx + 5] +
                    q_vec_ptr[d_idx + 6] * k_vec_ptr[d_idx + 6] +
                    q_vec_ptr[d_idx + 7] * k_vec_ptr[d_idx + 7];
            }

            for (; d_idx + 4 <= D; d_idx += 4) {
                result +=
                    q_vec_ptr[d_idx] * k_vec_ptr[d_idx] +
                    q_vec_ptr[d_idx + 1] * k_vec_ptr[d_idx + 1] +
                    q_vec_ptr[d_idx + 2] * k_vec_ptr[d_idx + 2] +
                    q_vec_ptr[d_idx + 3] * k_vec_ptr[d_idx + 3];
            }

            for (; d_idx < D; ++d_idx) {
                result += q_vec_ptr[d_idx] * k_vec_ptr[d_idx];
            }
        } else {
            result = T(-INFINITY);
        }

        uint out_idx = target_idx * K_rel + k_rel_idx;
        out[out_idx] = result;
        r   r    local_qk_perfr   r   ry   nameinput_namesoutput_namessource                   @   TWDinputstemplategridthreadgroupoutput_shapesoutput_dtypesr   )r#   r'   r(   metal_kernelrj   rd   rg   )r   r   r   rq   KERNELrG   rH   S_qr   r-   S_koutput_shape
grid_dim_x
grid_dim_y
grid_dim_z	kernel_fntg_ytg_xoutputss                      r   rk   z,RelPositionMultiHeadLocalAttention.matmul_qk   s   cJ ww1c1ww1c11c1q519-US[
UQY

GG(( c
	 ) 
	 J'
J'
8z1%Dz3'D#Xz1%Dz3'D"Wz2&Dz2&Dz2&Dz2&D"9DBYDAXDAXDtQ<D4|4|q6aggaa
 j*5tQ'.77)
 qzr   probc                    d}|j                   \  }}}}|j                   \  }	}	}
}t        j                  j                  dddgdg|      }||||f}|}|}||z  }t	        |d      }t	        |d|z        }t        |d	      }t        |d	      } |||gd
|j                  fd|fd|fd|fg|||f||d	f|g|j                  g      }|d   S )Na  
        // D, W, D_v are provided as constant
        uint B = prob_shape[0];
        uint H = prob_shape[1];
        uint S_p = prob_shape[2];
        uint S_v = v_shape[2];
        uint K_rel = 2 * W + 1;

        uint d_idx = thread_position_in_grid.x;
        uint s_p_idx = thread_position_in_grid.y;
        uint bh_idx = thread_position_in_grid.z;  // merged

        if (d_idx >= D_v || s_p_idx >= S_p || bh_idx >= (B * H)) {
            return;
        }

        uint b_idx = bh_idx / H;
        uint h_idx = bh_idx % H;

        T current_sum = 0.0f;

        // p[b, h, s_p, k_rel]
        uint P_H_stride = S_p * K_rel;
        uint P_B_stride = H * P_H_stride;

        // v[b, h, s_v, d]
        uint V_H_stride = S_v * D_v;
        uint V_B_stride = H * V_H_stride;

        // out[b, s_p, h, d]
        uint O_S_stride = D_v * H;
        uint O_B_stride = S_p * O_S_stride;

        uint stick_p_v_idx = S_v - S_p + s_p_idx;
        // stick to right (assuming S_v >= S_p)

        uint k = 0;
        // hand unrolling
        for (; k + 16 <= K_rel; k += 16) {
            float prob_vals[16], v_vals[16];
            int s_v_indices[16];
            bool valid[16];

            for (uint i = 0; i < 16; ++i) {
                s_v_indices[i] = int(stick_p_v_idx) + int(k + i) - int(W);
                valid[i] = (s_v_indices[i] >= 0 && s_v_indices[i] < S_v);
                if (valid[i]) {
                    uint prob_idx = b_idx * P_B_stride + h_idx * P_H_stride + s_p_idx * K_rel + (k + i);
                    uint v_idx = b_idx * V_B_stride + h_idx * V_H_stride + uint(s_v_indices[i]) * D_v + d_idx;
                    prob_vals[i] = prob[prob_idx];
                    v_vals[i] = v[v_idx];
                } else {
                    prob_vals[i] = 0.0f;
                    v_vals[i] = 0.0f;
                }
            }

            current_sum +=
                prob_vals[0] * v_vals[0] + prob_vals[1] * v_vals[1] +
                prob_vals[2] * v_vals[2] + prob_vals[3] * v_vals[3] +
                prob_vals[4] * v_vals[4] + prob_vals[5] * v_vals[5] +
                prob_vals[6] * v_vals[6] + prob_vals[7] * v_vals[7] +
                prob_vals[8] * v_vals[8] + prob_vals[9] * v_vals[9] +
                prob_vals[10] * v_vals[10] + prob_vals[11] * v_vals[11] +
                prob_vals[12] * v_vals[12] + prob_vals[13] * v_vals[13] +
                prob_vals[14] * v_vals[14] + prob_vals[15] * v_vals[15];
        }

        for (; k + 8 <= K_rel; k += 8) {
            for (uint i = 0; i < 8; ++i) {
                int s_v_idx_signed = int(stick_p_v_idx) + int(k + i) - int(W);
                if (s_v_idx_signed >= 0 && s_v_idx_signed < S_v) {
                    uint s_v_idx = uint(s_v_idx_signed);
                    uint prob_idx = b_idx * P_B_stride + h_idx * P_H_stride + s_p_idx * K_rel + (k + i);
                    uint v_idx = b_idx * V_B_stride + h_idx * V_H_stride + s_v_idx * D_v + d_idx;
                    current_sum += prob[prob_idx] * v[v_idx];
                }
            }
        }

        for (; k + 4 <= K_rel; k += 4) {
            for (uint i = 0; i < 4; ++i) {
                int s_v_idx_signed = int(stick_p_v_idx) + int(k + i) - int(W);
                if (s_v_idx_signed >= 0 && s_v_idx_signed < S_v) {
                    uint s_v_idx = uint(s_v_idx_signed);
                    uint prob_idx = b_idx * P_B_stride + h_idx * P_H_stride + s_p_idx * K_rel + (k + i);
                    uint v_idx = b_idx * V_B_stride + h_idx * V_H_stride + s_v_idx * D_v + d_idx;
                    current_sum += prob[prob_idx] * v[v_idx];
                }
            }
        }

        for (; k < K_rel; ++k) {
            int s_v_idx_signed = int(stick_p_v_idx) + int(k) - int(W);
            if (s_v_idx_signed >= 0 && s_v_idx_signed < S_v) {
                uint s_v_idx = uint(s_v_idx_signed);
                uint prob_idx = b_idx * P_B_stride + h_idx * P_H_stride + s_p_idx * K_rel + k;
                uint v_idx = b_idx * V_B_stride + h_idx * V_H_stride + s_v_idx * D_v + d_idx;
                current_sum += prob[prob_idx] * v[v_idx];
            }
        }

        uint out_idx =
            b_idx * O_B_stride + s_p_idx * O_S_stride + h_idx * D_v + d_idx;

        context_out[out_idx] = current_sum;
        local_pv_matmulr   r   context_outr|   r   i   r    r   r   r   D_vr   r   )r#   r'   r(   r   rd   rj   rg   )r   r   r   rq   r   rG   rH   S_pK_relr-   S_vr   r   r   r   r   r   r   r   r   s                       r   rp   z,RelPositionMultiHeadLocalAttention.matmul_pv  s   jX  ::1c51c3GG(("'	 ) 
	 33'

U
:r":tt|,4|4|!9DJJ'#qC<%Nj*5tQ'.::,
 qzr   )TNNr   r   r1   )r2   r3   r4   r5   r_   r'   r6   tupler   r0   rk   rp   r7   r8   s   @r   ra   ra      s?   
 &*&*(2  	
 HHtO HHtO CHo. $( $O$88O$ 88O$ 88	O$
 DO$ hhoO$ 
O$b`288 ` `S `RXX `DKbhh K288 K K Kr   ra   c            	            e Zd Z	 	 d
dededef fdZd Zddej                  dede	ej                  ej                  f   fd	Z
 xZS )RelPositionalEncodingd_modelmax_lenscale_inputc                     |dz  dk(  r|dkD  sJ t         |           || _        || _        |rt	        j
                  | j                        nd| _        | j                          y )Nr   r   g      ?)r
   r   r   r   mathsqrtr   calculate_pe)r   r   r   r   r   s       r   r   zRelPositionalEncoding.__init__(  sZ     {aGaK//0;TYYt||,
r   c                 t   t        j                  | j                  dz
  | j                   dt         j                        }t        j                  |d      j                  t         j                        }t        j                  t        j                  d| j                  dt         j                        t        j                  d      | j                  z   z        }t        j                  d| j                  z  dz
  | j                  ft         j                        }t        j                  ||z        |d d dd df<   t        j                  ||z        |d d dd df<   t        j                  |d      j                  t         j                        | _        t        j                  | j                         y Nr    rO   rf   )axisr   r   g     @)r'   aranger   int32rX   rm   float32expr   r   logr?   sincos_peevalr   	positionsdiv_termpes       r   r   z"RelPositionalEncoding.calculate_pe6  s-   IIdllQ.rR	NN915<<RZZH	66IIaq

;!DLL012
 XXq4<<'!+T\\:"**MffY121add7ffY121add7>>"1-44RZZ@
r   rB   offsetr   c                 ^   |j                   d   |z   }|| j                  kD  r|dz   | _        | j                          || j                  z  }| j                  j                   d   }|dz  |dz
  z
  }|dz  |dz
  z   dz   }| j                  d d ||f   j                  |j                        }||fS )Nr    r   )r#   r   r   r   r   rm   rg   )r   rB   r   	input_len
buffer_len	start_idxend_idxr   s           r   r0   zRelPositionalEncoding.__call__G  s    GGAJ'	t||#$q=DL

NXX^^A&
!Oy1}5	/Y]3a7((1i//077@'zr   )  Tr   )r2   r3   r4   r5   r_   r   r   r'   r6   r   r0   r7   r8   s   @r   r   r   '  sa      	  	""(( C bhh>P8Q r   r   c                        e Zd Z	 	 	 ddedededeeef   f fdZd Zddej                  ded	eej                  ej                  f   fd
Z
 xZS )LocalRelPositionalEncodingr   r   r   rb   c                 F    |\  | _         | _        t        |   |||       y )N)left_contextright_contextr
   r   )r   r   r   r   rb   r   s        r   r   z#LocalRelPositionalEncoding.__init__Z  s&     1=-4-';7r   c                    t        j                  | j                  | j                   dz
  dt         j                        }t        j
                  |d      j                  t         j                        }t        j                  t        j                  d| j                  dt         j                        t        j                  d      | j                  z   z        }t        j                  | j                  | j                  z   dz   | j                  ft         j                        }t        j                  ||z        |d d dd df<   t        j                  ||z        |d d dd df<   t        j
                  |d      j                  t         j                        | _        t        j                   | j                         y r   )r'   r   r   r   r   rX   rm   r   r   r   r   r   r?   r   r   r   r   r   s       r   r   z'LocalRelPositionalEncoding.calculate_pee  sD   II 2 22Q6"((
	 NN915<<RZZH	66IIaq

;!DLL012
 XX!3!33a7Fbjj
 ffY121add7ffY121add7>>"1-44RZZ@
r   rB   r   r   c                     || j                   z  }| j                  | j                  z   dz   }| j                  d d d |f   j	                  |j
                        }||fS )Nr    )r   r   r   r   rm   rg   )r   rB   r   r   r   s        r   r0   z#LocalRelPositionalEncoding.__call__z  sY    

N##d&8&881<((1hwh;'..qww7'zr   )r   Tr   r   )r2   r3   r4   r5   r_   r   r   r   r'   r6   r0   r7   r8   s   @r   r   r   Y  sv      (2	8	8 	8 		8
 CHo	8*"(( C bhh>P8Q r   r   )r   mlx.corecorer'   mlx.nnr   Moduler   r:   ra   r   r    r   r   <module>r      s]      *" *"Z\"$6 \"~Q)F Qh/BII /d'!6 'r   