+
    &j8                     4   ^ RI t ^ RIHt ^ RIHtHtHt ^ RIHt	 ^ RI
Ht R tR tR t ! R R	]4      t ! R
 R]4      t ! R R]4      t ! R R]4      t ! R R]4      t ! R R]4      t ! R R]4      t ! R R]4      t ! R R]4      t ! R R]4      tR# )    N)
accumulate)OptionalTupleUnion)Modulec                     \        V \        \        34      '       d'   \        V 4      V8w  d   \	        V4      h\        V 4      # \        V \
        4      '       g   \	        V4      hV .V,          # )N)
isinstancelisttuplelen
ValueErrorint)xnmsgs   &&&m/Users/jameslopez/projects/CWCArchive/cwc-podcast/.venv/lib/python3.14/site-packages/mlx/nn/layers/pooling.py_value_or_listr      sS    !dE]##q6Q;S/!Awao37N    c                 x   V^ ,          .p\        VR,          V4       F-  w  rEVP                  WE,          4       VP                  V4       K/  	  VP                  VR,          4       \        V4      ^,
          p^ .\        ^V^4      O\        ^V^4      OVNpV P	                  V4      p V P                  V4      p V # )r      NN)zipappendr   rangereshape	transpose)r   shapewindow_shape	new_shapesw	last_axis
axis_orders   &&&     r    _non_overlapping_sliding_windowsr%      s    q
IE"I|,  - U2YI"IQeAy!,QuQ	1/EQyQJ			)A	JAHr   c                 6   V P                   ^8  d   \        RV P                    R24      hV P                  ^R p\        V4      \        V4      u;8X  d   \        V4      8X  g3   M \        R\        V4       R\        V4       R\        V4       R24      hV P                  p\        ;QJ d&    R \        W1V4       4       F  '       d   K   RM	  R	M! R \        W1V4       4       4      '       d   \        WV4      # \        \        \        \        \        VR,           4      \        P                  4      4      4      4      R
,          pV^ ,          .pT\        W1V4       UUU	u. uF  w  rxp	Wx,
          V	,          ^,           NK  	  up	pp,          pWa,          pWdR,          .,          pVR,          p
T
\        V^R V4       UU	u. uF  w  rW,          NK  	  up	p,          p
W^R ,          p
WRR ,          p
\        P                  ! WV
4      # u up	ppi u up	pi )   zcTo extract sliding windows at least 1 spatial dimension (3 total) is needed but the input only has z dimensions.zTo extract sliding windows the window shapes and strides must have the same number of spatial dimensions as the signal but the signal has z dims and the window shape has z and strides have .c              3   V   "   T F  w  rpW#8H  ;'       d    W,          ^ 8H  x  K!  	  R# 5ir   N ).0sizewindowstrides   &   r   	<genexpr>#_sliding_windows.<locals>.<genexpr>8   s1      $S D& 	//T]a//$Ss   ))FTr   :Nr   NNr   )r   )ndimr   r   r   allr   r%   r
   reversedr   operatormulmx
as_strided)r   r   window_stridesspatial_dimsr   stridesfinal_shaper-   r.   r/   final_strides	og_strides   &&&         r   _sliding_windowsr?   '   s   vvz::;&&O
 	

 771R=L\!2Ic.6II|$%%DSEVDW X  #N 34A7
 	
 GGE
s $'N$Ssss $'N$S   0,GG8DHUT\,BHLL!QRSTUWXG 8*K$'N$S$S D& 
6!A%%$S K K"I;K BKM47"~4V4V0y	4V M Qr]"MRS\!M==77s   #"H
Hc                   >   a a ] tR t^Tt oV 3R ltR tR tRtVtV ;t	# )_Poolc                   < \         SV `  4        Wn        W n        W0n        W@n        WPn        \        \        \        V P                  4      ) ^,
          R^4      4      V n
        R# )r   Nr   )super__init___pooling_function_kernel_size_stride_padding_padding_valuer   r   r   _axes)selfpooling_functionkernel_sizer/   paddingpadding_value	__class__s   &&&&&&r   rD   _Pool.__init__U   sR    !1'+5#d&7&7"8!81!<b!DE
r   c                    \        V P                  4      p\        V P                  4      p\         ;QJ d    . R  V P                   4       F  NK  	  5M! R  V P                   4       4      pRV RV RV 2# )c              3   2   "   T F  q^ ,          x  K  	  R# 5ir*   r+   r,   ps   & r   r0   $_Pool._extra_repr.<locals>.<genexpr>b   s     /AQ44s   zkernel_size=z	, stride=z
, padding=)r   rF   rG   rH   )rK   ksstpds   &   r   _extra_repr_Pool._extra_repr_   s_    4$$%4<< U//UU///bT2$j==r   c                   \         ;QJ d&    R  V P                   4       F  '       g   K   RM	  RM! R  V P                   4       4      '       d>   \        P                  ! VR.V P                  ,           R.,           V P                  R7      p\        WP                  V P                  4      pV P                  WP                  4      # )c              3   8   "   T F  q^ ,          ^ 8  x  K  	  R# 5ir*   r+   rT   s   & r   r0   !_Pool.__call__.<locals>.<genexpr>g   s     /Ataxs   TF)constant_values)r   r   )
anyrH   r7   padrI   r?   rF   rG   rE   rJ   )rK   r   s   &&r   __call___Pool.__call__f   s    3//333////4==(F83 $ 3 3A
 Q 1 14<<@%%a44r   )rJ   rF   rH   rI   rE   rG   )
__name__
__module____qualname____firstlineno__rD   rZ   rb   __static_attributes____classdictcell____classcell__rP   __classdict__s   @@r   rA   rA   T   s     F>5 5r   rA   c                   B   a a ] tR t^qt oRV3R lV 3R llltRtVtV ;t# )_Pool1dc          	         < V ^8  d   QhRS[ S[S[S[,          3,          RS[S[ S[S[S[,          3,          ,          RS[ S[S[S[,          3,          /#    rM   r/   rN   r   r   r   r   )formatrl   s   "r   __annotate___Pool1d.__annotate__r   s\     X X 3c
?+	X
 sE#J/0X sE#J'Xr   c                B  < \        V 4      P                  pR p\        V^VP                  VR4      4      pVe   \        V^VP                  VR4      4      pMTp\        V^VP                  VR4      4      pV Uu. uF  qV3NK  	  pp\        S	V `  WWEV4       R# u upi )z<[{}] '{}' must be an integer or a tuple containing 1 integerrM   Nr/   rN   typerd   r   rs   rC   rD   
rK   rL   rO   rM   r/   rN   
class_namer   rU   rP   s
   &&&&&&   r   rD   _Pool1d.__init__r   s     $Z((
L$CJJz=A
 #FAszz*h/OPF F !SZZ
I-NO#*+7aq67+)W ,   ;Br+   Nr   rd   re   rf   rg   rD   rh   ri   rj   rk   s   @@r   rn   rn   q        X X Xr   rn   c                   B   a a ] tR t^t oRV3R lV 3R llltRtVtV ;t# )_Pool2dc                   < V ^8  d   QhRS[ S[S[S[S[3,          3,          RS[S[ S[S[S[S[3,          3,          ,          RS[S[ S[S[S[S[3,          3,          ,          /# rp   rr   )rs   rl   s   "r   rt   _Pool2d.__annotate__   sp     X X 3c3h/0	X
 sE#s(O345X %U38_ 456Xr   c                B  < \        V 4      P                  pR p\        V^VP                  VR4      4      pVe   \        V^VP                  VR4      4      pMTp\        V^VP                  VR4      4      pV Uu. uF  qV3NK  	  pp\        S	V `  WWEV4       R# u upi )z=[{}] '{}' must be an integer or a tuple containing 2 integersrM   Nr/   rN   rw   ry   s
   &&&&&&   r   rD   _Pool2d.__init__        $Z((
M$CJJz=A
 #FAszz*h/OPF F !SZZ
I-NO#*+7aq67+)W ,r|   r+   r}   r~   rk   s   @@r   r   r      r   r   r   c                   B   a a ] tR t^t oRV3R lV 3R llltRtVtV ;t# )_Pool3dc                   < V ^8  d   QhRS[ S[S[S[S[S[3,          3,          RS[S[ S[S[S[S[S[3,          3,          ,          RS[S[ S[S[S[S[S[3,          3,          ,          /# rp   rr   )rs   rl   s   "r   rt   _Pool3d.__annotate__   sy     X X 3c3m 445	X
 sE#sC-$889:X %U3S=%9 9:;Xr   c                B  < \        V 4      P                  pR p\        V^VP                  VR4      4      pVe   \        V^VP                  VR4      4      pMTp\        V^VP                  VR4      4      pV Uu. uF  qV3NK  	  pp\        S	V `  WWEV4       R# u upi )z=[{}] '{}' must be an integer or a tuple containing 3 integersrM   Nr/   rN   rw   ry   s
   &&&&&&   r   rD   _Pool3d.__init__   r   r|   r+   r}   r~   rk   s   @@r   r   r      r   r   r   c                   F   a a ] tR t^t oRtRV3R lV 3R llltRtVtV ;t# )	MaxPool1da  Applies 1-dimensional max pooling.

Spatially downsamples the input by taking the maximum of a sliding window
of size ``kernel_size`` and sliding stride ``stride``.

Args:
    kernel_size (int or tuple(int)): The size of the pooling window kernel.
    stride (int or tuple(int), optional): The stride of the pooling window.
        Default: ``kernel_size``.
    padding (int or tuple(int), optional): How much negative infinity
        padding to apply to the input. The padding amount is applied to
        both sides of the spatial axis. Default: ``0``.

Examples:
    >>> import mlx.core as mx
    >>> import mlx.nn.layers as nn
    >>> x = mx.random.normal(shape=(4, 16, 5))
    >>> pool = nn.MaxPool1d(kernel_size=2, stride=2)
    >>> pool(x)
c          	         < V ^8  d   QhRS[ S[S[S[,          3,          RS[S[ S[S[S[,          3,          ,          RS[ S[S[S[,          3,          /# rp   rr   )rs   rl   s   "r   rt   MaxPool1d.__annotate__   sZ     N N3c
?+N sE#J/0N sE#J'	Nr   c                \   < \         SV `  \        P                  \	        R 4      ) WV4       R# infNrC   rD   r7   maxfloatrK   rM   r/   rN   rP   s   &&&&r   rD   MaxPool1d.__init__   "     	%,WMr   r+   r}   	rd   re   rf   rg   __doc__rD   rh   ri   rj   rk   s   @@r   r   r      s     *N N Nr   r   c                   F   a a ] tR t^t oRtRV3R lV 3R llltRtVtV ;t# )	AvgPool1da  Applies 1-dimensional average pooling.

Spatially downsamples the input by taking the average of a sliding window
of size ``kernel_size`` and sliding stride ``stride``.

Args:
    kernel_size (int or tuple(int)): The size of the pooling window kernel.
    stride (int or tuple(int), optional): The stride of the pooling window.
        Default: ``kernel_size``.
    padding (int or tuple(int), optional): How much zero padding to apply to
        the input. The padding amount is applied to both sides of the spatial
        axis. Default: ``0``.

Examples:
    >>> import mlx.core as mx
    >>> import mlx.nn.layers as nn
    >>> x = mx.random.normal(shape=(4, 16, 5))
    >>> pool = nn.AvgPool1d(kernel_size=2, stride=2)
    >>> pool(x)
c          	         < V ^8  d   QhRS[ S[S[S[,          3,          RS[S[ S[S[S[,          3,          ,          RS[ S[S[S[,          3,          /# rp   rr   )rs   rl   s   "r   rt   AvgPool1d.__annotate__   sZ     C C3c
?+C sE#J/0C sE#J'	Cr   c                H   < \         SV `  \        P                  ^ WV4       R# r*   rC   rD   r7   meanr   s   &&&&r   rD   AvgPool1d.__init__        	!['Br   r+   r}   r   rk   s   @@r   r   r      s     *C C Cr   r   c                   F   a a ] tR t^t oRtRV3R lV 3R llltRtVtV ;t# )	MaxPool2da1  Applies 2-dimensional max pooling.

Spatially downsamples the input by taking the maximum of a sliding window
of size ``kernel_size`` and sliding stride ``stride``.

The parameters ``kernel_size``, ``stride``, and ``padding`` can either be:

* a single ``int`` -- in which case the same value is used for both the
  height and width axis.
* a ``tuple`` of two ``int`` s -- in which case, the first ``int`` is
  used for the height axis, the second ``int`` for the width axis.

Args:
    kernel_size (int or tuple(int, int)): The size of the pooling window.
    stride (int or tuple(int, int), optional): The stride of the pooling
        window. Default: ``kernel_size``.
    padding (int or tuple(int, int), optional): How much negative infinity
        padding to apply to the input. The padding is applied on both sides
        of the height and width axis. Default: ``0``.

Examples:
    >>> import mlx.core as mx
    >>> import mlx.nn.layers as nn
    >>> x = mx.random.normal(shape=(8, 32, 32, 4))
    >>> pool = nn.MaxPool2d(kernel_size=2, stride=2)
    >>> pool(x)
c                   < V ^8  d   QhRS[ S[S[S[S[3,          3,          RS[S[ S[S[S[S[3,          3,          ,          RS[S[ S[S[S[S[3,          3,          ,          /# rp   rr   )rs   rl   s   "r   rt   MaxPool2d.__annotate__  sn     N N3c3h/0N sE#s(O345N %U38_ 456	Nr   c                \   < \         SV `  \        P                  \	        R 4      ) WV4       R# r   r   r   s   &&&&r   rD   MaxPool2d.__init__  r   r   r+   r}   r   rk   s   @@r   r   r      s     8N N Nr   r   c                   F   a a ] tR tRt oRtRV3R lV 3R llltRtVtV ;t# )	AvgPool2di  a(  Applies 2-dimensional average pooling.

Spatially downsamples the input by taking the average of a sliding window
of size ``kernel_size`` and sliding stride ``stride``.

The parameters ``kernel_size``, ``stride``, and ``padding`` can either be:

* a single ``int`` -- in which case the same value is used for both the
  height and width axis.
* a ``tuple`` of two ``int`` s -- in which case, the first ``int`` is
  used for the height axis, the second ``int`` for the width axis.

Args:
    kernel_size (int or tuple(int, int)): The size of the pooling window.
    stride (int or tuple(int, int), optional): The stride of the pooling
        window. Default: ``kernel_size``.
    padding (int or tuple(int, int), optional): How much zero
        padding to apply to the input. The padding is applied on both sides
        of the height and width axis. Default: ``0``.

Examples:
    >>> import mlx.core as mx
    >>> import mlx.nn.layers as nn
    >>> x = mx.random.normal(shape=(8, 32, 32, 4))
    >>> pool = nn.AvgPool2d(kernel_size=2, stride=2)
    >>> pool(x)
c                   < V ^8  d   QhRS[ S[S[S[S[3,          3,          RS[S[ S[S[S[S[3,          3,          ,          RS[S[ S[S[S[S[3,          3,          ,          /# rp   rr   )rs   rl   s   "r   rt   AvgPool2d.__annotate__:  sn     C C3c3h/0C sE#s(O345C %U38_ 456	Cr   c                H   < \         SV `  \        P                  ^ WV4       R# r*   r   r   s   &&&&r   rD   AvgPool2d.__init__:  r   r   r+   r}   r   rk   s   @@r   r   r     s     8C C Cr   r   c                   F   a a ] tR tRt oRtRV3R lV 3R llltRtVtV ;t# )	MaxPool3diC  a|  Applies 3-dimensional max pooling.

Spatially downsamples the input by taking the maximum of a sliding window
of size ``kernel_size`` and sliding stride ``stride``.

The parameters ``kernel_size``, ``stride``, and ``padding`` can either be:

* a single ``int`` -- in which case the same value is used for the depth,
  height, and width axis.
* a ``tuple`` of three ``int`` s -- in which case, the first ``int`` is used
  for the depth axis, the second ``int`` for the height axis, and the third
  ``int`` for the width axis.

Args:
    kernel_size (int or tuple(int, int, int)): The size of the pooling window.
    stride (int or tuple(int, int, int), optional): The stride of the pooling
        window. Default: ``kernel_size``.
    padding (int or tuple(int, int, int), optional): How much negative infinity
        padding to apply to the input. The padding is applied on both sides
        of the depth, height and width axis. Default: ``0``.

Examples:
    >>> import mlx.core as mx
    >>> import mlx.nn.layers as nn
    >>> x = mx.random.normal(shape=(8, 16, 32, 32, 4))
    >>> pool = nn.MaxPool3d(kernel_size=2, stride=2)
    >>> pool(x)
c                   < V ^8  d   QhRS[ S[S[S[S[S[3,          3,          RS[S[ S[S[S[S[S[3,          3,          ,          RS[S[ S[S[S[S[S[3,          3,          ,          /# rp   rr   )rs   rl   s   "r   rt   MaxPool3d.__annotate__a  sw     N N3c3m 445N sE#sC-$889:N %U3S=%9 9:;	Nr   c                \   < \         SV `  \        P                  \	        R 4      ) WV4       R# r   r   r   s   &&&&r   rD   MaxPool3d.__init__a  r   r   r+   r}   r   rk   s   @@r   r   r   C  s     :N N Nr   r   c                   F   a a ] tR tRt oRtRV3R lV 3R llltRtVtV ;t# )	AvgPool3dij  as  Applies 3-dimensional average pooling.

Spatially downsamples the input by taking the average of a sliding window
of size ``kernel_size`` and sliding stride ``stride``.

The parameters ``kernel_size``, ``stride``, and ``padding`` can either be:

* a single ``int`` -- in which case the same value is used for the depth,
  height, and width axis.
* a ``tuple`` of three ``int`` s -- in which case, the first ``int`` is used
  for the depth axis, the second ``int`` for the height axis, and the third
  ``int`` for the width axis.

Args:
    kernel_size (int or tuple(int, int, int)): The size of the pooling window.
    stride (int or tuple(int, int, int), optional): The stride of the pooling
        window. Default: ``kernel_size``.
    padding (int or tuple(int, int, int), optional): How much zero
        padding to apply to the input. The padding is applied on both sides
        of the depth, height and width axis. Default: ``0``.

Examples:
    >>> import mlx.core as mx
    >>> import mlx.nn.layers as nn
    >>> x = mx.random.normal(shape=(8, 16, 32, 32, 4))
    >>> pool = nn.AvgPool3d(kernel_size=2, stride=2)
    >>> pool(x)
c                   < V ^8  d   QhRS[ S[S[S[S[S[3,          3,          RS[S[ S[S[S[S[S[3,          3,          ,          RS[S[ S[S[S[S[S[3,          3,          ,          /# rp   rr   )rs   rl   s   "r   rt   AvgPool3d.__annotate__  sw     C C3c3m 445C sE#sC-$889:C %U3S=%9 9:;	Cr   c                H   < \         SV `  \        P                  ^ WV4       R# r*   r   r   s   &&&&r   rD   AvgPool3d.__init__  r   r   r+   r}   r   rk   s   @@r   r   r   j  s     :C C Cr   r   )r5   	itertoolsr   typingr   r   r   mlx.corecorer7   mlx.nn.layers.baser   r   r%   r?   rA   rn   r   r   r   r   r   r   r   r   r+   r   r   <module>r      s       ) )  %	 *8Z5F 5:Xe X0Xe X0Xe X0N N>C C>#N #NL#C #CL$N $NN$C $Cr   