
    i                        d dl mZmZmZmZ d dlmZ 	 	 	 	 ddej                  deee	ee	df   ee	   f      deee
ee
   ee
df   f      dedee   d	ej                  fd
Z	 	 ddej                  de	dedee   d	ej                  f
dZy)    )ListOptionalTupleUnionNinputsize.scale_factormodealign_cornersreturnc                    | j                   }|dk  rt        d| d      |dz
  }||t        d      ||t        d      |t        |t        t        f      s|g|z  }|t        |t        t        f      s|g|z  }|ag }t        |      D ]Q  }t        dt        t        j                  | j                  |dz      ||   z                    }|j                  |       S |dk(  rt        | |d   ||      S t        d	| d      )
a  Interpolate array with correct shape handling.

    Args:
        input (mx.array): Input tensor [N, C, ...] where ... represents spatial dimensions
        size (int or tuple): Output size
        scale_factor (float or tuple): Multiplier for spatial size
        mode (str): 'nearest' or 'linear'
        align_corners (bool): If True, align corners of input and output tensors
       z+Expected at least 3D input (N, C, D1), got D   z2Only one of size or scale_factor should be definedz+One of size or scale_factor must be defined   r   z/Only 1D interpolation currently supported, got )ndim
ValueError
isinstancelisttuplerangemaxintmxceilshapeappendinterpolate1d)	r   r   r	   r
   r   r   spatial_dimsi	curr_sizes	            u/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/mlx_audio/tts/models/interpolate.pyinterpolater#      s0     ::DaxFtfANOO!8L L4MNN	,.FGG 
4$ ?v$
<$(O$~4 ||$AAs2775;;q1u+=Q+O#PQRIKK	" % qUDGT=AA=l^1M
 	
    c                    | j                   \  }}}|dk  rd}|dk  rd}|dk(  r|dk(  rt        j                  dg      }ng||z  }t        j                  t        j                  |      |z        j                  t        j                        }t        j                  |d|dz
        }| dddd|f   S |r'|dkD  r"t        j                  |      |dz
  |dz
  z  z  }	nG|dk(  rt        j                  dg      }	n+t        j                  |      ||z  z  }	|s|	d||z  z  z   dz
  }	|dk(  rt        j                  | |||f      }
|
S t        j                  |	      j                  t        j                        }t        j                  |dz   |dz
        }|	|z
  }| dddd|f   }| dddd|f   }|d|z
  ddddf   z  ||ddddf   z  z   }
|
S )z 1D interpolation implementation.r   nearestr   Ng        g      ?)
r   r   arrayfloorarangeastypeint32clipbroadcast_tominimum)r   r   r
   r   
batch_sizechannelsin_widthindicesscalexoutputx_lowx_highx_fracy_lowy_highs                   r"   r   r   9   s    &+[["J( ax!|y19hhsmGtOEhhryy67>>rxxHGgggq(Q,7GQ7]## IIdO1:;19#A		$8d?3A x$//#5 1}Xt(DEHHQKrxx(EZZ	8a<0FYF !Q+E1a< F a&j$a-006F4q=<Q3QQFMr$   )NNr&   N)linearN)typingr   r   r   r   mlx.corecorer   r'   r   floatstrboolr#   r    r$   r"   <module>rC      s    / / 
 >BKO$(0
880

5eCHotCy89
:0
 5UU5#:5F!FGH0
 	0

 D>0
 XX0
l $(	3883
3 3 D>	3
 XX3r$   