+
    &jw                     t    ^ RI t ^ RIHt ^ RIHt ^ RIHt  ! R R]4      t ! R R]4      t	 ! R R	]4      t
R# )
    N)Union)Modulec                   R   a a ] tR t^
t oRtRV3R lV 3R llltR tR tRtVt	V ;t
# )Conv1dan  Applies a 1-dimensional convolution over the multi-channel input sequence.

The channels are expected to be last i.e. the input shape should be ``NLC`` where:

* ``N`` is the batch dimension
* ``L`` is the sequence length
* ``C`` is the number of input channels

Args:
    in_channels (int): The number of input channels
    out_channels (int): The number of output channels
    kernel_size (int): The size of the convolution filters
    stride (int, optional): The stride when applying the filter.
        Default: ``1``.
    padding (int, optional): How many positions to 0-pad the input with.
        Default: ``0``.
    dilation (int, optional): The dilation of the convolution.
    groups (int, optional): The number of groups for the convolution.
        Default: ``1``.
    bias (bool, optional): If ``True`` add a learnable bias to the output.
        Default: ``True``
c                J   < V ^8  d   QhRS[ RS[ RS[ RS[ RS[ RS[ RS[ RS[/# 	   in_channelsout_channelskernel_sizestridepaddingdilationgroupsbias)intbool)format__classdict__s   "q/Users/jameslopez/projects/CWCArchive/cwc-podcast/.venv/lib/python3.14/site-packages/mlx/nn/layers/convolution.py__annotate__Conv1d.__annotate__"   s[        	
         c	                  < \         S
V `  4        W,          ^ 8w  d   \        RV RV R24      h\        P                  ! ^W,          ,          4      p	\
        P                  P                  V	) V	W#W,          3R7      V n        V'       d   \
        P                  ! V34      V n
        WPn        W`n        W@n        Wpn        R# )r   The number of input channels (-) must be divisible by the number of groups ()lowhighshapeN)super__init__
ValueErrormathsqrtmxrandomuniformweightzerosr   r   r   r   r   selfr
   r   r   r   r   r   r   r   scale	__class__s   &&&&&&&&& r   r#   Conv1d.__init__"   s     	1$0 >66<XQ@ 
 		!{89:ii''k.CD ( 

 ,1DI r   c                N   V P                   P                  R	,          V P                  ,           RV P                   P                  ^ ,           RV P                   P                  ^,           RV P                   RV P                   RV P
                   RV P                   RRV 9    2# )
   , , kernel_size=	, stride=
, padding=, dilation=	, groups=, bias=r   r*   r!   r   r   r   r   r-   s   &r   _extra_reprConv1d._extra_reprC   s    {{  $t{{232dkk6G6G6J5K L;;,,Q/0	$++ G||nK ?kk] #dN#	%	
r   c                    \         P                  ! WP                  V P                  V P                  V P
                  V P                  4      pR V 9   d   W P                  ,           pV# r   )r'   conv1dr*   r   r   r   r   r   r-   xys   && r   __call__Conv1d.__call__L   H    II{{DKKt}}dkk
 T>IIAr   r   r   r   r   r   r*   r2   r   r2   r2   T__name__
__module____qualname____firstlineno____doc__r#   r=   rE   __static_attributes____classdictcell____classcell__r/   r   s   @@r   r   r   
   s$     . B
 r   r   c                   R   a a ] tR t^Ut oRtRV3R lV 3R llltR tR tRtVt	V ;t
# )Conv2da  Applies a 2-dimensional convolution over the multi-channel input image.

The channels are expected to be last i.e. the input shape should be ``NHWC`` where:

* ``N`` is the batch dimension
* ``H`` is the input image height
* ``W`` is the input image width
* ``C`` is the number of input channels

Args:
    in_channels (int): The number of input channels.
    out_channels (int): The number of output channels.
    kernel_size (int or tuple): The size of the convolution filters.
    stride (int or tuple, optional): The size of the stride when
        applying the filter. Default: ``1``.
    padding (int or tuple, optional): How many positions to 0-pad
        the input with. Default: ``0``.
    dilation (int or tuple, optional): The dilation of the convolution.
    groups (int, optional): The number of groups for the convolution.
        Default: ``1``.
    bias (bool, optional): If ``True`` add a learnable bias to the
        output. Default: ``True``
c                   < V ^8  d   QhRS[ RS[ RS[S[ S[3,          RS[S[ S[3,          RS[S[ S[3,          RS[S[ S[3,          RS[ RS[/# r   r   r   tupler   )r   r   s   "r   r   Conv2d.__annotate__n   s     # ## # 3:&	#
 c5j!# sEz"# U
## # #r   c	                  < \         S
V `  4        W,          ^ 8w  d   \        RV RV R24      h\        R W4V34      w  r4p\        P
                  ! ^W^ ,          ,          V^,          ,          ,          4      p	\        P                  P                  V	) V	V.VOW,          N5R7      V n	        V'       d   \        P                  ! V34      V n        WPn        W@n        W`n        Wpn        R# )r   r   r   r   c                 8    \        V \        4      '       d   W 3# T # N
isinstancer   rC   s   &r   <lambda>!Conv2d.__init__.<locals>.<lambda>   s    
1c 2 2qf99r   r   N)r"   r#   r$   mapr%   r&   r'   r(   r)   r*   r+   r   r   r   r   r   r,   s   &&&&&&&&& r   r#   Conv2d.__init__n   s     	1$0 >66<XQ@ 
 (+9'*(
$W 		!{^;k!nLMNii''E+E{/DE ( 

 ,1DI r   c                N   V P                   P                  R
,          V P                  ,           RV P                   P                  ^ ,           RV P                   P                  R,           RV P                   RV P                   RV P
                   RV P                   RR	V 9    2# )r2   r3   r4   :r2      Nr5   r6   r7   r8   r9   r   r:   r;   r<   s   &r   r=   Conv2d._extra_repr   s    {{  $t{{232dkk6G6G6J5K L;;,,S12)DKK= I||nK ?kk] #dN#	%	
r   c                    \         P                  ! WP                  V P                  V P                  V P
                  V P                  4      pR V 9   d   W P                  ,           pV# r@   )r'   conv2dr*   r   r   r   r   r   rB   s   && r   rE   Conv2d.__call__   rG   r   rH   rI   rJ   rS   s   @@r   rU   rU   U   s$     0# #J
 r   rU   c                   R   a a ] tR t^t oRtRV3R lV 3R llltR tR tRtVt	V ;t
# )Conv3da  Applies a 3-dimensional convolution over the multi-channel input image.

The channels are expected to be last i.e. the input shape should be ``NDHWC`` where:

* ``N`` is the batch dimension
* ``D`` is the input image depth
* ``H`` is the input image height
* ``W`` is the input image width
* ``C`` is the number of input channels

Args:
    in_channels (int): The number of input channels.
    out_channels (int): The number of output channels.
    kernel_size (int or tuple): The size of the convolution filters.
    stride (int or tuple, optional): The size of the stride when
        applying the filter. Default: ``1``.
    dilation (int or tuple, optional): The dilation of the convolution.
    padding (int or tuple, optional): How many positions to 0-pad
        the input with. Default: ``0``.
    bias (bool, optional): If ``True`` add a learnable bias to the
        output. Default: ``True``
c                   < V ^8  d   QhRS[ RS[ RS[S[ S[3,          RS[S[ S[3,          RS[S[ S[3,          RS[S[ S[3,          RS[/# )r	   r
   r   r   r   r   r   r   rW   )r   r   s   "r   r   Conv3d.__annotate__   su     ! !! ! 3:&	!
 c5j!! sEz"! U
#! !r   c                  < \         S	V `  4        \        R  W4V34      w  r4p\        P                  ! ^W^ ,          ,          V^,          ,          V^,          ,          ,          4      p\
        P                  P                  V) VV.VOVN5R7      V n        V'       d   \
        P                  ! V34      V n
        WPn        W@n        W`n        R# )c                 :    \        V \        4      '       d   W V 3# T # r\   r]   r_   s   &r   r`   !Conv3d.__init__.<locals>.<lambda>   s    :a#5#5qQi<1<r   r   N)r"   r#   rb   r%   r&   r'   r(   r)   r*   r+   r   r   r   r   )
r-   r
   r   r   r   r   r   r   r.   r/   s
   &&&&&&&& r   r#   Conv3d.__init__   s     	'*<'*(
$W 		1~-A>QOP
 ii'';+;{; ( 

 ,1DI r   c                4   V P                   P                  R	,          V P                  ,           RV P                   P                  ^ ,           RV P                   P                  R,           RV P                   RV P                   RV P
                   RRV 9    2# )
r2   r3   r4   :r2      Nr5   r6   r7   r9   r   r:   r;   r<   s   &r   r=   Conv3d._extra_repr   s    {{  $t{{232dkk6G6G6J5K L;;,,S12)DKK= I||nK ?dN#%	
r   c                    \         P                  ! WP                  V P                  V P                  V P
                  4      pR V 9   d   W P                  ,           pV# r@   )r'   conv3dr*   r   r   r   r   rB   s   && r   rE   Conv3d.__call__   s=    IIadkk4<<OT>IIAr   )r   r   r   r   r*   )r2   r   r2   TrJ   rS   s   @@r   rk   rk      s#     .! !>
 r   rk   )r%   typingr   mlx.corecorer'   mlx.nn.layers.baser   r   rU   rk    r   r   <module>r}      s?       %HV HVMV M`CV Cr   