Ë
    –\;j‡  ã                   óX   — d dl Zd dlZddlmZ ddlmZ  G d„ de«      Z G d„ de«      Zy)	é    Né   )ÚBaseObserver)ÚObserverFactoryc                   ó*   ‡ — e Zd ZdZdˆ fd„	Zd„ Zˆ xZS )ÚGroupWiseWeightObserveraÞ  
    It collects channel-wise maximum absolute values of target weights.
    Args:
        bit_length(int, optional): Number of bits to represent an quantized integer in binary.
        dtype(str, optional): The data type of input tensor.
        name (str, optional): This parameter is used by developers to print debugging information. \
            For details, please refer to :ref:`api_guide_Name`. Default is None.
    Examples:
       .. code-block:: python
            from paddle.quantization import QuantConfig
            from paddle.quantization.quanters import AbsMaxChannelWiseWeightObserver
            quanter = AbsMaxChannelWiseWeightObserver()
            q_config = QuantConfig(activation=None, weight=quanter)
    c                 ó&   •— t         ‰| �  |¬«       y )N)Ú
quant_bits)ÚsuperÚ__init__)Úselfr	   Ú
group_sizeÚ	__class__s      €úpG:\00. PROJECTS\API\Inventory\templateJSON\kerjaOCR\Lib\site-packages\paddle/quantization/observers/groupwise.pyr   z GroupWiseWeightObserver.__init__'   s   ø€ Ü‰Ñ JÐÕ/ó    c                 ó   — t         S ©N)ÚGroupWiseWeightObserverLayer©r   s    r   Ú
_get_classz"GroupWiseWeightObserver._get_class*   s   € Ü+Ð+r   ©é   é€   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   Ú__classcell__©r   s   @r   r   r      s   ø„ ñõ0ö,r   r   c                   ób   ‡ — e Zd Zdˆ fd„	Zd„ Zd„ Zdefd„Zdefd„Zd„ Z	d„ Z
d	„ Zd
„ Zd„ Zˆ xZS )r   c                 óv   •— t         ‰| �  «        || _        || _        || _        d | _        d | _        d | _        y r   )r
   r   r	   r   Ú_layerÚ_maxÚ_scaleÚ_zero_point)r   Úlayerr	   r   r   s       €r   r   z%GroupWiseWeightObserverLayer.__init__/   s9   ø€ Ü‰ÑÔØ$ˆŒØ$ˆŒØˆŒØˆŒ	ØˆŒØˆÕr   c                 ó2   — | j                  |«      | _        |S r   )Ú_cal_abs_maxr"   )r   Úinputss     r   Úforwardz$GroupWiseWeightObserverLayer.forward8   s   € Ø×%Ñ% fÓ-ˆŒ	Øˆr   c                 óŽ  — |j                   }| j                  dk(  s| j                  dk(  sJ d«       ‚|j                   d   | j                  z  dk(  sJ d«       ‚t        |j                   «      dk(  sJ d«       ‚|j                  ddg«      j	                  |d   |d   | j                  z  | j                  g«      }t        j                  t        j                  |«      d¬	«      j                  d
«      }t        j                  |t        j                  d«      k(  t        j                  d«      |«      }|j                  ddg«      }|S )zeUse group_size to group the input, then use the
        absmax method to calculate the scale
        é@   r   z!group_size only support 64 or 128r   z-group_size must be a factor of input channelsr   z Currently only support 2D tensoré   )ÚaxisÚfloat32g:Œ0âŽyE>)Úshaper   ÚlenÚ	transposeÚreshapeÚpaddleÚmaxÚabsÚcastÚwhereÚnpr.   )r   r(   Úinput_shapeÚinput_processedÚabs_max_valuess        r   r'   z)GroupWiseWeightObserverLayer._cal_abs_max<   s,  € ð —l‘lˆà�O‰O˜rÒ! T§_¡_¸Ò%;ð	/à.ó	/Ø;ð �L‰L˜‰O˜dŸo™oÑ-°Ò2ð	;à:ó	;Ø2ä�6—<‘<Ó  AÒ%ÐIÐ'IÓIÐ%Ø ×*Ñ*¨A¨q¨6Ó2×:Ñ:Ø˜‰^˜[¨™^¨t¯©Ñ>ÀÇÁÐPó
ˆô  Ÿ™¤F§J¡J¨Ó$?ÀaÔH×MÑMØó
ˆô  Ÿ™ØœbŸj™j¨›mÑ+¬R¯Z©Z¸Ó-=¸~ó
ˆð (×1Ñ1°1°a°&Ó9ˆØÐr   Úreturnc                  ó   — y)Ng        © r   s    r   Ú	min_valuez&GroupWiseWeightObserverLayer.min_valueU   s   € Ør   c                 ó   — | j                   S r   )r"   r   s    r   Ú	max_valuez&GroupWiseWeightObserverLayer.max_valueX   s   € Ø�y‰yÐr   c                 ó   — | j                   S r   )Ú_quant_bitsr   s    r   Ú
bit_lengthz'GroupWiseWeightObserverLayer.bit_length[   s   € Ø×ÑÐr   c                  ó   — y)Néÿÿÿÿr>   r   s    r   Ú
quant_axisz'GroupWiseWeightObserverLayer.quant_axis^   s   € Ør   c                 ó†   — | j                   €| j                  | _         t        j                  | j                   «      | _        y)z$Compute thresholds for MAX function.N)r#   r"   r3   Ú
zeros_liker$   r   s    r   Úcal_thresholdsz+GroupWiseWeightObserverLayer.cal_thresholdsa   s.   € à�;‰;ÐØŸ)™)ˆDŒKÜ!×,Ñ,¨T¯[©[Ó9ˆÕr   c                 óR   — | j                   €| j                  «        | j                   S )zReturn output scales.)r#   rJ   r   s    r   Úscalesz#GroupWiseWeightObserverLayer.scalesg   s"   € à�;‰;ÐØ×ÑÔ!Ø�{‰{Ðr   c                 óR   — | j                   €| j                  «        | j                   S )zReturn output zero points.)r$   rJ   r   s    r   Úzero_pointsz(GroupWiseWeightObserverLayer.zero_pointsm   s&   € à×ÑÐ#Ø×ÑÔ!Ø×ÑÐr   r   )r   r   r   r   r)   r'   Úfloatr?   rA   rD   rG   rJ   rL   rN   r   r   s   @r   r   r   .   sC   ø„ õ òòð2˜5ó ð˜5ó ò òò:òö r   r   )	Únumpyr8   r3   Úbase_observerr   Úfactoryr   r   r   r>   r   r   Ú<module>rS      s-   ðó ã å (Ý %ô,˜oô ,ô.C  <õ C r   