
    (HJj                        d dl mZmZ d dlmZ d dlmZ ej                  d        Z	d Z
 e
       Zdej                  dej                  dej                  dej                  d	ej                  d
ej                  dej                  dej                  deeef   fdZd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ef   deej                     deej                     de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ef   deej                     deej                     fdZy)    )OptionalTupleNc                     | j                  t        j                        } t        j                  | |z         } t        j
                  | |d   |d         S )Nr      )astypemxfloat32nnsoftplusclip)dtdt_biastime_step_limits      [/Users/ahmed/devFolder/claude-voice/.venv/lib/python3.12/site-packages/mlx_lm/models/ssm.py
compute_dtr      sC    	2::	B	R'\	"B772q)?1+=>>    c                      t         j                  j                         sy d} t         j                  j	                  dg dddg|       S )Na  
        auto n = thread_position_in_grid.z;
        auto h_idx = n % H;
        auto g_idx = n / G;
        constexpr int n_per_t = Ds / 32;

        auto x = X + n * Dh;
        out += n * Dh;
        auto i_state = state_in + n * Dh * Ds;
        auto o_state = state_out + n * Dh * Ds;

        // C and B have shape [batch, group, state_dim]
        // C and B need to be offset by group size
        auto C_ = C + g_idx * Ds;
        auto B_ = B + g_idx * Ds;

        auto ds_idx = thread_position_in_threadgroup.x;
        auto d_idx = thread_position_in_grid.y;

        auto dt_ = static_cast<float>(dt[n]);
        auto A = -fast::exp(static_cast<float>(A_log[h_idx]));
        auto dA = fast::exp(A * dt_);

        float acc = 0.0;
        auto x_ = static_cast<float>(x[d_idx]);

        for (int i = 0; i < n_per_t; ++i) {
            auto s_idx = n_per_t * ds_idx + i;
            auto idx = d_idx * Ds + s_idx;
            auto dB_by_x = x_ * dt_ * static_cast<float>(B_[s_idx]);
            auto state = dA * i_state[idx] + dB_by_x;
            o_state[idx] = static_cast<U>(state);
            acc += state * C_[s_idx];
        }
        acc = simd_sum(acc);
        if (thread_index_in_simdgroup == 0) {
            out[d_idx] = static_cast<T>(acc + x_ * D[h_idx]);
        }
    
ssm_kernel)XA_logBCDr   state_inout	state_out)nameinput_namesoutput_namessource)r   metalis_availablefastmetal_kernel)r    s    r   make_ssm_kernelr%      sL    88  "&FN 77C[)	    r   hidden_statesr   r   r   r   r   r   stater   c	                    | j                   \  }	}
}}| j                  }|j                  }|j                   dd  \  }}t        |||      }t        | ||||||gd|fd|fd|fd|fd|fd||z  fgd|||	z  fd	|	d
||f|j                   g||g      S )NTUDhDsHG    )r0      r   r   )inputstemplategridthreadgroupoutput_shapesoutput_dtypes)shapedtyper   _ssm_kernel)r&   r   r   r   r   r   r   r'   r   n_hd
input_type
state_typehbdss                    r   ssm_update_kernelrC   C   s     $$JAq!Q$$JJWWRS\FB	B	1BuaAr59**1I2J!H!r'N
 !QU^1a|U[[1!:. r   c                 P   | j                   d   }|t        j                  |d      }| |z  } t        j                  | d   |d      } t        j                  | d      } t        j
                  | d      }|/t        j                  |dd d d f   |d   z  |t        d             }|S )Nr   .Naxisr)   .inf)r8   r   expand_dimsrepeattrilcumsumwherefloat)xmasklx_segsums       r   segsumrT   d   s    	A~~dA&H
		!I,+A
2Ayy$H88dAi0(U5\M
 Or   rP   rQ   lengthsstepreturnc                    
  j                   \  }|j                   \  }}t        |||      }z  t        j                  |      j	                  |j
                         }||j                  ddd      z  }|j                  |d       z  }
 f	d}g }t        d|      D ]h  } ||dd||z   f   |dd||z   f   |dd||z   f   |dd||z   f   ||	dn|	d||z   f         \  }}

z
  
|j                  |       j t        j                  |d       |j                  ddd      z  z   }||fS )a5  SSD-SSM forward pass.

    Args:
        x: Input of shape (batch_size, seq_len, num_heads, head_dim).
        dt: Time deltas of shape (seq_len, num_heads,).
        A_log: State transition of shape (num_heads,).
        B: Input mixing of shape (batch_size, seq_len, num_groups, n).
        C: Output mixing of shape (batch_size, seq_len, num_groups, n).
        D: Residual connection.
        dt_bias: Bias for time deltas of shape (num_heads,).
        time_step_limit: Minimum and maximum value for time deltas.
        mask: Optional multiplicative mask.
        lengths: Optional lenghts of sequences, assumed to be the full length if unspecified.
        step: Step size for processing x.

    Code modified from
    https://github.com/cartesia-ai/edge/blob/main/cartesia-mlx/cartesia_mlx/layers/ssd/ops.py

    r   rE   c                 "  	 | j                   d   }t        j                  |d      }t        j                  |dd      |z  }t        j                  |d      }t        j
                  t        |j                  dd      |            }t        j                  ||z  d      }	|	| j                  dd      z  }
t        j                  |
dd      }
\t        j                  t        j                        dz
  d      }t        j                  |d      }t        j                  ||d      }n|d d d d dd d d f   }|j                  dd	dd      }t        j                  |z  d      j                  dd	      }| |z  }|j                  dd      j                  dd	      }||z  }|t        j
                  t        j                  |d
            }||d d dd d d d f   |z  z  }|j                  |dd      }|j                  df      |z  j                  d      j                  dd	      }|
|d   |z  z  }
0|.t        j                   t        j                  dk  d      ||      }|
j#                  j$                        |fS )Nr   )r         r   rZ   rG   )rQ   r   )r   rZ   r[   rE   r[   r)   rF   )r8   r   	transposeswapaxesrK   exprT   rL   maximumminimumrJ   take_along_axisrM   reshapesqueezeflattenrN   r   r9   )dtxdtAr   r   r'   rQ   sCBdecaysurrogate_attention_matrixyposdtxdecay
next_stateexp_dtA_cumsumy_prevbr>   dhgr=   rU   repeatsrV   rP   s                   r   _stepzssm_attn.<locals>._step   s]   IIaLLLL)[[Aq!A%YYr7+vcll1a0t<=%'WWR%Z%;"&a);;KK1a **RZZ6:A>C..i0C&&uc:E!QQ,'E1a+IIaaa(11!Q7;$$Q*33Aq9\
VVBIIc$;<N.B4)=>FFJ		!Q1a+A1a"a89A=FFrJRRSTVWX  	*V33A5#4w{I6zJ xx *,,r   r   N.rG   )
r8   r   r   r^   r   r9   rb   rangeappendconcatenate)rP   r   r   r   r   r   r   r'   r   rQ   rU   rV   rR   r<   Arf   re   ru   ysirk   rq   r>   rr   rs   r=   rt   s   `         ``         @@@@@@r   ssm_attnr|   s   s   B ''KAq!RJAq!Q	B	1B1fG			bhh	''A
qyyAr"
"C
**Q1a
 1
$C)- )-V 
B1a1q4x< 1q4x< aQXoaQXoLDd3AH+<&=
5 nG
		!  	r"Q1aA)>%>>Ae8Or   c                    | j                   d   }|dkD  sE|Ct        j                         t        j                  k7  st        j                  j                         st        | |||||||||	|
      S t        | ||||||||	      S )Nr   )rQ   rU   )r8   r   default_devicegpur!   r"   r|   rC   )r&   r   r   r   r   r   r   r'   r   rQ   rU   seq_lens               r   
ssm_updater      s     !!!$G!="&&(xx$$&
 	
 !

 
	
r   )N)NgMbP?g      Y@NN   )Nr   NN)typingr   r   mlx.corecorer   mlx.nnr
   compiler   r%   r:   arrayrO   rC   rT   intr|   r    r   r   <module>r      sr   "   ? ?/d 8888 
xx 
xx	
 
xx 	 XX 88 5%<(B. !%+9#"&c	xxc88c 
xxc 
xx	c
 
xxc 	c XXc BHHc 5%<(c 288
c bhhc c 288RXXc\ !%+9#"&,
88,
88,
 
xx,
 
xx	,

 
xx,
 	,
 XX,
 BHH,
 5%<(,
 288
,
 bhh,
r   