
    (HJj                        d dl mZ d dlmZmZ d dlmZ d dlm	Z	  eej                  d      d        ZddZ edd	      Z edd	      Z edd	      Z edd	      Zej                  	 dd
ej"                  dej"                  dej"                  dej"                  dej"                  dej"                  deej"                     deej"                  ej"                  f   fd       Z	 dd
ej"                  dej"                  dej"                  dej"                  dej"                  dej"                  deej"                     deej"                  ej"                  f   fdZ	 	 dd
ej"                  dej"                  dej"                  dej"                  dej"                  deej"                     deej"                     deej"                  ej"                  f   fdZ	 	 	 dd
ej"                  dej"                  dej"                  dej"                  dej"                  dej"                  dej"                  deej"                     deej"                     dedeej"                  ej"                  f   fdZy)    )partial)OptionalTupleNT)	shapelessc                     t        j                  t        j                  | j                  t         j                               t	        j
                  ||z         z        S N)mxexpastypefloat32nnsoftplus)A_logadt_biass      c/Users/ahmed/devFolder/claude-voice/.venv/lib/python3.12/site-packages/mlx_lm/models/gated_delta.py	compute_gr      s<    66266%,,rzz233bkk!g+6NNOO    Fc                 4   t         j                  j                         sy | rdnd}|r	d}d}d}d}nd}d}d	}d
}d| d| d| d| d| d}g d}| r|j                  d       d}	|r|	dz  }	| r|	dz  }	t         j                  j                  d|	 |ddg|      S )Nzmask[b_idx * T + t]truez// g: [B, T, Hv, Dk]z-auto g_ = g + (b_idx * T * Hv + hv_idx) * Dk;z	g_[s_idx]zg_ += Hv * Dk;z// g: [B, T, Hv]zauto g_ = g + b_idx * T * Hv;z
g_[hv_idx]z	g_ += Hv;a  
        auto n = thread_position_in_grid.z;
        auto b_idx = n / Hv;
        auto hv_idx = n % Hv;
        auto hk_idx = hv_idx / (Hv / Hk);
        constexpr int n_per_t = Dk / 32;

        // q, k: [B, T, Hk, Dk]
        auto q_ = q + b_idx * T * Hk * Dk + hk_idx * Dk;
        auto k_ = k + b_idx * T * Hk * Dk + hk_idx * Dk;

        // v, y: [B, T, Hv, Dv]
        auto v_ = v + b_idx * T * Hv * Dv + hv_idx * Dv;
        y += b_idx * T * Hv * Dv + hv_idx * Dv;

        auto dk_idx = thread_position_in_threadgroup.x;
        auto dv_idx = thread_position_in_grid.y;

        // state_in, state_out: [B, Hv, Dv, Dk]
        auto i_state = state_in + (n * Dv + dv_idx) * Dk;
        auto o_state = state_out + (n * Dv + dv_idx) * Dk;

        float state[n_per_t];
        for (int i = 0; i < n_per_t; ++i) {
          auto s_idx = n_per_t * dk_idx + i;
          state[i] = static_cast<float>(i_state[s_idx]);
        }

        z	
        zb
        auto beta_ = beta + b_idx * T * Hv;

        for (int t = 0; t < T; ++t) {
          if (z) {
            float kv_mem = 0.0f;
            for (int i = 0; i < n_per_t; ++i) {
              auto s_idx = n_per_t * dk_idx + i;
              state[i] = state[i] * a  ;
              kv_mem += state[i] * k_[s_idx];
            }
            kv_mem = simd_sum(kv_mem);

            auto delta = (v_[dv_idx] - kv_mem) * beta_[hv_idx];

            float out = 0.0f;
            for (int i = 0; i < n_per_t; ++i) {
              auto s_idx = n_per_t * dk_idx + i;
              state[i] = state[i] + k_[s_idx] * delta;
              out += state[i] * q_[s_idx];
            }
            out = simd_sum(out);
            if (thread_index_in_simdgroup == 0) {
              y[dv_idx] = static_cast<InT>(out);
            }
          } else {
            y[dv_idx] = static_cast<InT>(0);
          }
          // Increment data pointers to next time step
          q_ += Hk * Dk;
          k_ += Hk * Dk;
          v_ += Hv * Dv;
          y += Hv * Dv;
          z
          beta_ += Hv;
        }
        for (int i = 0; i < n_per_t; ++i) {
          auto s_idx = n_per_t * dk_idx + i;
          o_state[s_idx] = static_cast<StT>(state[i]);
        }
    )qkvgbetastate_inTmask _vec_maskgated_delta_stepy	state_out)nameinput_namesoutput_namessource)r	   metalis_availableappendfastmetal_kernel)
has_mask
vectorizedmask_source	g_commentg_setupg_access	g_advancer(   inputssuffixs
             r   _make_gated_delta_kernelr7      s    88  "+3'K *	A$	&	1		8 
 			  m % &.J /2 + }EFL ;FfF&'77x(;'	    r   )r.   r/   r   r   r   r   r   stater   returnc                    |}|j                   dk(  r|d   }n2|j                   dk(  r|ddddf   }nt        d|j                         ||z  }||ddddf   z  j                  d      }	||	z
  |d	   z  }
||ddddf   |
d	   z  z   }|| ddddf   z  j                  d      }|.t	        j
                  |d
      }t	        j                  |||      }|j                  | j                        |fS )a  
    Ops-based reference implementation for a single recurrent step.

    Shapes:
      - q, k: [B, H, Dk]
      - v: [B, H, Dv]
      - g: [B, H] or [B, H, Dk]
      - beta: [B, H]
      - state: [B, H, Dv, Dk]
    Returns:
      - y: [B, H, Dv]
      - new_state: [B, H, Dv, Dk]
       ).NN   .NzUnsupported gating shape axis).N)   r;   r<   )	ndim
ValueErrorshapesumr	   expand_dimswherer   dtype)r   r   r   r   r   r8   r   	old_statedecaykv_memdeltar#   s               r   _gated_delta_step_opsrL   ~   s   2 Ivv{/"	
1#tQ,4QWWI>??EMEaT1o%***3FZ4	?*EAc4lOeI&666E	3a<	 %%2%.A~~d3ui088AGGe##r   c           
         |j                   \  }}}	}
|j                   dd  \  }}| j                  }|j                  }|j                  dk(  r)t        }| ||||||g}|@t        }|j                  |       n(t        }| ||||||g}|t        }|j                  |        ||d|fd|fd|
fd|fd|	fd|fgd	|||z  fd
||||f|j                   g||g      S )Nr;      InTStTDkDvHkHv    )rU   rN   r@   )r5   templategridthreadgroupoutput_shapesoutput_dtypes)rC   rG   rA   _gated_delta_kernel_vec_gated_delta_kernel_vec_maskedr+   _gated_delta_kernel_gated_delta_kernel_masked)r   r   r   r   r   r8   r   Br   rS   rQ   rT   rR   
input_type
state_typekernelr5   s                    r   gated_delta_kernelrc      s    77LAq"bWWQR[FBJJvv{(Q1dE1-3FMM$$Q1dE1-/FMM$JJ2J2J2J2J
 "a"f1b"~u{{3!:. r   c                    | j                   \  }}}	}
|j                   dd \  }}|)t        j                  ||||
ft        j                        }||	z  x}dkD  r.t        j                  | |d      } t        j                  ||d      }g }t        |      D ]U  }t        | dd|f   |dd|f   |dd|f   |dd|f   |dd|f   ||dn|dd|f         \  }}|j                  |       W t        j                  |d      }||fS )a  
    Ops-based reference implementation for prompt prefill (sequential loop).
    Supports both scalar and vectorized gating.

    Shapes:
      - q, k: [B, T, Hk, Dk]
      - v: [B, T, Hv, Dv]
      - g: [B, T, Hv] (scalar) or [B, T, Hv, Dk] (vectorized)
      - beta: [B, T, Hv]
      - state: [B, Hv, Dv, Dk]
    Returns:
      - y: [B, T, Hv, Dv]
      - state: [B, Hv, Dv, Dk]
    NrG   r@   r>   )	rC   r	   zerosr   repeatrangerL   r+   stack)r   r   r   r   r   r8   r   r_   r   rS   rQ   rT   rR   repeat_factorystr#   s                    r   gated_delta_opsrn      s   . 77LAq"bWWRS\FB}!RR

;r!Q&IIa+IIa+	B1X(adGadGadGadGAJLDd1a4j
5 			!  	!Ae8Or   r   br   r   
use_kernelc
           	         t        j                  |      }
t        |||      }|L| j                  \  }}}}|j                  dd  \  }}t        j                  ||||ft         j
                        }|	rCt        j                         t         j                  k7  st         j                  j                         st        | ||||
||      S t        | ||||
||      S )Nre   rf   )r	   sigmoidr   rC   rg   r   default_devicegpur)   r*   rn   rc   )r   r   r   r   ro   r   r   r8   r   rp   r   r   r_   _rS   rQ   rT   rR   s                     r   gated_delta_updaterv     s     ::a=D%G$A}ww1b"B!RR

;**,6bhh>S>S>Uq!Q4==aAq$t<<r   )FFr   )NN)NNT)	functoolsr   typingr   r   mlx.corecorer	   mlx.nnr   compiler   r7   r]   r^   r[   r\   arrayrL   rc   rn   boolrv    r   r   <module>r      s    "   	t$P %PfR /%P 5tPUV 2EdS !9d" 
   $)$	xx)$	xx)$ 
xx)$ 
xx	)$
 (()$ 88)$ 288
)$ 288RXX)$ )$f  $(	xx(	xx( 
xx( 
xx	(
 ((( 88( 288
( 288RXX(b !%#-	xx-	xx- 
xx- 
xx	-
 ((- BHH- 288
- 288RXX-p !%#=	xx=	xx= 
xx= 
xx	=
 
xx= 88= XX= BHH= 288
= = 288RXX=r   