
    i                     ^   d Z ddlZddlmZmZ ddej                  deeef   deeef   dej                  fdZ	 dd	ej                  d
eedf   deeef   deeef   dej                  f
dZ	dde
dede
fdZdedefdZde
deedf   defdZ G d d      Zddej                  dedefdZy)zQ
ConvMLP Utilities
=================
Vedic-inspired optimized utility functions.
    N)TupleUnionxkernel_sizestridesreturnc           
         | j                   \  }}}}|\  }}|\  }	}
||z
  |	z  dz   }||z
  |
z  dz   }t        j                  ||||||f| j                        }t	        |      D ]D  }t	        |      D ]4  }||	z  }||
z  }| dddd|||z   |||z   f   |dddddddd||f<   6 F |j                  dddddd      j                  |||z  ||z  |z        }|j                  ddd      S )	a  
    Image to Column conversion (Vedic Block Multiplication).
    
    Inspired by Urdhva-Tiryagbyham (vertically and crosswise).
    Converts 2D convolution to efficient matrix multiplication.
    
    Parameters
    ----------
    x : np.ndarray
        Input tensor of shape (batch, channels, height, width)
    kernel_size : tuple
        (kernel_height, kernel_width)
    strides : tuple
        (stride_h, stride_w)
    
    Returns
    -------
    col : np.ndarray
        Unfolded patches of shape (batch, channels * kh * kw, out_h * out_w)
    
    Example
    -------
    >>> patches = im2col(x, kernel_size=(3, 3), strides=(1, 1))
    >>> # patches[:, :, 0] = first patch flattened
       dtypeNr               )shapenpzerosr   range	transposereshape)r   r   r   batchchannelshwkhkwshswout_hout_wcolijh_startw_starts                     ;/home/per/Documents/sklearn ConvMLP module/convmlp/utils.pyim2colr'      s.   4 GGE8QFBFBVNQEVNQE ((E8RUE:!''
JC5\ Tu 	TA"fG"fG$%aGGBJ,>PR
@R&R$SC1aAq !	TT --1aAq
)
1
1%SUXZHZ
[C==Aq!!    r!   x_shape.c           	         |\  }}}}|\  }}	|\  }
}||z
  |
z  dz   }||	z
  |z  dz   }| j                  ddd      j                  ||||||	      } t        j                  || j                        }t        |      D ]C  }t        |      D ]3  }||
z  }||z  }|dddd|||z   |||	z   fxx   | dd||f   z  cc<   5 E |S )a  
    Column to Image conversion (Reverse of im2col).
    
    Parameters
    ----------
    col : np.ndarray
        Column data of shape (batch, channels * kh * kw, out_h * out_w)
    x_shape : tuple
        Original input shape (batch, channels, height, width)
    kernel_size : tuple
        (kernel_height, kernel_width)
    strides : tuple
        (stride_h, stride_w)
    
    Returns
    -------
    x : np.ndarray
        Reconstructed image of shape x_shape
    r
   r   r   r   N)r   r   r   r   r   r   )r!   r)   r   r   r   r   r   r   r   r   r   r   r   r    r   r"   r#   r$   r%   s                      r&   col2imr+   ;   s   * $E8QFBFBVNQEVNQE --1a
 
(
(uhB
OC 			*A 5\ Lu 	LA"fG"fGaGGBJ&
(::;s1a7|K;	LL Hr(   layersstart_rfc                 x   |g}|}| D ])  }t        |d      rt        |d      r|j                  \  }}t        |j                  t              r|j                  n|j                  |j                  f\  }}||dz
  t        j                  |D 	cg c]   }	t        |	t              r|	j                  n|	" c}	      z  z   }|j                  |       t        |d      st        |j                  t              r|j                  n|j                  |j                  f\  }
}||
z  }|j                  |       , |S c c}	w )a  
    Calculate receptive field at each layer.
    
    Based on the "coconut sellers" concept - each layer adds to the receptive field.
    
    Parameters
    ----------
    layers : list
        List of layer objects with kernel_size and strides attributes
    start_rf : int
        Initial receptive field size
    
    Returns
    -------
    receptive_fields : list
        Receptive field size at each layer
    r   r   r
   	pool_size)	hasattrr   
isinstancer   tupler   prodappendr/   )r,   r-   receptive_fields
current_rflayerr   r   r   r   lphpws               r&   calculate_receptive_fieldr;   g   s,   $ !zJ 05-(WUI-F&&FB&0&FU]]U]]\a\i\iLjFB#rAvCS:U>? HRRSUZG[!))ab:b :U 2V 'V VJ##J/UK((25??E(JU__QVQ`Q`bgbqbqPrFB#bJ##J/0 :Us   %D7nc                 j    dt        t        j                  t        j                  |                   z  S )z Return the next power of 2 >= n.r   )intr   ceillog2)r<   s    r&   next_power_of_2rA      s#    BGGBGGAJ'(((r(   conv_layersinput_shapec                    d}|}| D ]  }t        |t              r[|j                  \  }}|j                  |      dd \  }}||z  |z  |z  |d   z  |j                  z  }	||	z  }|j                  |      }nt        |t
        t        f      st        |j                  t              r|j                  n|j                  |j                  f\  }
}|d   |d   |d   |
z  |d   |z  f} |S )z
    Calculate total FLOPs for a convolutional network.
    
    Inspired by Bhaskara's wheel method - counting all operations.
    r   Nr
   r   r   )	r1   Conv2Dr   _compute_output_shapefilters	MaxPool2D	AvgPool2Dr/   r2   )rB   rC   total_flopscurrent_shaper7   r   r   r   r    flopsr9   r:   s               r&   calculate_flopsrO      s    KM KeV$&&FB 66}EbcJLE5 EMB&+mB.??%--OE5 K!77FM	956(25??E(JU__QVQ`Q`bgbqbqPrFB*1-}Q/?(+r1=3Cr3IKMK r(   c                       e Zd ZdZd
defdZdej                  fdZdej                  dej                  fdZ	dej                  dej                  fdZ
y	)NikhilamQuantizerz
    Nikhilam Quantization for fast inference.
    
    Based on the Vedic "deficiency from base" method.
    Quantizes weights to low-precision for faster computation.
    num_bitsc                 .    || _         d | _        d | _        y N)rR   scale
zero_point)selfrR   s     r&   __init__zNikhilamQuantizer.__init__   s     
r(   weightsc                     |j                         |j                         z
  d| j                  z  dz
  z  | _        |j                         | _        | S )z"Calculate quantization parameters.r   r
   )maxminrR   rU   rV   rW   rY   s     r&   fitzNikhilamQuantizer.fit   s?    kkmgkkm3T]]8JQ8NO
!++-r(   r   c                     | j                   | j                  |       t        j                  || j                  z
  | j                   z        j                  t        j                        S )z"Quantize weights to low precision.)rU   r^   r   roundrV   astypeint8r]   s     r&   quantizezNikhilamQuantizer.quantize   sH    ::HHWxx4??2djj@AHHQQr(   c                 t    |j                  t        j                        | j                  z  | j                  z   S )z#Restore weights from low precision.)ra   r   float32rU   rV   r]   s     r&   
dequantizezNikhilamQuantizer.dequantize   s'    ~~bjj)DJJ6HHr(   N)   )__name__
__module____qualname____doc__r>   rX   r   ndarrayr^   rc   rf    r(   r&   rQ   rQ      s_     
2:: R

 Rrzz RI"** I Ir(   rQ   expected_channelsc                 ,   t        | j                        dk(  ryt        | j                        dk(  ryt        | j                        dk(  r3|r0| j                  d   |k7  rt        d| d| j                  d          yt        d	| j                         )
zValidate input tensor shape.r   Fr   r   r
   z	Expected z channels, got TzInvalid input shape: )lenr   
ValueError)r   rn   s     r&   validate_input_shaperr      s    
177|q	QWW		QWW	/@!@y):(;?177ST:,WXX0	:;;r(   ))r
   r
   )r
   rT   )rk   numpyr   typingr   r   rl   r>   r'   r+   listr;   rA   rO   rQ   boolrr   rm   r(   r&   <module>rw      s    -"bjj -"uS#X -"sCx -"^`^h^h -"b '-)

 )U38_ )5c? )#s(O)13)X"d "c "$ "J)s )s )
 E#s(O  6I I<<BJJ <3 <$ <r(   