
    9iu                        d Z ddlZddlmZmZmZ ddlmZm	Z	m
Z
 ddlmZ ddlmZmZmZmZ ddlmZ ddlZdd	lmZmZ  G d
 d      Z G d de      Z G d de      Z G d de      Z G d de      Z G d de      Z G d de      Z G d de      Z G d de      Z  G d de      Z! G d de      Z" G d  d!e      Z# G d" d#e      Z$y)$zO
ConvMLP Layers
==============
Vedic-inspired optimized layer implementations.
    N)BaseEstimatorClassifierMixinRegressorMixin)	check_X_ycheck_arraycheck_is_fitted)unique_labels)TupleOptionalUnionList)	lru_cache   )im2colcol2imc                   6    e Zd ZdZd Zd	dZd Zd Zd Zd Z	y)
Layerz)Base class for all neural network layers.c                      d| _         d| _        y )NTF)	trainablebuiltselfs    </home/per/Documents/sklearn ConvMLP module/convmlp/layers.py__init__zLayer.__init__   s    
    c                     t         NNotImplementedErrorr   xtrainings      r   forwardzLayer.forward       !!r   c                     t         r   r   r   grad_outputs     r   backwardzLayer.backward    r$   r   c                 0    || _         || _        d| _        |S )z<Default no-op build for layers without trainable parameters.T)input_shapeoutput_shaper   r   r*   s     r   buildzLayer.build#   s    &'
r   c                     i S r    r   s    r   
get_paramszLayer.get_params*   s    	r   c                 j    |j                         D ]  \  }}t        | |      st        | ||       ! | S r   )itemshasattrsetattr)r   paramskeyvalues       r   
set_paramszLayer.set_params-   s8     ,,. 	*JCtS!c5)	* r   NF)
__name__
__module____qualname____doc__r   r#   r(   r-   r0   r8   r/   r   r   r   r      s#    3""r   r   c                       e Zd ZdZ	 	 	 	 	 	 d"dedeeeeef   f   dedeeef   dededed	ef fd
Z	deedf   fdZ
d Zdej                  dej                  fdZd Zd#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j                  fdZdej                  dej                  fdZdej                  dej                  fdZdej                  dej                  fdZdej                  dej                  fdZdej                  dej                  fdZdefdZd$dej                  d eej                     fd!Z xZS )%Conv2Da0  
    2D Convolutional Layer with Vedic-optimized computation.
    
    Implements multiple computation strategies:
    - im2col (Vedic block multiplication)
    - FFT convolution (frequency domain)
    - Winograd minimal filtering
    - Nikhilam quantized inference
    
    Parameters
    ----------
    filters : int
        Number of convolution filters (output channels).
    kernel_size : tuple or int
        Size of the convolution kernel.
    strides : int or tuple, default=1
        Stride length.
    padding : str or int, default='same'
        Padding mode: 'same', 'valid', or integer.
    activation : str, default='relu'
        Activation function.
    use_bias : bool, default=True
        Whether to use bias term.
    method : str, default='im2col'
        Computation method: 'im2col', 'fft', 'winograd', 'direct'.
    kernel_initializer : str, default='he_normal'
        Weight initialization method.
    
    Example
    -------
    >>> layer = Conv2D(filters=32, kernel_size=(3, 3), padding='same')
    >>> output = layer.forward(input_array)
    filterskernel_sizestridespadding
activationuse_biasmethodkernel_initializerc	                    t         	|           || _        t        |t              r||fn|| _        t        |t              r||fn|| _        || _        || _        || _	        || _
        || _        d | _        d | _        d | _        i | _        y r   )superr   r@   
isinstanceintrA   rB   rC   rD   rE   rF   rG   kernelbiasr*   _cache)
r   r@   rA   rB   rC   rD   rE   rF   rG   	__class__s
            r   r   zConv2D.__init__[   s     	9CKQT9UK5[f-7-E)7$ "4 	 r   r*   .c                    || _         | j                  dk(  r<t        j                  d|d   t        j                  | j
                        z  z        }n:| j                  dk(  r)t        j                  d|d   | j                  z   z        }nd}t        j                  j                  | j                  |d   g| j
                   j                  t        j                        |z  | _        | j                  r4t        j                  | j                  t        j                        | _        d| _        | j!                  |      | _        | j"                  S )	z(Initialize weights based on input shape.	he_normalg       @glorot_uniformg{Gz?dtypeT)r*   rG   npsqrtprodrA   r@   randomrandnastypefloat32rL   rE   zerosrM   r   _compute_output_shape_output_shape)r   r*   stds      r   r-   zConv2D.buildx   s   & ""k1''#R2774;K;K3L!LMNC$$(88''#R4<<!?@ACCiiooLLO
 
 &
s	# ==RZZ@DI
 "77D!!!r   c                     |dd \  }}| j                   \  }}| j                  \  }}t        | j                  t              r| j                  dk(  rCt        t        j                  ||z              }t        t        j                  ||z              }	nt        t        j                  ||z
  dz   |z              }t        t        j                  ||z
  dz   |z              }	nnt        t        j                  |d| j                  z  z   |z
  dz   |z              }t        t        j                  |d| j                  z  z   |z
  dz   |z              }	t        |      dk(  r|d   | j                  ||	fS | j                  ||	fS )z$Calculate output spatial dimensions.Nsamer         r   )
rA   rB   rJ   rC   strrK   rW   ceillenr@   )
r   r*   hwkhkwshswout_hout_ws
             r   r_   zConv2D._compute_output_shape   sO   231!!BBdllC(||v%BGGAFO,BGGAFO,BGGQVaZ2$567BGGQVaZ2$567 Q%5!5!:Q!>" DEFEQ%5!5!:Q!>" DEFE?B;?OST?TAeU; 	1\\5%0	1r   r!   returnc           	         | j                   dk(  r|S | j                   dk(  rR|j                  dd \  }}| j                  \  }}t        d|dz
  |z  |z
  |z         }t        d|dz
  |z  |z
  |z         }n| j                   x}}|dz  }|dz  }	t	        |j                        dk(  r&t        j                  |d	d	|||z
  f|	||	z
  ffd
      S t        j                  |d	d	|||z
  f|	||	z
  ffd
      S )zApply padding to input.validrd   rc   Nr   r   re   rf   r   r   constantmode)rC   shaperB   maxri   rW   pad)
r   r!   rj   rk   rn   ro   pad_hpad_wpad_toppad_lefts
             r   
_pad_inputzConv2D._pad_input   s   <<7"H<<6!7723<DAq\\FBBFa<!+b01EBFa<!+b01E LL(EE1*A:qww<166!ffw.H&(89;AKM M 66!ffw.H&(89;AKM Mr   c                     | j                   dk(  ry| j                   dk(  r$| j                  d   dz  | j                  d   dz  fS | j                   | j                   fS )zCalculate padding values.rt   ru   rd   r   re   r   )rC   rA   r   s    r   _get_paddingzConv2D._get_padding   s[    <<7"\\V###A&!+T-=-=a-@A-EEE<<--r   r"   c                 *   | j                   s| j                  |j                         || _        | j	                  |      }| j
                  dk(  r| j                  |      }nt| j
                  dk(  r| j                  |      }nS| j
                  dk(  r| j                  |      }n2| j
                  dk(  r| j                  |      }n| j                  |      }| j                  r!|| j                  j                  dddd      z  }| j                  |      }|S )a{  
        Forward pass using the configured method.
        
        Parameters
        ----------
        x : np.ndarray
            Input tensor of shape (batch, channels, height, width)
        training : bool
            Whether in training mode (affects dropout, etc.)
        
        Returns
        -------
        output : np.ndarray
            Convolved output
        r   fftwinograddirectr   rR   )r   r-   ry   _inputr   rF   _forward_im2col_forward_fft_forward_winograd_forward_directrE   rM   reshape_apply_activation)r   r!   r"   x_paddedoutputs        r   r#   zConv2D.forward   s      zzJJqww ??1% ;;("))(3F[[E!&&x0F[[J&++H5F[[H$))(3F))(3F ==dii''2q!44F ''/r   r   c                 8   |j                   \  }}}}| j                  \  }}| j                  \  }}	||z
  |z  dz   }
||z
  |	z  dz   }t        j                  || j
                  |
|ft        j                        }t        |      D ]  }t        | j
                        D ]y  }t        d||z
  dz   |      D ]a  }t        d||z
  dz   |	      D ]I  }||dd|||z   |||z   f   }t        j                  || j                  |   z        |||||z  ||	z  f<   K c {  |S )z'Direct convolution (naive but correct).r   rU   r   N)
ry   rA   rB   rW   r^   r@   r]   rangesumrL   )r   r   batchchannelsrj   rk   rl   rm   rn   ro   rp   rq   r   bfijpatchs                     r   r   zConv2D._forward_direct   sD    (xA!!BBRB"RB"5$,,u=RZZPu 	TA4<<( Tq!b&1*b1 TA"1a"fqj"5 T (Aq2vq2v)= >57VVEDKKPQN<R5Sq!QUArE12TTT	T r   c                    |j                   \  }}}}| j                  \  }}| j                  \  }}	||z
  |z  dz   }
||z
  |	z  dz   }t        || j                  | j                        }| j                  j                  | j                  d      }t        j                  |j                  ddd      |j                        }|j                  ddd      j                  || j                  |
|      }|S )u,  
        Vedic Block Multiplication (im2col) Method.
        
        Inspired by Urdhva-Tiryagbyham (vertically and crosswise).
        Unfold image patches into columns for efficient matrix multiplication.
        
        Time Complexity: O(n²) matrix mult instead of O(n³) nested loops
        r   rR   r   re   )ry   rA   rB   r   rL   r   r@   rW   matmul	transposeT)r   r   r   r   rj   rk   rl   rm   rn   ro   rp   rq   patches
kernel_colr   s                  r   r   zConv2D._forward_im2col  s     !)xA!!BBRB"RB" 4#3#3T\\B [[((r:
 7,,Q15z||D !!!Q*225$,,uUr   c           	         |j                   \  }}}}| j                  \  }}| j                  \  }}	dt        t	        j
                  t	        j                  ||z   dz
                    z  }
dt        t	        j
                  t	        j                  ||z   dz
                    z  }t        j                  j                  ||
|f      }t	        j                  || j                  ||z
  dz   ||z
  dz   ft        j                        }t        | j                        D ]  }t	        j                  |
|ft        j                        }| j                  |   |d|d|f<   t        j                  j                  |      }||z  }t        j                  j                  ||
|f      ddd||z
  dz   d||z
  dz   f   |dd|f<    |dddddd|dd|	f   S )u   
        FFT Convolution Method.
        
        Based on convolution theorem: conv in spatial = multiply in frequency.
        Optimal for large kernels (k >= 7).
        
        Time Complexity: O(n log n) instead of O(n²)
        re   r   )srU   N)ry   rA   rB   rK   rW   rh   log2r   rfft2r^   r@   r]   r   rL   irfft2)r   r   r   r   rj   rk   rl   rm   rn   ro   fft_hfft_wX_fftr   r   kernel_paddedK_fftout_ffts                     r   r   zConv2D._forward_fft-  s    !)xA!!BB SR!!4566SR!!4566 X%85$,,B
AFQJGrzzZt||$ 
	YAHHeU^2::FM&*kk!nM#2#ss(#FFLL/E emG 66==UEN=CAw"QwPWQRSUQUVWQWPWDWXF1a4L
	Y aDbD$B$&''r   c                    | j                   dk7  r&t        j                  d       | j                  |      S t	        j
                  g dg dg dg dgt        j                        }t	        j
                  g dg d	g d
g dg dgt        j                        }t	        j
                  ddgddgddgddggt        j                        }|j                  \  }}}}| j                  \  }	}
|dz
  |	z  dz   }|dz
  |
z  dz   }t	        j                  || j                  ||ft        j                        }t	        j                  d|| j                  d   |      }d}d}t        |      D ]  }t        | j                        D ]  }t        d|dz
  dz   |	dz        D ]  }t        d|dz
  dz   |
dz        D ]  }||dd|||z   |||z   f   }t	        j                  d||j                  d|||      |      }||dddddddf   z  }t	        j                  d|||      }||	z  }||
z  }|d   |||||dz   ||dz   f<       |S )z
        Winograd Minimal Filtering Method.
        
        Reduces multiplication count for small kernels (3x3).
        Uses precomputed transformation matrices.
        
        For 3x3 kernel on 4x4 output: 16 mults -> 4 multiplies.
        )   r   z;Winograd optimized for 3x3 kernels. Falling back to im2col.)r   r   r   )      ?r   r   )r         r   )r   r   r   rU   )r   r   rR   r   )r   r   r   r   )r   rR   r   r   )r   r   r   r   )r   r   r   rR   r   r   rR   r   zij,jk,lk->ilrf   re   Nzij,bchw,kw->bcijkzij,bcijk,kj->bcklru   )rA   warningswarnr   rW   arrayr]   ry   rB   r^   r@   einsumrL   r   r   )r   r   GBAr   r   rj   rk   rn   ro   rp   rq   r   G_kerneltile_htile_wr   r   r   r   tileB_tilemresultout_iout_js                              r   r   zConv2D._forward_winogradQ  s    v%MMWX''11 HH	

  HH
  HHFFGG	

  !)xABQ2!Q2!5$,,u=RZZP 99^QAB u 	RA4<<( Rq!a%!)R!V4 RA"1a!eaia8 R'1a&j!AfH*(DE "$+>4<<PQS[]cekClno!p #XdAq!m%<< "$+>1a!H !"R !REKD\q!U57]E%'MAB!RRR	R* r   c           	         | j                   dk(  rt        j                  d|      S | j                   dk(  r1ddt        j                  t        j                  |dd             z   z  S | j                   dk(  rt        j
                  |      S | j                   dk(  rIt        j                  |t        j                  |dd	
      z
        }|t        j                  |dd	
      z  S | j                   | j                   dk(  r|S |S )zApply activation function.relur   sigmoidr     tanhsoftmaxTaxiskeepdimslinear)rD   rW   maximumexpclipr   rz   r   )r   r!   exp_xs      r   r   zConv2D._apply_activation  s    ??f$::a##__	)BFFBGGAtS$9#9::;;__&771:__	)FF1rvvaa$??@E266%a$???__$8(CHHr   r'   c                 &   | j                  |      }| j                  rt        j                  |d      }| j	                  | j
                        }|j                  \  }}}}| j                  \  }}	| j                  \  }
}|j                  dd \  }}|j                  dddd      j                  || j                  d	      }t        || j                  | j                        }t        j                  ||j                  ddd            }|j                  ddd      j                  | j                  |||	      }| j                  j                  | j                  d	      }t        j                  ||j                        }t!        |||||f| j                  | j                        }| j#                         \  }}|dkD  s|dkD  rOt%        |j                        d
k(  r|dddd|| xs d|| xs df   }|S |dd|| xs d|| xs df   }|S |}|S )zv
        Backward pass for gradient computation.
        
        Uses the same im2col method for efficiency.
        r   re   r   r   rc   Nr   r   r   re   rR   rf   )_activation_gradrE   rW   r   r   r   ry   rA   rB   r   r   r@   r   r   rL   r   r   r   ri   )r   r'   	grad_biasr   r   r   rj   rk   rl   rm   rn   ro   rp   rq   grad_out_reshapedr   grad_kernelgrad_kernel_colgrad_patchesgrad_input_paddedr|   r}   
grad_inputs                          r   r(   zConv2D.backward  s    ++K8 =={;I ??4;;/ (xA!!BB"((-u (11!Q1=EEeT\\[]^ 4#3#3T\\Bii 173D3DQ13MN!++Aq!4<<T\\8UWY[\ ++--dllB?yy!2O4E4EF"<%1a1H$($4$4dllD ((*u19	$**+q0.q!UE6>T5I5RWQWQ_[_K_/_`
 	 /q%$2FufnX\H\/\]
  +Jr   c                 n    | j                   dk(  r|| j                  dkD  z  S | j                   dv r|dz  S |S )z-Compute gradient through activation function.r   r   )r   r   g      ?)rD   r   r&   s     r   r   zConv2D._activation_grad  s?    ??f$$++/22__ 33%%r   c           
          | j                   | j                  | j                  | j                  | j                  | j
                  | j                  | j                  | j                  d	S )N	r@   rA   rB   rC   rD   rE   rF   rL   rM   r   r   s    r   r0   zConv2D.get_params  sM    ||++||||//kkkkII

 
	
r   rL   rM   c                     |j                  t        j                        | _        |0| j                  r$|j                  t        j                        | _        d| _        y)zSet layer weights.NT)r\   rW   r]   rL   rE   rM   r   )r   rL   rM   s      r   set_weightszConv2D.set_weights  s=    mmBJJ/BJJ/DI
r   )r   rd   r   Tr   rQ   r9   r   )r:   r;   r<   r=   rK   r   r
   rg   boolr   r-   r_   rW   ndarrayr   r   r#   r   r   r   r   r   r(   r   dictr0   r   r   __classcell__rO   s   @r   r?   r?   8   s    L #) "- 3c3h/0 	
 sCx     :"sCx "61*MBJJ M2:: M..+ +t +

 +Z

 rzz (

 rzz >"(RZZ "(BJJ "(HG"** G GR2:: "**  ,BJJ ,2:: ,\BJJ 2:: 
D 
"** HRZZ4H r   r?   c                        e Zd ZdZddedee   def fdZdeedf   fdZdeedf   fd	Z	dd
e
j                  dede
j                  fdZde
j                  de
j                  fdZ xZS )	MaxPool2Dz3Max pooling layer with Vedic-optimized computation.	pool_sizerB   rC   c                    t         |           t        |t              r||fn|| _        |xs | j                  | _        t        | j
                  t              r| j
                  | j
                  fn| j
                  | _        || _        y r   rI   r   rJ   rK   r   rB   rC   r   r   rB   rC   rO   s       r   r   zMaxPool2D.__init__  h    3=i3M)Y/S\0$..7A$,,PS7Tdll3Z^ZfZfr   r*   .c                 b    || _         | j                  |      | _        d| _        | j                  S NTr*   r_   r+   r   r,   s     r   r-   zMaxPool2D.build  0    & 66{C
   r   c                     |dd  \  }}| j                   \  }}| j                  \  }}||z
  d| j                  z  z   |z  dz   }||z
  d| j                  z  z   |z  dz   }	t        |      dk(  r|d   |d   ||	fS |d   ||	fS Nrc   re   r   rf   r   r   rB   rC   ri   
r   r*   rj   rk   phpwrn   ro   rp   rq   s
             r   r_   zMaxPool2D._compute_output_shape      231BBR!dll**r1A5R!dll**r1A5{q NKNE5AAAu--r   r!   r"   rr   c           	         || _         |j                  \  }}}}| j                  \  }}| j                  \  }	}
||z
  d| j                  z  z   |	z  dz   }||z
  d| j                  z  z   |
z  dz   }| j                  dkD  rHt        j                  |dd| j                  | j                  f| j                  | j                  ffd      }t        j                  ||||ft
        j                        }t        |      D ]U  }t        |      D ]E  }||	z  ||
z  }}|d d d d |||z   |||z   f   }t        j                  |d	      |d d d d ||f<   G W |S 
Nre   r   r   ru   rv   rw   rU   re   r   r   )r   ry   r   rB   rC   rW   r{   r^   r]   r   rz   r   r!   r"   r   r   rj   rk   r   r   rn   ro   rp   rq   r   r   r   h_startw_startr   s                      r   r#   zMaxPool2D.forward  sb    !xABBR!dll**r1A5R!dll**r1A5<<!q66DLL$,,+G<<68>HJA 5(E59Lu 	@A5\ @#$r61r6!Q
 2GGBJ4FFG%'VVE%?q!Qz"@	@ r   r'   c                    |j                   \  }}}}| j                  \  }}| j                  \  }}	t        j                  | j
                        }
t        |      D ]  }t        |      D ]  }||z  ||	z  }}| j
                  d d d d |||z   |||z   f   }|t        j                  |dd      k(  }t        |      D ]<  }t        |      D ],  }|
|||||z   |||z   fxx   |||f   |||||f   z  z  cc<   . >   |
S )Nr   Tr   )ry   r   rB   rW   
zeros_liker   r   rz   )r   r'   r   r   rp   rq   r   r   rn   ro   r   r   r   r   r   r   max_maskr   cs                      r   r(   zMaxPool2D.backward-  s<   (3(9(9%xBB]]4;;/
u 	EA5\ E#$r61r6 Aq''"**<ggbj>P$PQ!RVVE%NN u EA"8_ E"1a);WWRZ=O#OP$QTN[Aq!-DDEPEEE	E r   re   Nr   r9   r:   r;   r<   r=   rK   r   r   r
   r-   r_   rW   r   r   r#   r(   r   r   s   @r   r   r     s    =# HSM SV !sCx !	.sCx 	. t 

 .BJJ 2:: r   r   c                        e Zd ZdZddedee   def fdZdeedf   fdZdeedf   fd	Z	dd
e
j                  dede
j                  fdZde
j                  de
j                  fdZ xZS )	AvgPool2DzAverage pooling layer.r   rB   rC   c                    t         |           t        |t              r||fn|| _        |xs | j                  | _        t        | j
                  t              r| j
                  | j
                  fn| j
                  | _        || _        y r   r   r   s       r   r   zAvgPool2D.__init__H  r   r   r*   .c                 b    || _         | j                  |      | _        d| _        | j                  S r   r   r,   s     r   r-   zAvgPool2D.buildO  r   r   c                     |dd  \  }}| j                   \  }}| j                  \  }}||z
  d| j                  z  z   |z  dz   }||z
  d| j                  z  z   |z  dz   }	t        |      dk(  r|d   |d   ||	fS |d   ||	fS r   r   r   s
             r   r_   zAvgPool2D._compute_output_shapeU  r   r   r!   r"   rr   c           	         || _         |j                  \  }}}}| j                  \  }}| j                  \  }	}
||z
  d| j                  z  z   |	z  dz   }||z
  d| j                  z  z   |
z  dz   }| j                  dkD  rHt        j                  |dd| j                  | j                  f| j                  | j                  ffd      }t        j                  ||||ft
        j                        }t        |      D ]U  }t        |      D ]E  }||	z  ||
z  }}|d d d d |||z   |||z   f   }t        j                  |d	      |d d d d ||f<   G W |S r   )r   ry   r   rB   rC   rW   r{   r^   r]   r   meanr   s                      r   r#   zAvgPool2D.forward`  sb    !xABBR!dll**r1A5R!dll**r1A5<<!q66DLL$,,+G<<68>HJA 5(E59Lu 	AA5\ A#$r61r6!Q
 2GGBJ4FFG%'WWU%@q!Qz"A	A r   r'   c           
      ^   |j                   \  }}}}| j                  \  }}| j                  \  }}	t        j                  | j
                        }
||z  }t        |      D ]K  }t        |      D ];  }||z  ||	z  }}|d d d d ||dz   ||dz   f   |z  |
d d d d |||z   |||z   f<   = M |
S Nr   )ry   r   rB   rW   r   r   r   )r   r'   r   r   rp   rq   r   r   rn   ro   r   r   r   r   r   r   s                   r   r(   zAvgPool2D.backwardw  s    (3(9(9%xBB]]4;;/
G	u 	@A5\ @#$r61r61a!eQqsU 23i? 1a!3WWRZ5GGH@	@ r   r   r9   r   r   s   @r   r  r  E  s     # HSM SV !sCx !	.sCx 	. t 

 .BJJ 2:: r   r  c                        e Zd ZdZ fdZdeedf   fdZddej                  de
dej                  fd	Zd
ej                  dej                  fdZ xZS )Flattenz*Flatten layer to convert 4D to 2D for MLP.c                 >    t         |           d | _        d | _        y r   )rI   r   r*   r+   r   rO   s    r   r   zFlatten.__init__  s     r   r*   .c                     || _         t        |      dk(  r#|d   t        j                  |dd        f| _        nt        j                  |      f| _        d| _        | j                  S Nrf   r   r   T)r*   ri   rW   rY   r+   r   r,   s     r   r-   zFlatten.build  s^    &{q !,QQR1I JD!#!5 7D
   r   r!   r"   rr   c                 b    |j                   | _        |j                  |j                   d   d      S )Nr   rR   )ry   _input_shaper   r    s      r   r#   zFlatten.forward  s'    GGyyR((r   r'   c                 8    |j                  | j                        S r   )r   r  r&   s     r   r(   zFlatten.backward  s    ""4#4#455r   r9   )r:   r;   r<   r=   r   r
   rK   r-   rW   r   r   r#   r(   r   r   s   @r   r  r    s\    4!
!sCx !) )t )

 )6BJJ 62:: 6r   r  c                        e Zd ZdZd
def fdZddej                  dedej                  fdZ	dej                  dej                  fd	Z
 xZS )	Dropout2Dz2D Dropout layer.ratec                 >    t         |           || _        d | _        y r   )rI   r   r  mask)r   r  rO   s     r   r   zDropout2D.__init__  s    		r   r!   r"   rr   c                     |ryt         j                  j                  dd| j                  z
  |j                        j                  t         j                        | _        || j                  z  d| j                  z
  z  S |S r	  )rW   rZ   binomialr  ry   r\   r]   r  r    s      r   r#   zDropout2D.forward  sZ    		**1a$))mQWWELLRZZXDItyy=A		M22r   r'   c                 @    || j                   z  d| j                  z
  z  S r	  )r  r  r&   s     r   r(   zDropout2D.backward  s    TYY&!dii-88r   )r   r9   )r:   r;   r<   r=   floatr   rW   r   r   r#   r(   r   r   s   @r   r  r    sO    U 
 t 

 9BJJ 92:: 9r   r  c                        e Zd ZdZddedef fdZdeedf   fdZdde	j                  d	ed
e	j                  fdZde	j                  d
e	j                  fdZ xZS )BatchNorm2Dz"Batch normalization for 2D inputs.momentumepsilonc                 v    t         |           || _        || _        d | _        d | _        d | _        d | _        y r   )rI   r   r  r  gammabetarunning_meanrunning_var)r   r  r  rO   s      r   r   zBatchNorm2D.__init__  s:     
	 r   r*   .c                 n   |d   }t        j                  |t         j                        | _        t        j                  |t         j                        | _        t        j                  |t         j                        | _        t        j                  |t         j                        | _        d| _        |S )NrT   rU   T)	rW   onesr]   r!  r^   r"  r#  r$  r   )r   r*   r   s      r   r-   zBatchNorm2D.build  sr    r?WWXRZZ8
HHXRZZ8	HHXRZZ@7782::>
r   r!   r"   rr   c                    |rt        j                  |dd      }t        j                  |dd      }| j                  | j                  j                  ddd      z  d| j                  z
  |j                         z  z   | _        | j                  | j                  j                  ddd      z  d| j                  z
  |j                         z  z   | _        n<| j                  j                  dddd      }| j                  j                  dddd      }|| _        || _	        || _
        ||z
  t        j                  || j                  z         z  }| j                  j                  dddd      |z  | j                  j                  dddd      z   S )Nr   Tr   rR   r   )rW   r  varr  r#  r   squeezer$  _mean_varr   rX   r  r!  r"  )r   r!   r"   r  r(  
normalizeds         r   r#   zBatchNorm2D.forward  s]   7719t<D&&T:C $0A0A0I0I"aQR0S S !DMM 1T\\^C!DD#}}t/?/?/G/GAq/QQ 4==0CKKMA BD $$,,QAq9D""**1b!Q7C
	$h"''#*<"==
zz!!!RA.;dii>O>OPQSUWXZ[>\\\r   r'   c                    | j                   j                  dddd      }| j                  | j                  z
  t	        j
                  | j                  | j                  z         z  }||z  }t	        j                  |j                  dd        }t	        j                  |dd      |z  }t	        j                  ||z  dd      dz  t	        j                  | j                  | j                  z   d      z  }||z
  ||z  dz  |z  z
  }|S )	Nr   rR   re   r   Tr   r   g      )r!  r   r   r*  rW   rX   r+  r  rY   ry   r   power)	r   r'   r!  r,  grad_normalizedN	grad_meangrad_varr   s	            r   r(   zBatchNorm2D.backward  s    

""1b!Q/kkDJJ."''$))dll:R2SS
%-GGK%%ab)*FF?TJQN	66/J6YQUVY]]88DII4d;< %y0:3H13Lq3PP
r   )g?gh㈵>r9   )r:   r;   r<   r=   r  r   r
   rK   r-   rW   r   r   r#   r(   r   r   s   @r   r  r    sn    ,   u  sCx ] ]t ]

 ](BJJ 2:: r   r  c                   v     e Zd ZdZ fdZdeedf   fdZd
dej                  de
dej                  fd	Z xZS )GlobalAvgPool2Dz>Global average pooling - reduces each channel to single value.c                 "    t         |           y r   )rI   r   r  s    r   r   zGlobalAvgPool2D.__init__  s    r   r*   .c                     || _         t        |      dk(  r|d   |d   f| _        n|d   f| _        d| _        | j                  S r  )r*   ri   r+   r   r,   s     r   r-   zGlobalAvgPool2D.build  sN    &{q !,QQ @D!,Q 1D
   r   r!   r"   rr   c                 0    t        j                  |d      S )Nr   r   )rW   r  r    s      r   r#   zGlobalAvgPool2D.forward  s    wwqv&&r   r9   )r:   r;   r<   r=   r   r
   rK   r-   rW   r   r   r#   r   r   s   @r   r4  r4    sA    H!sCx !' 't '

 'r   r4  c                   p     e Zd ZdZdeedf   f fdZd	dej                  de	dej                  fdZ
 xZS )
ReshapezReshape layer.target_shape.c                 0    t         |           || _        y r   )rI   r   r:  )r   r:  rO   s     r   r   zReshape.__init__  s    (r   r!   r"   rr   c                 X     |j                   g |j                  d d | j                   S r	  )r   ry   r:  r    s      r   r#   zReshape.forward
  s+    qyy:!''"1+:(9(9::r   r9   )r:   r;   r<   r=   r
   rK   r   rW   r   r   r#   r   r   s   @r   r9  r9    s<    )U38_ ); ;t ;

 ;r   r9  c                       e Zd ZdZd	dej
                  dedej
                  fdZdej
                  dej
                  fdZy)
ReLUzReLU activation layer.r!   r"   rr   c                 B    |dkD  | _         t        j                  d|      S )Nr   )_maskrW   r   r    s      r   r#   zReLU.forward  s    U
zz!Qr   r'   c                      || j                   z  S r   )r@  r&   s     r   r(   zReLU.backward  s    TZZ''r   Nr9   	r:   r;   r<   r=   rW   r   r   r#   r(   r/   r   r   r>  r>    sB        t  

  (BJJ (2:: (r   r>  c                       e Zd ZdZd	dej
                  dedej
                  fdZdej
                  dej
                  fdZy)
SigmoidzSigmoid activation layer.r!   r"   rr   c           	          ddt        j                  t        j                  |dd             z   z  | _        | j                  S )Nr   r   r   )rW   r   r   _outputr    s      r   r#   zSigmoid.forward   s6    A4(='= >>?||r   r'   c                 @    || j                   z  d| j                   z
  z  S r	  rF  r&   s     r   r(   zSigmoid.backward$  s    T\\)Q-=>>r   Nr9   rB  r/   r   r   rD  rD    sB    # t 

 ?BJJ ?2:: ?r   rD  c                       e Zd ZdZd	dej
                  dedej
                  fdZdej
                  dej
                  fdZy)
TanhzTanh activation layer.r!   r"   rr   c                 N    t        j                  |      | _        | j                  S r   )rW   r   rF  r    s      r   r#   zTanh.forward+  s    wwqz||r   r'   c                 ,    |d| j                   dz  z
  z  S )Nr   re   rH  r&   s     r   r(   zTanh.backward/  s    a$,,!"3344r   Nr9   rB  r/   r   r   rJ  rJ  (  sB      t 

 5BJJ 52:: 5r   rJ  c                       e Zd ZdZd	dej
                  dedej
                  fdZdej
                  dej
                  fdZy)
SoftmaxzSoftmax activation layer.r!   r"   rr   c                     t        j                  |t        j                  |dd      z
        }|t        j                  |dd      z  | _        | j                  S )NrR   Tr   )rW   r   rz   r   rF  )r   r!   r"   r   s       r   r#   zSoftmax.forward6  sE    q266!"t<<=rvve"tDD||r   r'   c                     |S r   r/   r&   s     r   r(   zSoftmax.backward;  s    r   Nr9   rB  r/   r   r   rN  rN  3  sB    # t 

 
BJJ 2:: r   rN  )%r=   numpyrW   sklearn.baser   r   r   sklearn.utils.validationr   r   r   sklearn.utils.multiclassr	   typingr
   r   r   r   	functoolsr   r   utilsr   r   r   r?   r   r  r  r  r  r4  r9  r>  rD  rJ  rN  r/   r   r   <module>rX     s     G G L L 2 / /   ! D|U |FG GT@ @N6e 629 9$6% 6r'e '&;e ;(5 (?e ?55 5
e 
r   