
    (HJjx                        d Z ddlZddlZg dZddlmZ ddlmZ ddlm	Z
 ddlZ ed      dLd       Z ed      dLd       Z ed      dLd	       Z ed      dLd
       ZeeeeedZdej$                  dededdfdZdedededededeej$                  ej$                  f   fdZdej$                  dej$                  dej$                  dej$                  fdZdej$                  dej$                  dej$                  dej$                  fdZdej$                  dedej$                  fdZ	 	 dMdej$                  dedededef
dZdej$                  dededej$                  fdZdej$                  d edej$                  fd!Z	 	 	 	 	 	 dNd"e
j>                  ez  fd#Z 	 	 	 	 	 	 dOd$Z! ed      	 	 	 	 	 dPd%ed&ed'ed(ed)ee   d*ee   d+ed,e"de
j>                  fd-       Z# G d. d/      Z$	 dQd0e
j>                  d1ed2ede
j>                  fd3Z%d4e
j>                  de
j>                  fd5Z&d6e
j>                  de
j>                  fd7Z'd8edefd9Z(d:e
j>                  d;ed<ed=e"de
j>                  f
d>Z)d?ed@edAedBedCedee
j>                  e
j>                  f   fdDZ*	 	 	 	 	 	 	 	 	 	 dRd:e
j>                  d%edEedFedGedHedIedJed=e"dBedCede
j>                  fdKZ+y)SzPure audio processing utilities - no TTS/STT imports.

This module contains only audio processing functions (window functions, STFT, mel filterbanks)
that can be imported without pulling in TTS or STT dependencies.
    N)hanninghammingblackmanbartlettSTR_TO_WINDOW_FNstftistft
ISTFTCachemel_filtersintegrated_loudnesslfilternormalize_loudnessnormalize_peakcompute_deltas_kaldimel_scale_kaldiinverse_mel_scale_kaldiget_mel_banks_kaldicompute_fbank_kaldi)	lru_cache)Optional)maxsizec                     |r| n| dz
  }t        j                  t        |       D cg c]4  }ddt        j                  dt        j
                  z  |z  |z        z
  z  6 c}      S c c}w )zHanning (Hann) window.

    Args:
        size: Window length
        periodic: If True, use periodic window (for spectral analysis)
             ?   mxarrayrangemathcospisizeperiodicdenomns       W/Users/ahmed/devFolder/claude-voice/.venv/lib/python3.12/site-packages/mlx_audio/dsp.pyr   r   '   s_     D$(E88@EdL1DHHQ[1_u455	6L L   9A&c                     |r| n| dz
  }t        j                  t        |       D cg c]4  }ddt        j                  dt        j
                  z  |z  |z        z  z
  6 c}      S c c}w )zHamming window.

    Args:
        size: Window length
        periodic: If True, use periodic window (for spectral analysis)
    r   HzG?q=
ףp?r   r   r#   s       r(   r   r   5   s_     D$(E88BG+N+QtxxDGGa% 788	8+N Nr)   c                 6   |r| n| dz
  }t        j                  t        |       D cg c]d  }ddt        j                  dt        j
                  z  |z  |z        z  z
  dt        j                  dt        j
                  z  |z  |z        z  z   f c}      S c c}w )zBlackman window.r   gzG?r   r   g{Gz?   r   r#   s       r(   r   r   C   s     D$(E88
 4[		
 ! DHHQ[1_u4556TXXa$''kAo5667 !		
 	
s   A)Bc                     |r| n| dz
  }t        j                  t        |       D cg c]  }ddt        ||dz  z
        z  |z  z
   c}      S c c}w )zBartlett (triangular) window.r   r   )r   r   r   absr#   s       r(   r   r   Q   sR     D$(E88tMAQSUQY//%77MNNMs   !A)hannr   r   r   r   datarate
block_sizereturnc                 X   t        | t        j                        st        d      t        j                  | j
                  t        j                        st        d      | j                  dk(  r| j                  d   dkD  rt        d      | j                  d   ||z  k  rt        d      y )	Nz#Data must be of type numpy.ndarray.zData must be floating point.r   r      z&Audio must have five channels or less.r   z3Audio must have length greater than the block size.)	
isinstancenpndarray
ValueError
issubdtypedtypefloatingndimshape)r2   r3   r4   s      r(   _validate_loudness_audiorA   a   s    dBJJ'>??==R[[1788yyA~$**Q-!+ABBzz!}zD((NOO )    gain_dbq_factorcenter_freqfilter_typec                    d| dz  z  }dt         j                  z  ||z  z  }t        j                  |      d|z  z  }|dk(  rF||dz   |dz
  t        j                  |      z  z   dt        j                  |      z  |z  z   z  }d|z  |dz
  |dz   t        j                  |      z  z   z  }	||dz   |dz
  t        j                  |      z  z   dt        j                  |      z  |z  z
  z  }
|dz   |dz
  t        j                  |      z  z
  dt        j                  |      z  |z  z   }d|dz
  |dz   t        j                  |      z  z
  z  }|dz   |dz
  t        j                  |      z  z
  dt        j                  |      z  |z  z
  }n|dk(  rrdt        j                  |      z   dz  }dt        j                  |      z    }	dt        j                  |      z   dz  }
d|z   }dt        j                  |      z  }d|z
  }nt        d	|       t        j                  ||	|
g      |z  t        j                  |||g      |z  fS )
N
   g      D@       @
high_shelfr   r   	high_passzUnsupported filter type: )r    r"   sinr!   sqrtr;   r9   r   )rC   rD   rE   r3   rF   	amplitudeomegaalphab0b1b2a0a1a2s                 r(   _biquad_coefficientsrX   o   sZ    w~&I$''M[4/0EHHUOsX~.El"]1}/0$))I&&./

 )^	A)a-488E?1RRS]1}/0$))I&&./
 ]1}/0$))I&&./ 	
 9q=Y]dhhuo$EEF]1}/0$))I&&./ 	
 
	#$((5/!Q&488E?"#$((5/!Q&Y$((5/!Y4[MBCC88RRL!B&"b"(>(CCCrB   bac                 L   t        j                  | t         j                        } t        j                  |t         j                        }t        j                  |      }|j                  dk7  rt	        d      |j
                  dk(  s|d   dk(  rt	        d      | j
                  dk(  rt        j                  |      S | |d   z  } ||d   z  }t        j                  |j                  | j                  |j                        }|j                  |d      }t        j                  ||      }t        t        |      t        |             dz
  }|dk(  r| d   |z  S t        j                  ||      }t        |      D ]  \  }}	| d   |	z  |d   z   }
t        d|      D ]C  }|t        |       k  r| |   |	z  nd}|t        |      k  r||   |
z  nd}||   |z   |z
  ||dz
  <   E |t        |       k  r| |   |	z  nd}|t        |      k  r||   |
z  nd}||z
  |d	<   |
||<    |S )
zApply a 1-D causal linear filter.

    This implements the standard direct-form II transposed recurrence for the
    1-D DSP paths used in this package.
    r=   r   z#dsp.lfilter only supports 1-D inputr   z4filter denominator must have a non-zero leading termFcopy        )r9   asarrayfloat64r?   r;   r$   
zeros_likeresult_typer=   astype
empty_likemaxlenzeros	enumerater   )rY   rZ   r2   r=   xy	state_lenstater'   sampleoutputifeedforwardfeedbacks                 r(   r   r      s    	

1BJJ'A


1BJJ'A::dDyyA~>??vv{adaiOPPvv{}}T""	AaDA	AaDANN4::qww8EE&A
au%ACFCF#a'IA~taxHHYe,Eq\	61q)q)$A+,s1v:!A$-3K()CF
qtf}H 8k1H<E!a%L % 093q6/Aa	lV+s,5A,>1Y<&(C(*b	! " HrB   c                     t        | ||      S N)r   )rY   rZ   r2   s      r(   _apply_lfilterrv      s    1arB   c                 l   t        j                  | t         j                  d      }t        ddt	        j
                  d      z  d|d      \  }}t        dd	d
|d      \  }}t        |j                  d         D ]8  }t        |||d d |f         |d d |f<   t        |||d d |f         |d d |f<   : |S )NT)r=   r^   g      @r   r   g     p@rJ   r_   r   g      C@rL   )	r9   r   rb   rX   r    rN   r   r@   rv   )r2   r3   weightedhigh_shelf_bhigh_shelf_ahigh_pass_bhigh_pass_achannels           r(   _k_weight_audior~      s    xxBJJT:H!5Q1vt\"L,  4CdD+VK*+-,G(< 
G  .hq'z&: 
G	 , OrB   overlapc                 6	   t        j                  | d      }t        |||       |j                  dk(  r|j	                  |j
                  d   d      }t        ||      }|j
                  d   }|j
                  d   }g d}d}d|z
  }	||z  }
t        t        j                  |
|z
  ||	z  z        dz         }t        j                  d|      }t        j                  ||ft         j                        }t        |      D ]q  }|D ]j  }t        |||	z  z  |z        }t        |||	z  dz   z  |z        }d||z  z  t        j                  t        j                  ||||f               z  |||f<   l s t        j                          5  t        j"                  d	t$        
       |D cg c]R  }ddt        j&                  t        j                  t        |      D cg c]  }||   |||f   z   c}            z  z   T }}}ddd       t)              D cg c]  \  }}||k\  r| }}}t        j                          5  t        j"                  d	t$        
       t        |      D cg c]*  }t        j*                  |D cg c]	  }|||f    c}      , }}}ddd       ddt        j&                  t        j                  t        |      D cg c]  }||   |   z   c}            z  z   dz
  }t)        |      D cg c]  \  }}||kD  r||kD  r| }}}t        j                          5  t        j"                  d	t$        
       t        j,                  t        j                  t        |      D cg c]*  }t        j*                  |D cg c]	  }|||f    c}      , c}}            }ddd       t        j.                  d	      5  t1        ddt        j&                  t        j                  t        |      D cg c]  }||   |   z   c}            z  z         cddd       S c c}w c c}}w # 1 sw Y   =xY wc c}}w c c}w c c}}w # 1 sw Y   xY wc c}w c c}}w c c}w c c}}w # 1 sw Y   xY wc c}w # 1 sw Y   yxY w)z>Measure integrated loudness in LUFS using BS.1770 K-weighting.Tr]   r   r   )      ?r   r   (\?r   g     Qr   r\   ignore)categoryg&1      $@N)divide)r9   r   rA   r?   reshaper@   r~   introundarangeri   rb   r   sumsquarewarningscatch_warningssimplefilterRuntimeWarninglog10rj   mean
nan_to_numerrstatefloat)r2   r3   r4   r   
input_datanum_channelsnum_sampleschannel_gainsabsolute_thresholdstepduration_seconds
num_blocksblock_indicesmean_squarer}   block_indexlowerupperblock_loudnessloudnessgated_blocksgated_mean_squarerelative_thresholds                          r(   r   r      s    $T*JZz:!''
(8(8(;Q?
 T2J##A&L""1%K/M=D"T)
#j0Z$5FGIAMJ IIa,M((L*5RZZHK&(K
kD&89D@AE
kD&81&<=DEE14
T8I1Jbff		*U5['%9:;O 1K,- ) ' 
	 	 	"h@  -
  - hh (-\':':G &g.Wk=Q1RR':	  - 	 
 
#$ &/~%>%>!K)) 	%>   
	 	 	"h@ !.
. GG,W,;[+!56,WX. 	 
 
# 	

((FF $)#6#6 "'*->w-GG#6

		
 
	   &/~%>%>!K((X8J-J 	%>   
	 	 	"h@MMHH $)#6 $7 GG 0</; ((<=/; $7

 
#  
H	%hh (-\':':G &g.1B71KK':	
 
&	%u
 
#	"  X
 
#	"	 
#	", 
&	%s   !P=(2P7P2.P7?P=Q
 *Q*QQ	QQQ(Q-0AQ>8Q8Q3	Q8'Q>3RR
R2P77P==QQQQ%3Q88Q>>R
RRinput_loudnesstarget_loudnessc                     ||z
  }t        j                  d|dz        }|| z  }t        j                  t        j                  |            dk\  rt	        j
                  d       |S )z-Normalize audio to a target loudness in LUFS.r         4@r   #Possible clipped samples in output.)r9   powerrg   r0   r   warn)r2   r   r   delta_loudnessgainrp   s         r(   r   r   T  sW     %~5N88D.4/0DD[F	vvbffVn$;<MrB   target_peak_dbc                    t        j                  t        j                  |             }t        j                  d|dz        |z  }|| z  }t        j                  t        j                  |            dk\  rt	        j
                  d       |S )z)Normalize audio to a target peak in dBFS.r   r   r   r   )r9   rg   r0   r   r   r   )r2   r   current_peakr   rp   s        r(   r   r   d  sf    66"&&,'L88D.4/0<?DD[F	vvbffVn$;<MrB   windowc                    ||dz  }||}t        |t              r<t        j                  |j	                               }|t        d|        ||      }n|}|j                  d   |k  r?||j                  d   z
  }	t        j                  |t        j                  |	f      gd      }dd}
|r |
| |dz  |      } d| j                  d   |z
  |z  z   }|dk  r%t        d| j                  d    d	| d
| d| d	      ||f}|df}t        j                  | ||      }t        j                  j                  ||z        S )Nr.   Unknown window function: r   axisc                     |dk(  rt        j                  | ||fg      S |dk(  r5| d|dz    d d d   }| |dz    d d d d   }t        j                  || |g      S t        d|       )Nconstantreflectr   r`   zInvalid pad_mode )r   padconcatenater;   )rk   paddingpad_modeprefixsuffixs        r(   _padzstft.<locals>._pad  s    z!66!w0122"q7Q;'"-F1~+DbD1F>>61f"5660
;<<rB   r   r   zInput is too short (length=z) for n_fft=z with hop_length=z and center=.r@   strides)r   )r8   strr   getr   r;   r@   r   r   ri   
as_stridedfftrfft)rk   n_fft
hop_length
win_lengthr   centerr   	window_fnwpad_sizer   
num_framesr@   r   framess                  r(   r   r   q  s_    aZ

&#$((8	8ABBj!wwqzE1771:%NNArxx45A>= EQJ)aggaj5(Z77JQ)!''!*\%HYZdYeeqrxqyyz{
 	
 E1oG]]1E7;F66;;vz""rB   c                 b   || j                   d   dz
  dz  }||dz  }t        |t              rBt        j	                  |j                               }|t        d|        ||dz         dd }n|}|j                   d   |k  r=t        j                  |t        j                  ||j                   d   z
  f      gd      }| j                   d   }	|	dz
  |z  |z   }
t        j                  |
      }t        j                  |
      }t        j                  j                  | d      j                  dd      }t        j                  |	      |z  }|dddf   t        j                  |      z   }|j                         }||z  j                         }|r||z  n|}t        j                  ||	f      j                         }|j                   |   j#                  |      }|j                   |   j#                  |      }t        j$                  |d	kD  ||z  |      }|r|||dz  | dz   }||d| }|S )
ak  Inverse Short-Time Fourier Transform.

    Args:
        x: Complex STFT output of shape (n_fft // 2 + 1, num_frames)
        hop_length: Hop length between frames (default: win_length // 4)
        win_length: Window length (default: (n_fft - 1) * 2)
        window: Window function name or array (default: "hann")
        center: If True, remove center padding (default: True)
        length: Target output length (default: None)
        normalized: If True, use window squared (COLA) normalization. If False, use simple window normalization (default: True)

    Returns:
        Reconstructed time-domain signal
    Nr   r   r.   r   r`   r   r   绽|=)r@   r8   r   r   r   r   r;   r   r   ri   r   irfft	transposer   flattentileataddwhere)rk   r   r   r   r   length
normalizedr   r   r   treconstructed
window_sumframes_timeframe_offsetsindicesindices_flatupdates_reconstructedwindow_normupdates_windows                       r(   r	   r	     s'   . ggaj1n)
1_
&#$((8	8ABBj1n%cr*wwqzJNNArxxaggaj)@(BCD1MJ	a:%
2AHHQKM!J &&,,qq,)33Aq9K IIj)J6MAtG$ryy'<<G??$L(1_557'1q5QKWW[:-8@@BN "$$\2667LMM|,00@J HHUMJ6M &.%jAoq8HI%gv.rB   sample_rater   n_melsf_minf_maxnorm	mel_scaleprecisec           	      X   	
 dd	dd
xs  dz  	
 f	d}|r`t        j                  t         j                        5   |t         j                        j	                  t         j
                        cddd       S  |t         j
                        S # 1 sw Y   xY w)u  Triangular mel filterbank as an mx.array of shape (n_mels, n_fft // 2 + 1).

    Args:
        precise: If True, compute the filterbank in float64 on the CPU stream
            and cast to float32 before returning. Default float32 path drifts
            ~5e-6 from a torchaudio float64 reference — enough to perturb the
            CTC decode in numerically sensitive models (e.g. granite_speech_nar).
            One-time cost at lru_cache miss; runtime use is unaffected.
    c                     |dk(  rdt        j                  d| dz  z         z  S d\  }}| |z
  |z  }d}||z
  |z  }t        j                  d      dz  }| |k\  r|t        j                  | |z        |z  z   }|S )	Nhtk     F@r        @r_   gP@     @@皙@      ;@)r    r   log)freqr   r   f_spmels
min_log_hzmin_log_mellogsteps           r(   	hz_to_melzmel_filters.<locals>.hz_to_mel  s    DJJsTE\'9::: %tu$
!E)T1((3-$&:$*;!<w!FFDrB   c           	          |dk(  rdd| dz  z  dz
  z  S d\  }}||| z  z   }d}||z
  |z  }t        j                  d      d	z  }t        j                  | |k\  |t        j                  || |z
  z        z  |      }|S )
Nr   r   r   r   r   r   r   r   r   )r    r   r   r   exp)r   r   r   r   freqsr   r   r   s           r(   	mel_to_hzzmel_filters.<locals>.mel_to_hz  s    DTF]3c9:: %tt#
!E)T1((3-$&K4++= >??

 rB   r   c                 l  	 dz  dz   }t        j                  ddz  ||       }       }       }t        j                  ||dz   |       } |      }|dd  |d d z
  }t        j                  |d      t        j                  |d      z
  }|d d d df    |d d z  }	|d d dd f   |dd  z  }
t        j                  t        j                  |	      t        j
                  |	|
            }dk(  r*d|ddz    |d  z
  z  }|t        j                  |d      z  }|j                  dd      S )	Nr   r   r   r\   r`   rK   slaneyrI   )r   linspaceexpand_dimsmaximumrc   minimummoveaxis)r=   n_freqs	all_freqsm_minm_maxm_ptsf_ptsf_diffslopesdown_slopes	up_slopes
filterbankenormr   r   r   r   r  r   r   r   r   s                r(   _buildzmel_filters.<locals>._build)  sT    1*q.KK;!#3WEJ	 %+%+E5&1*EB%+ qrU3BZ'q)BNN9a,HH q#2#v&"+51ab5MF12J.	ZZMM+&

;	(J

 85VaZ05&>ABE"..22J""1a((rB   N)r   )r   streamcpurb   re   float32)r   r   r   r   r   r   r   r   r  r   r  s   ```````  @@r(   r   r     sy    *" $[1_E) )@ YYrvv"**%,,RZZ8 "** s   3B  B)c                       e Zd ZdZd ZdededefdZdededed	ej                  def
d
Z		 	 ddej                  dej                  dededed	ej                  de
dedej                  fdZd Zd Zy)r
   z
    Advanced caching for iSTFT operations. Fully vectorized Overlap-Add for MLX.
    Handles multiple configurations efficiently.
    Automatically caches normalization buffers and position indices for maximum performance.
    c                      i | _         i | _        y ru   )norm_buffer_cacheposition_cacheselfs    r(   __init__zISTFTCache.__init__W  s    !# rB   r   frame_lengthr   c                     |||f}|| j                   vrZt        j                  |      dddf   |z  t        j                  |      dddf   z   }|j                  d      | j                   |<   | j                   |   S )z.Get cached position indices or create new onesNr`   )r  r   r   r   )r!  r   r#  r   key	positionss         r(   get_positionszISTFTCache.get_positions[  s    <4d)))		*%ag.;))L)$'23  (1'8'8'<D$""3''rB   r   r   r   c                    t        t        |j                                     }|||||f}|| j                  vr|j                  d   }|dz
  |z  |z   }	| j                  |||      }
|dz  }t        j                  |	t        j                        }t        j                  ||      }|j                  |
   j                  |      }t        j                  |d      }|| j                  |<   | j                  |   S )z1Get cached normalization buffer or create new oner   r   r   r\   r   )hashtupletolistr  r@   r'  r   ri   r  r   r   r   r	  )r!  r   r   r   r   r   window_hashr%  r#  ola_lenpositions_flatwindow_squarednorm_bufferwindow_sq_tileds                 r(   get_norm_bufferzISTFTCache.get_norm_bufferh  s     512j*k:Fd,,,!<<?L!A~3lBG!//
L*UN#QYN((7"**=K ggnjAO%..8<<_MK**[%8K*5D""3'%%c**rB   N	real_part	imag_partr   audio_lengthr5   c	                 Z   |j                   d   |k  rI||j                   d   z
  }	t        j                  |t        j                  |	f|j                        g      }|d|z  z   }
t        j
                  j                  |
j                  ddd      |d      }||z  }|j                   \  }}}|dz
  |z  |z   }| j                  |||||      }| j                  |||      }t        j                  |      |z  }|dddf   |dddf   z   }t        j                  ||z  t        j                        }|j                  |j                  d         j                  |j                  d            }|j                  ||      }||dddf   z  }|r|dz  }|dd|df   }||ddd|f   }|S )	a  
        iSTFT with automatic caching and vectorized overlap-add.

        Args:
            real_part: Real part of STFT output (batch, freq, time)
            imag_part: Imaginary part of STFT output (batch, freq, time)
            n_fft: FFT size
            hop_length: Hop length
            win_length: Window length
            window: Window function
            center: If True, remove center padding
            audio_length: Target audio length

        Returns:
            Reconstructed audio (batch, samples)
        r   r\   y              ?r   r   r`   r'   r   N)r@   r   r   ri   r=   r   r   r   r2  r'  r   r  r   r   r   )r!  r3  r4  r   r   r   r   r   r5  r   stft_complextime_frameswindowed_frames
batch_sizer   r#  r-  r0  r.  batch_offsetsglobal_indicesrp   	start_cuts                          r(   r	   zISTFTCache.istft  s   8 <<?U"&,,q/)C^^VRXXsfFLL-Q$RSF !2	>1ffll<#9#9!Q#BeRTlU &./>/D/D,
J>Z/,> **:z6:
 ++JjQ 		*-7'a0=D3II:/

C>11"56::?;R;RSU;VW
G4 +dAg.. 
IAyzM*F#A}},-FrB   c                 l    | j                   j                          | j                  j                          y)z$Clear all cached data to free memoryN)r  clearr  r   s    r(   clear_cachezISTFTCache.clear_cache  s&    $$&!!#rB   c                     t        | j                        t        | j                        t        | j                        t        | j                        z   dS )z"Get information about cached items)norm_buffersposition_indicestotal_cached_items)rh   r  r  r   s    r(   
cache_infozISTFTCache.cache_info  sK       6 67 #D$7$7 8"%d&<&<"=$%%&#'
 	
rB   )TN)__name__
__module____qualname____doc__r"  r   r'  r   r   r2  boolr	   rA  rF   rB   r(   r
   r
   P  s    !( (3 (C (++ + 	+
 + +F  C88C 88C 	C
 C C C C C 
CJ$

rB   r
   specgramr   modec                 0   |dk  rt        d|       | j                  }| j                  d|d         } | j                  d   }|dz
  dz  }t        ||dz   z  d|z  dz   z        dz  }|dk(  r]t	        j
                  | d	d	ddf   |d
      }t	        j
                  | d	d	dd	f   |d
      }t	        j                  || |gd
      }	nt	        j                  | d||fg      }	t	        j                  | |dz   |	j                        }
|	j                  d   d|z  z
  }t	        j                  ||f|	j                        }t        |      D ]6  }|	d	d	|||z   f   }||
z  }t	        j                  |d
      |z  |d	d	|f<   8 |j                  |      S )a  
    Compute delta coefficients of a spectrogram (Kaldi-compatible).

    The formula is:
    d_t = sum_{n=1}^{N} n * (c_{t+n} - c_{t-n}) / (2 * sum_{n=1}^{N} n^2)

    Args:
        specgram: MLX array of dimension (..., freq, time)
        win_length: The window length used for computing delta (default: 5)
        mode: Padding mode - "edge" or "constant" (default: "edge")

    Returns:
        MLX array of deltas of dimension (..., freq, time)
       zwin_length should be >= 3, got r`   r   r   r   g      @edgeNr   r   r   r\   )r;   r@   r   r   r   repeatr   r   r   r=   ri   r   r   )rM  r   rN  original_shapenum_featuresr'   r&   pad_left	pad_rightpaddedkernel_weights
time_stepsrp   rq   r   rx   s                   r(   r   r     s   " A~::,GHH^^NN2$67H>>!$L	aAA!q1u+Q+,s2E v~99Xa1f-qq9IIhq"#v.:	8Y ?aH6Aq6"23YYr1q5=Na1q5(JXX|Z0EF:1q:~--.N*vvhQ/%7q!t 
 >>.))rB   r   c                 >    dt        j                  d| dz  z         z  S )z/Convert frequency to mel scale (Kaldi formula).     @r   r   )r   r   )r   s    r(   r   r     s    BFF3-...rB   mel_freqc                 >    dt        j                  | dz        dz
  z  S )z/Convert mel scale to frequency (Kaldi formula).r   r\  r   )r   r  )r]  s    r(   r   r     s     BFF8f,-344rB   rk   c                 <    | dk(  rdS d| dz
  j                         z  S )z7Returns the smallest power of 2 that is greater than x.r   r   r   )
bit_length)rk   s    r(   _next_power_of_2ra    s%    Q15A!a%!3!3!555rB   waveformwindow_sizewindow_shift
snip_edgesc                    | j                   d   }|r&||k  rt        j                  d      S d||z
  |z  z   }n~||dz  z   |z  }|dz  |dz  z
  }|dkD  r@| d|dz    ddd   }|dkD  r| d| dz
  d   n| ddd   }t        j                  || |g      } n#| ddd   }t        j                  | | d |g      } t        j                  | ||f|df      S )zCExtract frames from waveform using strided windowing (Kaldi-style).r   rR  r   r   Nr`   r   )r@   r   ri   r   r   )	rb  rc  rd  re  r   mr   rV  rW  s	            r(   _get_strided_kaldirh    s    ..#K$88F##{*|;;LA-.<?Q!227C!G,TrT2H8;asdQh!34XbQRSUgEVI~~x9&EFH 2I~~x	&BCH==![)9LRSCTUUrB   num_binswindow_length_paddedsample_freqlow_freq	high_freqc                    | dkD  sJ d       |dz  dk(  sJ |dz  }d|z  }|dk  r||z  }d|cxk  r|k  rn J d|cxk  r|k  sJ  J ||z  }t        t        t        j                  |                  }t        t        t        j                  |                  }	|	|z
  | dz   z  }
t        j                  |       j                  dd      }|||
z  z   }||d	z   |
z  z   }||d
z   |
z  z   }t        |      }t        |t        j                  |      z        j                  dd      }||z
  ||z
  z  }||z
  ||z
  z  }t        j                  t        j                  d      t        j                  ||            }||j                         fS )a  
    Create Kaldi-compatible mel filterbank matrix.

    Args:
        num_bins: Number of mel bins
        window_length_padded: Padded window length (FFT size)
        sample_freq: Sample frequency in Hz
        low_freq: Low frequency cutoff
        high_freq: High frequency cutoff (0 or negative = relative to Nyquist)

    Returns:
        (bins, center_freqs): Mel filterbank matrix and center frequencies
    rP  zMust have at least 3 mel binsr   r   r   r_   r   r`   r   rI   )r   r   r   r   r   r   r   r	  ri   r
  squeeze)ri  rj  rk  rl  rm  num_fft_binsnyquistfft_bin_widthmel_low_freqmel_high_freqmel_freq_deltabin_idxleft_mel
center_mel	right_melcenter_freqsmelup_slope
down_slopebinss                       r(   r   r   3  s   ( a<888<!#q((('1,LKGCW	8%g%GGC),Fw,FGG,FGG"66M();<=L/"((9*=>?M#l2x!|DNii!))"a0Gg66H3. @@J#??I*:6L
-"))L*AA
B
J
J1b
QCh:#89Hc/i*&<=J::bhhqk2::h
#CDD%%'''rB   win_lenwin_incnum_melswin_typepreemphasisditherc                 6   | j                   dk(  r| d   } ||z  dz  }||z  dz  }t        ||z  dz        }t        ||z  dz        }t        |      }t        | |||      }|j                  d   dk(  rt        j                  d|f      S |dk7  r1t
        j                  j                  |j                        |z  }||z   }t        j                  |dd      }||z
  }|dk7  r>|d	d	ddf   }|d	d	dd	f   ||d	d	d	d
f   z  z
  }t        j                  ||gd      }|dk(  rKt        j                  |      }ddt        j                  dt
        j                  z  |z  |dz
  z        z  z
  }n|dk(  rKt        j                  |      }ddt        j                  dt
        j                  z  |z  |dz
  z        z  z
  }n{|dk(  rat        j                  |      }ddt        j                  dt
        j                  z  |z  |dz
  z        z  z
  }t        j                  |d      }nt        j                  |      }||z  }||k7  r||z
  }t        j                   |dd|fg      }t
        j"                  j%                  ||d      }t        j&                  |      dz  }t)        ||t+        |      |	|
      \  }}t        j                   |ddg      }t        j,                  ||j.                        }t        j0                  t        j2                  |d            }|S )a  
    Compute Kaldi-compatible log mel-filterbank features.

    Args:
        waveform: Input audio (1D array)
        sample_rate: Sample rate in Hz
        win_len: Window length in samples
        win_inc: Window shift in samples
        num_mels: Number of mel bins
        win_type: Window type ("hamming", "hanning", "povey", "rectangular")
        preemphasis: Preemphasis coefficient
        dither: Dither amount (0 to disable)
        snip_edges: If True, discard incomplete frames at edges
        low_freq: Low frequency cutoff
        high_freq: High frequency cutoff (0 = Nyquist)

    Returns:
        Log mel-filterbank features (time, num_mels)
    r   r   i  gMbP?r_   r   T)r   keepdimsNr`   r   r   r+   r,   r   r   poveyg333333?rR  r7  rI   )r   r   g:0yE>)r?   r   ra  rh  r@   r   ri   randomnormalr   r   r   r!   r"   r   onesr   r   r   r0   r   r   matmulTr   r	  )rb  r   r  r  r  r  r  r  re  rl  rm  frame_length_msframe_shift_mswindow_shift_samplesrc  padded_window_sizestrided_input
rand_gauss	row_means	first_col
other_colsr'   r   r1   r   
fft_resultspectrummel_energies_mel_featuress                                 r(   r   r   f  s   @ }}A;+d2O{*T1N{^;eCDkO3e;<K)+6 '+3ZM 1"xxH&& }YY%%m&9&9:VC
%
2 A=I!I-M c!!QqS&)	"1ab5)K-3B3:O,OO
	:'>QG 9IIk"rvva"%%i!m{Q&GHHH	Y	IIk"sRVVAIM[1_$EFFF	W	IIk"S266!bee)a-;?"CDDD$%%!F*M [(${2}v7|.DE ].@qIJvvj!S(H *$eK&8(IOL! 66,(89L 99X|~~6L66"**\489LrB   )F)g?g      ?)i   NNr1   Tr   )NNr1   TNF)r   NNr   F)r7   rQ  )
i  i  i  <   r   g
ףp=
?r   Tr   r_   ),rJ  r    r   __all__	functoolsr   typingr   mlx.corecorer   numpyr9   r   r   r   r   r   r:   r   r   rA   r   r*  rX   r   rv   r~   r   r   r   r   r   r	   rK  r   r
   r   r   r   ra  rh  r   r   rL  rB   r(   <module>r     s    *      4
 
 4
 
 4
 
 4O O  P2:: PS Pe PPT P,D,D,D ,D 	,D
 ,D 2::rzz!",D^(rzz (bjj (

 (rzz (Vbjj RZZ rzz bjj "** C BJJ , 	q

**q

q
 q
 	q

 q
h
**  ZZ	 	 	U 	rzz 	 #0#
 HHsN0#j I^ 4
 !YYY Y 	Y
 E?Y 3-Y Y Y XXY YxD
 D
Z :@,*hh,*$',*36,*XX,*^/"(( /rxx /
5bhh 5288 5
6 6 6
VhhV%(V8;VIMVXXV20(0(0( 0( 	0(
 0( 288RXX0(j chhcc c 	c
 c c c c c c c XXcrB   