
    i%              #          d Z ddlmZ ddlmZ ddlmZmZm	Z	m
Z
 g dZ G d d      Zddddd	d
d
dd
d
ddddej                  dedee   dee   dedededededededededej                  fdZddddd	d
d
dd
dddddd
ddej                  dee   dee   dee   dedededededededee   deded ed!edej                  f"d"Zy)#u  Demucs spectral wrappers around ``mlx_spectro``.

Re-exports the core ``SpectralTransform`` from ``mlx_spectro`` and adds
``spectro`` / ``ispectro`` convenience functions that handle the
multi-dimensional tensor layouts used by Demucs models (3-D for STFT,
4-D and 5-D for iSTFT).

``CachedSpectralPair`` caches the transform and lazily creates
``compiled_pair()`` instances for repeated chunk sizes, giving 1.3–1.7x
speedup over the eager path.
    )OptionalN)SpectralTransform
WindowLikeget_transform_mlxresolve_fft_params)CachedSpectralPairr   spectroispectroc                   ,   e Zd ZdZ	 	 ddedee   ddfdZdededefd	Zd
e	j                  de	j                  fdZde	j                  dede	j                  fdZd
e	j                  de	j                  fdZde	j                  dede	j                  fdZy)r   u5  Cached STFT/iSTFT transform with compiled_pair() acceleration.

    Creates the ``SpectralTransform`` once and reuses it.  For repeated
    chunk sizes (the common case in Demucs inference), lazily builds and
    caches ``compiled_pair()`` graphs that eliminate Python dispatch overhead.

    Handles the Demucs multi-dimensional reshape internally:
      - stft:  [B, C, T]   → [B*C, T]  → stft → [B, C, F, N]
      - istft: [B, C, F, N] → [B*C, F, N] → istft → [B, C, T]
               [B, S, C, F, N] → [B*S*C, F, N] → istft → [B, S, C, T]
    Nn_fft
hop_lengthreturnc           
      f    t        ||d d      \  }}}t        |||ddddd       | _        i | _        y )Nr   hannTFr   r   
win_length	window_fnperiodiccenter
normalizedwindow)r   r   
_transform_compiled_cache)selfr   r   	eff_n_ffthopwins         h/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/demucs_mlx/spec_mlx.py__init__zCachedSpectralPair.__init__,   sI    
 1
D!L	3+	
 24    lengthbatchc                     || j                   vr-| j                  j                  |d|      }|| j                   |<   | j                   |   S )z:Get or create a compiled_pair for the given signal length.bfn)r!   layoutwarmup_batch)r   r   compiled_pair)r   r!   r"   pairs       r   	_get_pairzCachedSpectralPair._get_pair>   sS    ---??00e% 1 D ,0D  (##F++r    xc                    |j                   dk(  r|j                  \  }}}t        j                  |      j	                  ||z  |      }| j                  |||z        \  }} ||      }|j	                  |||j                  d   |j                  d         S | j                  t        |j                  d         t        |j                  d               \  }} ||      S )u   STFT with Demucs multi-dim support.

        Input [B, C, T] → output [B, C, F, N].
        Input [B, T]    → output [B, F, N].
                 r   ndimshapemx
contiguousreshaper)   int)	r   r*   BCTx2stft_fn_spec2s	            r   stftzCachedSpectralPair.stftG   s     66Q;ggGAq!q!))!a%3B1q51JGQBKE==Au{{1~u{{1~FF^^C$4c!''!*oF
qzr    zc                    |j                   dk(  r|j                  \  }}}}}t        j                  |      j	                  ||z  |z  ||      }| j                  |||z  |z        \  }	}
 |
|      }|j	                  ||||j                  d         S |j                   dk(  rz|j                  \  }}}}t        j                  |      j	                  ||z  ||      }| j                  |||z        \  }	}
 |
|      }|j	                  |||j                  d         S | j                  |t        |j                  d               \  }	}
 |
|      S )u   iSTFT with Demucs multi-dim support.

        Input [B, S, C, F, N] → output [B, S, C, T].
        Input [B, C, F, N]    → output [B, C, T].
        Input [B, F, N]       → output [B, T].
           r-      r   r0   )r   r?   r!   r7   Sr8   FNz2r<   istft_fnwav2s               r   istftzCachedSpectralPair.istftW   s,    66Q;GGMAq!Qq!))!a%!)Q:B..Q;KAxB<D<<1aA7766Q;JAq!Qq!))!a%A6B..Q7KAxB<D<<1djjm44nnVS_=8{r    c                 X   |j                   dk(  r|j                  \  }}}t        j                  |      j	                  ||z  |      }| j
                  j                  |      }|j	                  |||j                  d   |j                  d         S | j
                  j                  |      S )z?STFT using cached transform without compiled graphs (fallback).r,   r-   r.   )r1   r2   r3   r4   r5   r   r>   )r   r*   r7   r8   r9   r:   r=   s          r   
stft_eagerzCachedSpectralPair.stft_eagero   s    66Q;ggGAq!q!))!a%3BOO((,E==Au{{1~u{{1~FF##A&&r    c                 ^   |j                   dk(  r||j                  \  }}}}}t        j                  |      j	                  ||z  |z  ||      }| j
                  j                  ||      }	|	j	                  ||||	j                  d         S |j                   dk(  rw|j                  \  }}}}t        j                  |      j	                  ||z  ||      }| j
                  j                  ||      }	|	j	                  |||	j                  d         S | j
                  j                  ||      S )z@iSTFT using cached transform without compiled graphs (fallback).rA   )r!   r-   rB   )r1   r2   r3   r4   r5   r   rI   )
r   r?   r!   r7   rC   r8   rD   rE   rF   rH   s
             r   istft_eagerzCachedSpectralPair.istft_eagerx   s   66Q;GGMAq!Qq!))!a%!)Q:B??((F(;D<<1aA7766Q;JAq!Qq!))!a%A6B??((F(;D<<1djjm44$$Qv$66r    )i   N)__name__
__module____qualname____doc__r6   r   r   tupler)   r3   arrayr>   rI   rK   rM    r    r   r   r      s    
 $(44 SM4 
	4$, ,C ,E ,bhh 288  rxx   0'BHH ' '7RXX 7s 7rxx 7r    r   i   r   TF)r   r   r   r   r   r   r   r   onesidedreturn_complexpad
torch_liker*   r   r   r   r   r   r   r   r   rU   rV   rW   rX   r   c          
      $   |	st        d      |
st        d      t        t        |      ||t        |            \  }}}|rJ|rH|dz  }t        | j                  d         |k  r(t	        dt        | j                  d          d| d      t        ||||||||      }| j                  d	k(  rw| j                  \  }}}t        j                  |       j                  ||z  |      }|j                  |      }|j                  |||j                  d
   |j                  d         S | j                  d
k(  }|r
| dddf   } n'| j                  dk7  rt        d| j                         |j                  |       }|rt        j                  |d      S |S )u   Torch-compatible STFT with Demucs multi-dim support.

    Input shapes:
    - [T]         → output [F, N]
    - [B, T]      → output [B, F, N]
    - [B, C, T]   → output [B, C, F, N]   (Demucs layout)
    Only onesided=True supportedz"Only return_complex=True supportedr.   r/   z@stft: reflect padding requires input length > eff_n_fft//2 (len=z, pad=z).r   r,   r-   Nz,spectro expects [T], [B,T], or [B,C,T], got r   )axis)NotImplementedErrorr   r6   r2   RuntimeErrorr   r1   r3   r4   r5   r>   
ValueErrorsqueeze)r*   r   r   r   r   r   r   r   r   rU   rV   rW   rX   r   r   r   pad_amt	transformr7   r8   r9   r:   r=   orig_1dspecs                            r   r	   r	      s   . !"@AA!"FGG,E
J
CHIsC
 fq.qwwr{w&%%(%5$6fWIRI 
 "	I 	vv{''1a]]1%%a!eQ/r"}}Q5;;q>5;;q>BB ffkGdAgJ	
1GyQRR>>!D'.2::d#8D8r    auto)r   r   r   r   r   r   r   r   rU   rV   r!   rW   rX   safetyallow_fusedr?   r!   re   rf   c          
         |t        d      |	st        d      |
rt        d      t        |      }| t        | j                  d         }|dz
  dz  }t	        t        |      ||t        |            \  }}}t        ||||||||      }t        |t        |      t        |      |	      }| j                  d
k(  rr| j                  \  }}}}}t        j                  |       j                  ||z  |z  ||      } |j                  |fi |}|j                  ||||j                  d         S | j                  dk(  rm| j                  \  }}}}t        j                  |       j                  ||z  ||      } |j                  |fi |}|j                  |||j                  d         S | j                  dk(  }|r| dddddf   } n'| j                  dk7  rt        d| j                          |j                  | fi |}|r|d   S |S )u  Torch-compatible iSTFT with Demucs multi-dim support.

    Input shapes:
    - [F, N]         → output [T]
    - [B, F, N]      → output [B, T]
    - [B, C, F, N]   → output [B, C, T]       (Demucs layout)
    - [B, S, C, F, N] → output [B, S, C, T]   (Demucs bag layout)
    Nzhop_length requiredrZ   zOnly real output supportedr-   r.   r   )r!   rX   rf   re   rA   rB   r,   z@ispectro expects [F,N], [B,F,N], [B,C,F,N], or [B,S,C,F,N], got r   )r^   r\   r6   r2   r   r   dictboolr1   r3   r4   r5   rI   )r?   r   r   r   r   r   r   r   r   rU   rV   r!   rW   rX   re   rf   r   Fbinsr   r   ra   istft_kwr7   rC   r8   rD   rE   rF   rH   orig_2dwavs                                  r   r
   r
      s#   6 .//!"@AA!">??
j/C}AGGBK a,E
CSXIsC "	I 
#%	H 	vv{1aA]]1%%a!eaiA6yr.X.||Aq!TZZ]33 	vv{WW
1a]]1%%a!eQ2yr.X.||Aq$**Q-00 ffkGdAqjM	
1##$77)-
 	

 )//!
(x
(C3q6%#%r    )rQ   typingr   mlx.corecorer3   mlx_spectror   r   r   r   __all__r   rS   r6   strrj   r	   r
   rT   r    r   <module>ru      s  
   e7 e7V  $ $D9	xxD9 D9 	D9
 D9 D9 D9 D9 D9 D9 D9 D9 
D9 D9 XXD9T   $ $  #X&	xxX& C=X& 	X&
 X& X& X& X& X& X& X& X& SMX& 
X& X&  !X&" #X&$ XX%X&r    