Ë
    T^(h|u  ã                   óØ  — d dl Z d dlZd dlZd dlmZ d dlmZ ddlmZ  ej                  e
«      Z G d„ dej                  «      Z G d„ d	ej                  «      Z G d
„ dej                  «      Z G d„ dej                  «      Z G d„ dej                  «      Z G d„ dej                  «      Zdd„Zdd„Zdd„Z G d„ de«      Z G d„ de«      Z G d„ de«      Zdd„Z G d„ de«      Zy) é    N)Únn)ÚFunctioné   )Úloggingc                   ó>   ‡ — e Zd ZdZ	 	 	 	 	 	 	 	 	 dˆ fd„	Zdd„Zˆ xZS )ÚQuantEmbeddingaÞ  
    Quantized version of `torch.nn.Embedding`. Adds quantization-specific arguments on top of `torch.nn.Embedding`.

    Args:
        weight_bit (`int`, *optional*, defaults to `8`):
            Bitwidth for the quantized weight.
        momentum (`float`, *optional*, defaults to `0.95`):
            Momentum for updating the activation quantization range.
        quant_mode (`bool`, *optional*, defaults to `False`):
            Whether or not the layer is quantized.
    c                 óì  •— t         ‰| �  «        || _        || _        || _        || _        || _        || _        || _        t        j                  t        j                  ||g«      «      | _        | j                  dt        j                  d«      «       | j                  dt        j                  | j                  «      «       |	| _        |
| _        || _        d| _        t(        j*                  | _        y )NÚweight_scaling_factoré   Úweight_integerF)ÚsuperÚ__init__Únum_ÚdimÚpadding_idxÚmax_normÚ	norm_typeÚscale_grad_by_freqÚsparser   Ú	ParameterÚtorchÚzerosÚweightÚregister_bufferÚ
zeros_likeÚ
weight_bitÚmomentumÚ
quant_modeÚpercentile_modeÚSymmetricQuantFunctionÚapplyÚweight_function)ÚselfÚnum_embeddingsÚembedding_dimr   r   r   r   r   Ú_weightr   r   r   Ú	__class__s               €úe/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/ibert/quant_modules.pyr   zQuantEmbedding.__init__,   sÅ   ø€ ô 	‰ÑÔØ"ˆŒ	Ø ˆŒØ&ˆÔØ ˆŒØ"ˆŒØ"4ˆÔØˆŒä—l‘l¤5§;¡;°ÀÐ/NÓ#OÓPˆŒØ×ÑÐ4´e·k±kÀ!³nÔEØ×ÑÐ-¬u×/?Ñ/?ÀÇÁÓ/LÔMà$ˆŒØ ˆŒØ$ˆŒØ$ˆÔÜ5×;Ñ;ˆÕó    c           	      ó\  — | j                   sct        j                  j                  || j                  | j
                  | j                  | j                  | j                  | j                  «      d fS | j                  }|j                  j                  «       }|j                  «       j                  d«      }|j                  «       j                  d«      }t        | j                   ||d«      | _        | j%                  | j                  | j                   | j&                  | j"                  «      | _        t        j                  j                  || j(                  | j
                  | j                  | j                  | j                  | j                  «      }|| j"                  z  | j"                  fS )Nr   F)r   r   Ú
functionalÚ	embeddingr   r   r   r   r   r   ÚdataÚdetachÚminÚexpandÚmaxÚ$symmetric_linear_quantization_paramsr   r
   r"   r   r   )	r#   ÚxÚ	positionsÚincremental_stateÚwÚw_transformÚw_minÚw_maxÚemb_ints	            r(   ÚforwardzQuantEmbedding.forwardM   sT  € Ø�Šä—‘×'Ñ'ØØ—K‘KØ×$Ñ$Ø—M‘MØ—N‘NØ×+Ñ+Ø—K‘Kóð ðð ð �K‰KˆØ—f‘f—m‘m“oˆØ—‘Ó!×(Ñ(¨Ó+ˆØ—‘Ó!×(Ñ(¨Ó+ˆä%IÈ$Ï/É/Ð[`ÐbgÐinÓ%oˆÔ"Ø"×2Ñ2Ø�K‰K˜Ÿ™¨$×*>Ñ*>À×@ZÑ@Zó
ˆÔô —-‘-×)Ñ)ØØ×ÑØ×ÑØ�M‰MØ�N‰NØ×#Ñ#Ø�K‰Kó
ˆð ˜×3Ñ3Ñ3°T×5OÑ5OÐOÐOr)   )	NNç       @FFNé   çffffffî?F©NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r;   Ú__classcell__©r'   s   @r(   r   r      s1   ø„ ñ
ð  ØØØ ØØØØØõ<÷B"Pr)   r   c                   ó<   ‡ — e Zd ZdZdˆ fd„	Zd„ Z	 	 	 	 	 dd„Zˆ xZS )ÚQuantActap  
    Quantizes the given activation.

    Args:
        activation_bit (`int`):
            Bitwidth for the quantized activation.
        act_range_momentum (`float`, *optional*, defaults to `0.95`):
            Momentum for updating the activation quantization range.
        per_channel (`bool`, *optional*, defaults to `False`):
            Whether to or not use channel-wise quantization.
        channel_len (`int`, *optional*):
            Specify the channel length when set the *per_channel* True.
        quant_mode (`bool`, *optional*, defaults to `False`):
            Whether or not the layer is quantized.
    c                 óò  •— t         ‰| �  «        || _        || _        || _        || _        d| _        t        j                  | _	        | j
                  sš| j                  dt        j                  d«      «       | j                  dt        j                  d«      «       | j                  dt        j                  d«      «       | xj                  dz  c_        | xj                  dz  c_        y t        d«      ‚)NFÚx_minr   Úx_maxÚact_scaling_factorgñhãˆµøä>ú;per-channel mode is not currently supported for activation.)r   r   Úactivation_bitÚact_range_momentumr   Úper_channelÚ
percentiler    r!   Úact_functionr   r   r   rI   rJ   ÚNotImplementedError)r#   rM   rN   rO   Úchannel_lenr   r'   s         €r(   r   zQuantAct.__init__ƒ   s¼   ø€ Ü‰ÑÔà,ˆÔØ"4ˆÔØ$ˆŒØ&ˆÔØˆŒÜ2×8Ñ8ˆÔà×ÒØ× Ñ  ¬%¯+©+°a«.Ô9Ø× Ñ  ¬%¯+©+°a«.Ô9Ø× Ñ Ð!5´u·{±{À1³~ÔFØ�JŠJ˜$Ñ�JØ�JŠJ˜$ÑŽJä%Ð&cÓdÐdr)   c           
      óØ   — | j                   j                  › d| j                  › d| j                  › d| j                  j                  «       d›d| j                  j                  «       d›d�
S )Nz(activation_bit=z, quant_mode: z, Act_min: z.2fz, Act_max: ú))r'   r@   rM   r   rI   ÚitemrJ   )r#   s    r(   Ú__repr__zQuantAct.__repr__–   si   € à�~‰~×&Ñ&Ð'Ð'7¸×8KÑ8KÐ7Lð MØŸ?™?Ð+¨;°t·z±z·±Ó7HÈÐ6Mð NØŸ
™
Ÿ™Ó)¨#Ð.¨að1ð	
r)   c                 ó€  — |€|n||z   }| j                   �rÂ| j                  rJ d«       ‚| j                  rJ d«       ‚|j                  j	                  «       }|j                  j                  «       }	|	j                  «       j                  «       dk(  r!|j                  «       j                  «       dk(  sJ d«       ‚| j                  j	                  «       dkD  rF| j                  j                  «       dk  r)| j                  |z   | _        | j                  |	z   | _	        n¼| j                  dk(  rKt        j                  | j                  |«      | _        t        j
                  | j                  |	«      | _	        nb| j                  | j                  z  |d| j                  z
  z  z   | _        | j                  | j                  z  |	d| j                  z
  z  z   | _	        | j                  s|d fS |€| j                  n|}|€| j                  n|}	t        | j                  ||	| j                  ¬	«      | _        |€3| j!                  || j                  | j                  | j                  «      }
n.t"        j%                  ||| j                  | j                  ||«      }
| j                  j'                  d«      }|
|z  | j                  fS )
Nz:percentile mode is not currently supported for activation.rL   r   z5NaN detected when computing min/max of the activationg¢&ú|”ç¾g¢&ú|”ç>éÿÿÿÿr   )rO   )ÚtrainingrP   rO   r-   r/   r1   ÚisnanÚsumrI   rJ   rN   r   r   r2   rM   rK   rQ   ÚFixedPointMulr!   Úview)r#   r3   Úpre_act_scaling_factorÚidentityÚidentity_scaling_factorÚspecified_minÚspecified_maxÚx_actrI   rJ   Úquant_act_intÚcorrect_output_scales               r(   r;   zQuantAct.forward�   s[  € ð Ð%‘¨8°a©<ˆà�=‹=Ø—’ÐdÐ(dÓdÐ&Ø×'Ò'ÐfÐ)fÓfÐ'Ø—J‘J—N‘NÓ$ˆEØ—J‘J—N‘NÓ$ˆEà—;‘;“=×$Ñ$Ó&¨!Ò+°·±³×0AÑ0AÓ0CÀqÒ0Hð ØGóÐHð
 �z‰z�~‰~Ó 'Ò)¨d¯j©j¯n©nÓ.>ÀÒ.GØ!ŸZ™Z¨%Ñ/�”
Ø!ŸZ™Z¨%Ñ/�•
ð ×(Ñ(¨BÒ.Ü"ŸY™Y t§z¡z°5Ó9�”
Ü"ŸY™Y t§z¡z°5Ó9�•
à!ŸZ™Z¨$×*AÑ*AÑAÀEÈQÐQU×QhÑQhÑMhÑDiÑi�”
Ø!ŸZ™Z¨$×*AÑ*AÑAÀEÈQÐQU×QhÑQhÑMhÑDiÑi�”
à�ŠØ˜$�;Ðà+Ð3�—
’
¸ˆØ+Ð3�—
’
¸ˆä"FØ×Ñ ¨¸4×;KÑ;Kô#
ˆÔð "Ð)à ×-Ñ-¨a°×1DÑ1DÀdÇoÁoÐW[×WnÑWnÓo‰Mä)×/Ñ/ØØ&Ø×#Ñ#Ø×'Ñ'ØØ'óˆMð  $×6Ñ6×;Ñ;¸BÓ?ÐàÐ3Ñ3°T×5LÑ5LÐLÐLr)   )r>   FNF)NNNNN©r@   rA   rB   rC   r   rW   r;   rD   rE   s   @r(   rG   rG   r   s*   ø„ ñõ eò&
ð  $ØØ $ØØ÷<Mr)   rG   c                   ó8   ‡ — e Zd ZdZ	 dˆ fd„	Zˆ fd„Zdd„Zˆ xZS )ÚQuantLineara8  
    Quantized version of `torch.nn.Linear`. Adds quantization-specific arguments on top of `torch.nn.Linear`.

    Args:
        weight_bit (`int`, *optional*, defaults to `8`):
            Bitwidth for the quantized weight.
        bias_bit (`int`, *optional*, defaults to `32`):
            Bitwidth for the quantized bias.
        per_channel (`bool`, *optional*, defaults to `False`):
            Whether or not to use channel-wise quantization.
        quant_mode (`bool`, *optional*, defaults to `False`):
            Whether or not the layer is quantized.
    c                 ó’  •— t         ‰| �  «        || _        || _        t	        j
                  t        j                  ||g«      «      | _        | j                  dt        j                  | j                  «      «       | j                  dt        j                  | j                  «      «       |r\t	        j
                  t        j                  |«      «      | _        | j                  dt        j                  | j                  «      «       || _        || _        || _        || _        || _        d| _        t"        j$                  | _        y )Nr   Úfc_scaling_factorÚbias_integerF)r   r   Úin_featuresÚout_featuresr   r   r   r   r   r   r   Úbiasr   r   rO   Úbias_bitr   r    r!   r"   )	r#   rm   rn   ro   r   rp   rO   r   r'   s	           €r(   r   zQuantLinear.__init__ë   só   ø€ ô 	‰ÑÔØ&ˆÔØ(ˆÔä—l‘l¤5§;¡;°¸kÐ/JÓ#KÓLˆŒØ×ÑÐ-¬u×/?Ñ/?ÀÇÁÓ/LÔMØ×ÑÐ0´%·+±+¸d×>OÑ>OÓ2PÔQÙÜŸ™¤U§[¡[°Ó%>Ó?ˆDŒIØ× Ñ  ´×1AÑ1AÀ$Ç)Á)Ó1LÔMà$ˆŒØ$ˆŒØ&ˆÔØ ˆŒØ$ˆŒØ$ˆÔÜ5×;Ñ;ˆÕr)   c                 ód   •— t         ‰| �  «       }d|› d| j                  › d| j                  › d�}|S )Nú(z weight_bit=z, quant_mode=rU   )r   rW   r   r   )r#   Úsr'   s     €r(   rW   zQuantLinear.__repr__  s9   ø€ Ü‰GÑÓˆØ�ˆs�,˜tŸ™Ð/¨}¸T¿_¹_Ð<MÈQÐOˆØˆr)   c                 ó  — | j                   s8t        j                  j                  || j                  | j
                  ¬«      d fS |�|j                  dk(  sJ d«       ‚| j                  }|j                  j                  «       }| j                  r7t        j                  |dd ¬«      \  }}t        j                  |dd ¬«      \  }}n>|j                  «       j                  d«      }|j                  «       j                  d«      }t        | j                  ||| j                  «      | _        | j#                  | j                  | j                  | j$                  | j                   «      | _        | j                   |z  }| j
                  �-| j#                  | j
                  | j(                  d|«      | _        |j-                  dd«      }||z  }	t        j                  j                  |	| j&                  | j*                  ¬«      |z  |fS )N)r   ro   )r   z«Input activation to the QuantLinear layer should be globally (non-channel-wise) quantized. Please add a QuantAct layer with `per_channel = True` before this QuantAct layerr   )r   ÚoutFrY   )r   r   r+   Úlinearr   ro   Úshaper-   r.   rO   r   r/   r1   r0   r2   r   rk   r"   r   r   rp   rl   r^   )
r#   r3   Úprev_act_scaling_factorr6   r7   r8   Ú_r9   Úbias_scaling_factorÚx_ints
             r(   r;   zQuantLinear.forward  sÂ  € Ø�ŠÜ—=‘=×'Ñ'¨°$·+±+ÀDÇIÁIÐ'ÓNÐPTÐTÐTð 'Ð2Ð7N×7TÑ7TÐX\Ò7\ð 	
ð_ó	
Ð\ð
 �K‰KˆØ—f‘f—m‘m“oˆØ×ÒÜ—y‘y °!¸Ô>‰HˆE�1Ü—y‘y °!¸Ô>‰HˆE‘1à—O‘OÓ%×,Ñ,¨QÓ/ˆEØ—O‘OÓ%×,Ñ,¨QÓ/ˆEä!EÀdÇoÁoÐW\Ð^cÐei×euÑeuÓ!vˆÔØ"×2Ñ2Ø�K‰K˜Ÿ™¨$×*>Ñ*>À×@VÑ@Vó
ˆÔð #×4Ñ4Ð7NÑNÐà�9‰9Ð Ø $× 4Ñ 4°T·Y±YÀÇÁÈuÐViÓ jˆDÔà"9×">Ñ">¸qÀ"Ó"EÐØÐ+Ñ+ˆô �M‰M× Ñ  ¨t×/BÑ/BÈ×IZÑIZÐ Ó[Ð^qÑqØð
ð 	
r)   )Tr=   é    FF©Nrg   rE   s   @r(   ri   ri   Ü   s   ø„ ñð nsõ<ô,÷
#
r)   ri   c                   ó2   ‡ — e Zd ZdZdˆ fd„	Zd„ Zdd„Zˆ xZS )ÚIntGELUa}  
    Quantized version of `torch.nn.GELU`. Adds quantization-specific arguments on top of `torch.nn.GELU`.

    Args:
        quant_mode (`bool`, *optional*, defaults to `False`):
            Whether or not the layer is quantized.
        force_dequant (`str`, *optional*, defaults to `"none"`):
            Force dequantize the layer if either "gelu" or "nonlinear" is given.
    c                 ó0  •— t         ‰| �  «        || _        |dv rt        j	                  d«       d| _        | j                  st        j                  «       | _        d| _        d| _	        g d¢| _
        | j                  dxx   | j                  d   z  cc<   y )	N)Ú	nonlinearÚgeluzForce dequantize geluFgà-� ö?é   )g]mÅþ²{Ò¿gçû©ñÒMü¿r   é   r   )r   r   r   ÚloggerÚinfor   ÚGELUÚactivation_fnÚkÚconstÚcoeff)r#   r   Úforce_dequantr'   s      €r(   r   zIntGELU.__init__7  sv   ø€ Ü‰ÑÔØ$ˆŒàÐ1Ñ1Ü�K‰KÐ/Ô0Ø#ˆDŒOà�ŠÜ!#§¡£ˆDÔàˆŒØˆŒ
Ú)ˆŒ
Ø�
‰
�1‹˜Ÿ™ A™Ñ&Œr)   c                 óÖ  — t        j                  | j                  d   |z  «      }t        j                  | j                  d   |dz  z  «      }t        j                  |«      }t        j                  t        j
                  |«      | «      }|||z   dz  |z   z  }|dz  | j                  d   z  }t        j                  |d| j                  z  z  «      }|d| j                  z  z  }||fS ©Nr   r„   r   )	r   Úfloorr‹   Úsignr/   ÚabsÚ	floor_ster!   rŠ   )r#   r{   Úscaling_factorÚb_intÚc_intr�   Úabs_intÚy_ints           r(   Úint_erfzIntGELU.int_erfG  sÏ   € Ü—‘˜DŸJ™J q™M¨NÑ:Ó;ˆÜ—‘˜DŸJ™J q™M¨N¸AÑ,=Ñ=Ó>ˆÜ�z‰z˜%Ó ˆä—)‘)œEŸI™I eÓ,¨u¨fÓ5ˆØ˜ 5™¨QÑ.°Ñ6Ñ7ˆØ'¨Ñ*¨T¯Z©Z¸©]Ñ:ˆô —‘ ¨¨4¯:©:©Ñ 5Ó6ˆØ'¨!¨T¯Z©Z©-Ñ7ˆà�nÐ$Ð$r)   c                 óÆ   — | j                   s| j                  |«      d fS ||z  }| j                  ||| j                  z  «      \  }}d|z  }|||z   z  }||z  dz  }||z  |fS )Nç      ð?r„   )r   rˆ   r˜   r‰   )r#   r3   r“   r{   Úsigmoid_intÚsigmoid_scaling_factorÚ	shift_ints          r(   r;   zIntGELU.forwardV  s…   € Ø�ŠØ×%Ñ% aÓ(¨$Ð.Ð.à�NÑ"ˆØ.2¯l©l¸5À.ÐSW×SYÑSYÑBYÓ.ZÑ+ˆÐ+àÐ1Ñ1ˆ	à˜ yÑ0Ñ1ˆØ'Ð*@Ñ@À1ÑDˆà�~Ñ% ~Ð5Ð5r)   )TÚnoner}   )r@   rA   rB   rC   r   r˜   r;   rD   rE   s   @r(   r   r   ,  s   ø„ ñõ'ò %÷6r)   r   c                   ó6   ‡ — e Zd ZdZdˆ fd„	Zd„ Zd„ Zd„ Zˆ xZS )Ú
IntSoftmaxaØ  
    Quantized version of `torch.nn.Softmax`. Adds quantization-specific arguments on top of `torch.nn.Softmax`.

    Args:
        output_bit (`int`):
            Bitwidth for the layer output activation.
        quant_mode (`bool`, *optional*, defaults to `False`):
            Whether or not the layer is quantized.
        force_dequant (`str`, *optional*, defaults to `"none"`):
            Force dequantize the layer if either "softmax" or "nonlinear" is given.
    c                 ó‚  •— t         ‰| �  «        || _        d| _        || _        |dv rt
        j                  d«       d| _        t        d| j                  ¬«      | _        d| _	        d| _
        g d	¢| _        | j                  d
xx   | j                  d   z  cc<   | j                  dxx   | j                  d   z  cc<   y )Nr|   )r�   ÚsoftmaxzForce dequantize softmaxFé   ©r   gvqà-æ¿é   )gN„ª$ôëÖ?g¾Ã'|:ï?rš   r   r   r„   )r   r   Ú
output_bitÚmax_bitr   r…   r†   rG   ÚactÚx0rŠ   Úcoef)r#   r¦   r   rŒ   r'   s       €r(   r   zIntSoftmax.__init__r  s›   ø€ Ü‰ÑÔØ$ˆŒØˆŒØ$ˆŒàÐ4Ñ4Ü�K‰KÐ2Ô3Ø#ˆDŒOä˜B¨4¯?©?Ô;ˆŒØˆŒØˆŒ
Ú1ˆŒ	Ø�	‰	�!‹˜Ÿ	™	 !™Ñ$‹Ø�	‰	�!‹˜Ÿ	™	 !™Ñ$Œr)   c                 ó6  — t        j                  «       5  t        j                  | j                  d   |z  «      }t        j                  | j                  d   |dz  z  «      }d d d «       |z   |z  z   }| j                  d   |dz  z  }||fS # 1 sw Y   Œ-xY wrŽ   )r   Úno_gradr�   rª   )r#   r{   r“   r”   r•   Úzs         r(   Úint_polynomialzIntSoftmax.int_polynomialƒ  s–   € Ü�]‰]‹_ñ 	BÜ—K‘K §	¡	¨!¡¨~Ñ =Ó>ˆEÜ—K‘K §	¡	¨!¡¨~¸qÑ/@Ñ @ÓAˆE÷	Bð �U‰]˜eÑ# eÑ+ˆØŸ™ 1™¨¸Ñ(9Ñ9ˆØ�.Ð Ð ÷	Bð 	Bús   •ABÂBc                 óî  — t        j                  «       5  t        j                  | j                  |z  «      }d d d «       t        j                  || j
                  z  «      }t        j                  ||z  «      }|||z  z
  }| j                  ||«      \  }}t        j                  t        j                  |d| j
                  |z
  z  z  «      d¬«      }|d| j
                  z  z  }||fS # 1 sw Y   Œ´xY w)Nr„   r   ©r/   )
r   r¬   r�   r©   r1   rŠ   r’   r!   r®   Úclamp)r#   r{   r“   Úx0_intÚqÚrÚexp_intÚexp_scaling_factors           r(   Úint_expzIntSoftmax.int_exp‹  sÑ   € Ü�]‰]‹_ñ 	;Ü—[‘[ §¡¨>Ñ!9Ó:ˆF÷	;ä—	‘	˜% §¡¨fÑ!4Ó5ˆä�O‰O˜E F™NÓ+ˆØ�F˜Q‘JÑˆØ&*×&9Ñ&9¸!¸^Ó&LÑ#ˆÐ#Ü—+‘+œiŸo™o¨g¸¸d¿j¹jÈ1¹nÑ8MÑ.MÓNÐTUÔVˆØ+¨a°·±©mÑ;ˆØ˜Ð&Ð&÷	;ð 	;ús   •#C+Ã+C4c                 ó
  — | j                   s#t        j                  j                  |d¬«      d fS ||z  }|j	                  dd¬«      \  }}||z
  }| j                  ||«      \  }}| j                  ||«      \  }}||z  }|j                  dd¬«      }	t        j                  d| j                  z  |	z  «      }
t        j                  ||
z  d| j                  | j                  z
  z  z  «      }dd| j                  z  z  }||z  |fS )NrY   ©r   T)r   Úkeepdimr„   r   )r   r   r+   r¢   r1   r·   r¨   r\   r’   r!   r§   r¦   )r#   r3   r“   r{   Ú	x_int_maxry   rµ   r¶   ÚexpÚexp_int_sumÚfactors              r(   r;   zIntSoftmax.forward—  s  € Ø�ŠÜ—=‘=×(Ñ(¨°Ð(Ó3°TÐ9Ð9à�NÑ"ˆà—y‘y R°�yÓ6‰ˆ	�1Ø˜	Ñ!ˆØ&*§l¡l°5¸.Ó&IÑ#ˆÐ#ð #'§(¡(¨7Ð4FÓ"GÑˆÐØÐ*Ñ*ˆà—k‘k b°$�kÓ7ˆÜ—‘  D§L¡L¡°;Ñ!>Ó?ˆÜ—/‘/ '¨FÑ"2°Q¸4¿<¹<È$Ï/É/Ñ;YÑ5ZÑ"ZÓ[ˆØ˜Q §¡Ñ/Ñ/ˆØ˜Ñ'¨Ð7Ð7r)   )Frž   )	r@   rA   rB   rC   r   r®   r·   r;   rD   rE   s   @r(   r    r    e  s   ø„ ñ
õ%ò"!ò
'ö8r)   r    c                   ó8   ‡ — e Zd ZdZdˆ fd„	Zd„ Zd„ Zdd„Zˆ xZS )ÚIntLayerNormaû  
    Quantized version of `torch.nn.LayerNorm`. Adds quantization-specific arguments on top of `torch.nn.LayerNorm`.

    Args:
        output_bit (`int`, *optional*, defaults to `8`):
            Bitwidth for the layer output activation.
        quant_mode (`bool`, *optional*, defaults to `False`):
            Whether or not the layer is quantized.
        force_dequant (`str`, *optional*, defaults to `"none"`):
            Force dequantize the layer if either "layernorm" or "nonlinear" is given.
    c                 ó   •— t         ‰| �  «        || _        || _        t	        j
                  t        j                  |«      «      | _        t	        j
                  t        j                  |«      «      | _	        || _
        |dv rt        j                  d«       d| _
        | j                  dt        j                  d«      «       || _        d| _        d | _        t#        | j                  | j                  ¬«      | _        y )N)r�   Ú	layernormzForce dequantize layernormFÚshiftr   r|   r¤   )r   r   Únormalized_shapeÚepsr   r   r   r   r   ro   r   r…   r†   r   r¦   r§   Údim_sqrtrG   Ú
activation)r#   rÄ   rÅ   r¦   r   rŒ   r'   s         €r(   r   zIntLayerNorm.__init__¹  s¸   ø€ Ü‰ÑÔØ 0ˆÔØˆŒä—l‘l¤5§;¡;Ð/?Ó#@ÓAˆŒÜ—L‘L¤§¡Ð-=Ó!>Ó?ˆŒ	à$ˆŒØÐ6Ñ6Ü�K‰KÐ4Ô5Ø#ˆDŒOà×Ñ˜W¤e§k¡k°!£nÔ5Ø$ˆŒØˆŒØˆŒÜ" 4§?¡?¸t¿¹ÔOˆ�r)   c           	      ó  — t        j                  «       5  |dz  }t        j                  |dd¬«      }t        j                  t        j                  |d| j
                  z  z  «      «      j                  «       j                  «       }| j                  }t        j                  | j                  |«      | _        t        j                  dt        |«      › dt        | j                  «      › �«       d d d «       y # 1 sw Y   y xY w)Nr„   T©Úaxisrº   zDynamic shift adjustment: z -> )r   r¬   r\   Úlog2Úsqrtr§   Úceilr1   rÃ   r…   r†   Úint)r#   r—   Úy_sq_intÚvar_intrÃ   Ú	shift_olds         r(   Ú	set_shiftzIntLayerNorm.set_shiftÌ  s¼   € Ü�]‰]‹_ñ 	\Ø˜a‘xˆHÜ—i‘i ¨q¸$Ô?ˆGÜ—Z‘Z¤§
¡
¨7°Q¸¿¹±_Ñ+DÓ EÓF×KÑKÓM×RÑRÓTˆEØŸ
™
ˆIÜŸ™ 4§:¡:¨uÓ5ˆDŒJÜ�K‰KÐ4´S¸³^Ð4DÀDÌÈTÏZÉZËÐHYÐZÔ[÷	\÷ 	\ñ 	\ús   •CC8Ã8Dc                 ó¬   — | j                  |«       t        j                  |d| j                  z  z  «      }|dz  }t	        j
                  |dd¬«      }|S )z±
        This fallback function is called when overflow is detected during training time, and adjusts the `self.shift`
        to avoid overflow in the subsequent runs.
        r„   TrÉ   )rÒ   r’   r!   rÃ   r   r\   )r#   r—   Úy_int_shiftedrÏ   rÐ   s        r(   Úoverflow_fallbackzIntLayerNorm.overflow_fallbackÕ  sL   € ð
 	�‰�uÔÜ!Ÿ™¨°°4·:±:±Ñ(=Ó>ˆØ  !Ñ#ˆÜ—)‘)˜H¨1°dÔ;ˆØˆr)   c                 óŽ  — | j                   sx|j                  dd¬«      }||z
  }t        j                  |dz  dd¬«      }|t        j                  | j                  |z   «      z  }|| j
                  z  | j                  z   }|d fS | j                  €et        j                  |j                  d   t        j                  ¬«      }t        j                  |«      j                  |j                  «      | _        ||z  }t        j                  |j                  dd¬«      «      }||z
  }	t        j                  |	d| j                   z  z  «      }
|
dz  }t        j"                  |dd¬«      }| j$                  r[|j'                  «       d| j(                  z  k\  r;| j+                  |	«      }|j'                  «       d| j(                  z  dz   k  sJ d«       ‚t        j                  t        j                  |«      «      d| j                   z  z  }t        j                  d|z  «      }t        j                  |	|z  dz  «      }	| j                  dz  }| j                  j,                  j/                  «       | j
                  j,                  j/                  «       z  }t        j                  ||z  «      }|	|z   }	|| j
                  z  }|	|z  }||fS )	Nr„   TrÉ   )Údtypegš™™™™™¹?zfError detected in overflow handling: `var_int` exceeds `self.max_bit` (the maximum possible bit width)l        i   @)r   Úmeanr   rÌ   rÅ   r   ro   rÆ   Útensorrw   ÚfloatÚtoÚdeviceÚ	round_ster!   r’   rÃ   r\   rZ   r1   r§   rÕ   r-   r.   )r#   r3   r“   rØ   ÚyÚvarÚnr{   Úmean_intr—   rÔ   rÏ   rÐ   Ústd_intr¾   ro   Úbias_ints                    r(   r;   zIntLayerNorm.forwardà  sL  € Ø�ŠØ—6‘6˜q¨$�6Ó/ˆDØ�D‘ˆAÜ—*‘*˜Q ™T¨°4Ô8ˆCØ”E—J‘J˜tŸx™x¨#™~Ó.Ñ.ˆAØ�D—K‘K‘ $§)¡)Ñ+ˆAØ�d�7ˆNð �=‰=Ð Ü—‘˜QŸW™W Q™Z¬u¯{©{Ô;ˆAÜ!ŸJ™J q›M×,Ñ,¨Q¯X©XÓ6ˆDŒMð �NÑ"ˆÜ—?‘? 5§:¡:°1¸d :Ó#CÓDˆØ˜Ñ ˆÜ!Ÿ™¨°°4·:±:±Ñ(=Ó>ˆØ  !Ñ#ˆÜ—)‘)˜H¨1°dÔ;ˆð �=Š=à�{‰{‹}  4§<¡<¡Ò/Ø×0Ñ0°Ó7�Ø—{‘{“} q¨$¯,©,¡¸Ñ'<Ò<ð ðXóÐ<ô —/‘/¤%§*¡*¨WÓ"5Ó6¸¸D¿J¹J¹ÑFˆÜ—‘ ¨¡Ó1ˆÜ—‘ ¨¡°Ñ 2Ó3ˆØŸ™¨Ñ.ˆð �y‰y�~‰~×$Ñ$Ó&¨$¯+©+×*:Ñ*:×*AÑ*AÓ*CÑDˆÜ—?‘? 4¨.Ñ#8Ó9ˆà˜Ñ ˆØ'¨$¯+©+Ñ5ˆØ�NÑ"ˆà�.Ð Ð r)   )r=   Frž   r}   )	r@   rA   rB   rC   r   rÒ   rÕ   r;   rD   rE   s   @r(   rÀ   rÀ   ¬  s   ø„ ñ
õPò&\ò	÷.!r)   rÀ   c                 óT  — | j                   d   }t        |d|dz  z
  z  «      }t        ||z  dz  «      }t        j                  | |¬«      j                  }|dk(  r|dz  }n#t        j                  |  |¬«      j                   }|s |j                  «       }|j                  «       }||fS )aÆ  
    Calculate the percentile max and min values in a given tensor

    Args:
        input (`torch.Tensor`):
            The target tensor to calculate percentile max and min.
        lower_percentile (`float`):
            If 0.1, means we return the value of the smallest 0.1% value in the tensor as percentile min.
        upper_percentile (`float`):
            If 99.9, means we return the value of the largest 0.1% value in the tensor as percentile max.
        output_tensor (`bool`, *optional*, defaults to `False`):
            If True, this function returns tensors, otherwise it returns values.

    Returns:
        `Tuple(torch.Tensor, torch.Tensor)`: Percentile min and max value of *input*
    r   r   g{®Gáz„?)r‰   )rw   Úroundr   ÚkthvalueÚvaluesrV   )	ÚinputÚlower_percentileÚupper_percentileÚoutput_tensorÚinput_lengthÚlower_indexÚupper_indexÚupper_boundÚlower_bounds	            r(   Úget_percentile_min_maxrñ     s®   € ð" —;‘;˜q‘>€Lä˜¨Ð,<¸tÑ,CÑ(CÑDÓE€KÜ˜Ð'7Ñ7¸$Ñ>Ó?€Kä—.‘. ¨+Ô6×=Ñ=€Kà˜1ÒØ! A‘o‰ô —~‘~ u f°Ô<×CÑCÐCˆáØ!×&Ñ&Ó(ˆØ!×&Ñ&Ó(ˆØ˜Ð#Ð#r)   c                 óè  — t        | j                  «      dk(  r)|j                  dddd«      }|j                  dddd«      }n_t        | j                  «      dk(  r%|j                  dd«      }|j                  dd«      }n"|j                  d«      }|j                  d«      }|r3| j                  d|z  «      j	                  |«      j                  «        | S t        j                  d|z  | z  |z   «      S )a?  
    Quantize single-precision input tensor to integers with the given scaling factor and zeropoint.

    Args:
        input (`torch.Tensor`):
            Single-precision input tensor to be quantized.
        scale (`torch.Tensor`):
            Scaling factor for quantization.
        zero_pint (`torch.Tensor`):
            Shift for quantization.
        inplace (`bool`, *optional*, defaults to `False`):
            Whether to compute inplace or not.

    Returns:
        `torch.Tensor`: Linearly quantized value of *input* according to *scale* and *zero_point*.
    é   rY   r   r„   rš   )Úlenrw   r^   Úmul_Úadd_Úround_r   rå   )rè   ÚscaleÚ
zero_pointÚinplaces       r(   Úlinear_quantizerû   5  sÒ   € ô$ ˆ5�;‰;Ó˜1ÒØ—
‘
˜2˜q ! QÓ'ˆØ—_‘_ R¨¨A¨qÓ1‰
ä	ˆU�[‰[Ó	˜QÒ	Ø—
‘
˜2˜qÓ!ˆØ—_‘_ R¨Ó+‰
à—
‘
˜2“ˆØ—_‘_ RÓ(ˆ
áØ�
‰
�3˜‘;Ó×$Ñ$ ZÓ0×7Ñ7Ô9ØˆÜ�;‰;�s˜U‘{ UÑ*¨ZÑ7Ó8Ð8r)   c                 óÈ  — t        j                  «       5  d| dz
  z  dz
  }|rht        j                  t        j                  |j	                  «       |j	                  «       gd¬«      d¬«      \  }}t        j
                  |d¬«      |z  }nBt        |j	                  «       |j	                  «       «      }t        j
                  |d¬«      |z  }ddd«       |S # 1 sw Y   S xY w)a/  
    Compute the scaling factor with the given quantization range for symmetric quantization.

    Args:
        saturation_min (`torch.Tensor`):
            Lower bound for quantization range.
        saturation_max (`torch.Tensor`):
            Upper bound for quantization range.
        per_channel (`bool`, *optional*, defaults to `False`):
            Whether to or not use channel-wise quantization.

    Returns:
        `torch.Tensor`: Scaling factor that linearly quantizes the given range between *saturation_min* and
        *saturation_max*.
    r„   r   r¹   g:Œ0âŽyE>r°   N)r   r¬   r1   Ústackr‘   r±   )Únum_bitsÚsaturation_minÚsaturation_maxrO   rà   rø   ry   s          r(   r2   r2   X  sÃ   € ô$ 
�‰‹ñ 	5Ø�(˜Q‘,Ñ !Ñ#ˆáÜ—y‘y¤§¡¨n×.@Ñ.@Ó.BÀN×DVÑDVÓDXÐ-YÐ_`Ô!aÐghÔi‰HˆE�1Ü—K‘K ¨4Ô0°1Ñ4‰Eô ˜×*Ñ*Ó,¨n×.@Ñ.@Ó.BÓCˆEÜ—K‘K ¨4Ô0°1Ñ4ˆE÷	5ð €L÷	5ð €Lús   •B8CÃC!c                   ó0   — e Zd ZdZed„ «       Zed„ «       Zy)r    zw
    Class to quantize the given floating-point values using symmetric quantization with given range and bitwidth.
    c                 óÀ   — t        j                  d|j                  ¬«      }d|dz
  z  dz
  }t        |||d¬«      }t        j                  || |dz
  «      }|| _        |S )a6  
        Args:
            x (`torch.Tensor`):
                Floating point tensor to be quantized.
            k (`int`):
                Quantization bitwidth.
            percentile_mode (`bool`):
                Whether or not to use percentile calibration.
            scale (`torch.Tensor`):
                Pre-calculated scaling factor for *x*. Note that the current implementation of SymmetricQuantFunction
                requires pre-calculated scaling factor.

        Returns:
            `torch.Tensor`: Symmetric-quantized value of *input*.
        g        )rÜ   r„   r   F)rú   )r   rÙ   rÜ   rû   r±   rø   )Úctxr3   r‰   r   rø   rù   rà   Únew_quant_xs           r(   r;   zSymmetricQuantFunction.forward}  s_   € ô" —\‘\ #¨e¯l©lÔ;ˆ
à�!�a‘%‰L˜1ÑˆÜ% a¨°
ÀEÔJˆÜ—k‘k +°¨r°1°q±5Ó9ˆàˆŒ	ØÐr)   c                 ó  — | j                   }t        |j                  «      dk(  r|j                  dddd«      }n<t        |j                  «      dk(  r|j                  dd«      }n|j                  d«      }|j	                  «       |z  d d d d fS )Nró   rY   r   r„   )rø   rô   rw   r^   Úclone)r  Úgrad_outputrø   s      r(   ÚbackwardzSymmetricQuantFunction.backward—  s�   € à—	‘	ˆÜˆ{× Ñ Ó! QÒ&Ø—J‘J˜r 1 a¨Ó+‰Eä�×"Ñ"Ó# qÒ(Ø—J‘J˜r 1Ó%‰Eà—J‘J˜r“NˆEà× Ñ Ó" UÑ*¨D°$¸¸dÐBÐBr)   N©r@   rA   rB   rC   Ústaticmethodr;   r  © r)   r(   r    r    x  s1   „ ñð ñó ðð2 ñ
Có ñ
Cr)   r    c                   ó0   — e Zd ZdZed„ «       Zed„ «       Zy)r’   z;
    Straight-through Estimator(STE) for torch.floor()
    c                 ó,   — t        j                  |«      S r}   )r   r�   ©r  r3   s     r(   r;   zfloor_ste.forwardª  ó   € ä�{‰{˜1‹~Ðr)   c                 ó"   — |j                  «       S r}   ©r  ©r  r  s     r(   r  zfloor_ste.backward®  ó   € à× Ñ Ó"Ð"r)   Nr	  r  r)   r(   r’   r’   ¥  ó/   „ ñð ñó ðð ñ#ó ñ#r)   r’   c                   ó0   — e Zd ZdZed„ «       Zed„ «       Zy)rÝ   z;
    Straight-through Estimator(STE) for torch.round()
    c                 ó,   — t        j                  |«      S r}   )r   rå   r  s     r(   r;   zround_ste.forward¸  r  r)   c                 ó"   — |j                  «       S r}   r  r  s     r(   r  zround_ste.backward¼  r  r)   Nr	  r  r)   r(   rÝ   rÝ   ³  r  r)   rÝ   c                 óÆ  — | j                  «       }| j                  d«      } t        j                  | j	                  «       j                  «       «      \  }}g }|D ]i  }t        t        j                  |d|z  z  «      j                  t        j                  d«      t        j                  ¬«      «      }|j                  |«       Œk t        j                  |«      }t        |«      |z
  }t        j                  |«      j!                  | j"                  «      j                  |«      t        j                  |«      j!                  | j"                  «      j                  |«      fS )zü
    Decompose the scaling factor into mantissa and twos exponent.

    Args:
        scaling_factor (`torch.Tensor`):
            Target scaling factor to decompose.

    Returns:
        ``Tuple(torch.Tensor, torch.Tensor)`: mantisa and exponent
    rY   r„   Ú1)Úrounding)Úsizer^   ÚnpÚfrexpÚcpuÚnumpyrÎ   ÚdecimalÚDecimalÚquantizeÚROUND_HALF_UPÚappendÚarrayrÚ   r   Ú
from_numpyrÛ   rÜ   )Úinputsr§   Úshape_of_inputÚoutput_mÚoutput_eÚtmp_mÚmÚint_m_shifteds           r(   Úbatch_frexpr.  Á  s  € ð —[‘[“]€Nð �[‰[˜‹_€FäŸ™ &§*¡*£,×"4Ñ"4Ó"6Ó7Ñ€HˆhØ€EØò $ˆÜÜ�O‰O˜A  G¡Ñ,Ó-×6Ñ6´w·±ÀsÓ7KÔV]×VkÑVkÐ6Óló
ˆð 	�‰�]Õ#ð	$ô
 �x‰x˜‹€Hä�W‹~ Ñ(€Hô 	×Ñ˜Ó"×%Ñ% f§m¡mÓ4×9Ñ9¸.ÓIÜ×Ñ˜Ó"×%Ñ% f§m¡mÓ4×9Ñ9¸.ÓIðð r)   c                   ó6   — e Zd ZdZe	 	 dd„«       Zed„ «       Zy)r]   aQ  
    Function to perform fixed-point arithmetic that can match integer arithmetic on hardware.

    Args:
        pre_act (`torch.Tensor`):
            Input tensor.
        pre_act_scaling_factor (`torch.Tensor`):
            Scaling factor of the input tensor *pre_act*.
        bit_num (`int`):
            Quantization bitwidth.
        z_scaling_factor (`torch.Tensor`):
            Scaling factor of the output tensor.
        identity (`torch.Tensor`, *optional*):
            Identity tensor, if exists.
        identity_scaling_factor (`torch.Tensor`, *optional*):
            Scaling factor of the identity tensor *identity*, if exists.

    Returns:
        `torch.Tensor`: Output tensor(*pre_act* if *identity* is not given, otherwise the addition of *pre_act* and
        *identity*), whose scale is rescaled to *z_scaling_factor*.
    Nc                 ó  — t        |j                  «      dk(  rd„ }nd„ }|| _        d|dz
  z  dz
  }t        j                  «       5   ||«      }|� ||«      }|| _        t        j                  ||z  «      }	|j                  t        j                  «      }
|j                  t        j                  «      j                  t        j                  «      }|
|z  } ||«      }t        |«      \  }}|	j                  t        j                  «      |j                  t        j                  «      z  }t        j                  |d|z  z  «      }|�ít        j                  ||z  «      }|j                  t        j                  «      }
|j                  t        j                  «      j                  t        j                  «      }|
|z  } ||«      }t        |«      \  }}|j                  t        j                  «      |j                  t        j                  «      z  }t        j                  |d|z  z  «      }||z   }t        j                  |j                  t        j                  «      | dz
  |«      cd d d «       S # 1 sw Y   y xY w)Nr   c                 ó   — | S r}   r  ©r3   s    r(   ú<lambda>z'FixedPointMul.forward.<locals>.<lambda>  s   €  € r)   c                 ó(   — | j                  ddd«      S )Nr   rY   )r^   r2  s    r(   r3  z'FixedPointMul.forward.<locals>.<lambda>  s   €  §¡ q¨!¨RÓ 0€ r)   r„   r   r<   )rô   rw   r`   r   r¬   Úz_scaling_factorrå   ÚtypeÚdoublerÚ   r.  r±   )r  Úpre_actr_   Úbit_numr5  r`   ra   Úreshaperà   Úz_intÚ_AÚ_BÚ	new_scaler,  ÚeÚoutputÚwx_intÚm1Úe1Úoutput1s                       r(   r;   zFixedPointMul.forwardú  s  € ô Ð%×+Ñ+Ó,°Ò1Ù!‰Gá0ˆGØˆŒà�'˜A‘+Ñ Ñ"ˆä�]‰]‹_ñ !	DÙ%,Ð-CÓ%DÐ"ØÐ#Ù*1Ð2IÓ*JÐ'à#3ˆCÔ ä—K‘K Ð*@Ñ @ÓAˆEØ'×,Ñ,¬U¯\©\Ó:ˆBØ"×'Ñ'¬¯©Ó4×:Ñ:¼5¿<¹<ÓHˆBØ˜R™ˆIÙ 	Ó*ˆIä˜yÓ)‰DˆAˆqà—Z‘Z¤§¡Ó-°·±´u·|±|Ó0DÑDˆFÜ—[‘[ ¨3°©6Ñ!2Ó3ˆFàÐ#äŸ™ XÐ0GÑ%GÓH�à,×1Ñ1´%·,±,Ó?�Ø&×+Ñ+¬E¯K©KÓ8×>Ñ>¼u¿|¹|ÓL�Ø ™G�	Ù# IÓ.�	ä$ YÓ/‘��BØ Ÿ+™+¤e§l¡lÓ3°b·g±g¼e¿l¹lÓ6KÑK�ÜŸ+™+ g°°b±Ñ&9Ó:�à  6Ñ)�ä—;‘;˜vŸ{™{¬5¯;©;Ó7¸!¸¸a¹ÀÓC÷C!	D÷ !	Dò !	Dús   ÁH(I8É8Jc                 ó    — d }| j                   �|j                  «       | j                  z  }|j                  «       | j                  z  d d d d |d fS r}   )r`   r  r5  )r  r  Úidentity_grads      r(   r  zFixedPointMul.backward/  sU   € àˆØ�<‰<Ð#Ø'×-Ñ-Ó/°#×2FÑ2FÑFˆMØ× Ñ Ó" S×%9Ñ%9Ñ9¸4ÀÀtÈTÐS`ÐbfÐfÐfr)   r?   r	  r  r)   r(   r]   r]   ã  s<   „ ñð, ð Ø $ò2Dó ð2Dðh ñgó ñgr)   r]   )F)é   )r   r  r  r   r   Útorch.autogradr   Úutilsr   Ú
get_loggerr@   r…   ÚModuler   rG   ri   r   r    rÀ   rñ   rû   r2   r    r’   rÝ   r.  r]   r  r)   r(   ú<module>rL     sð   ðó$ ã Û Ý Ý #å ð 
ˆ×	Ñ	˜HÓ	%€ôPP�R—Y‘Yô PPôfgMˆr�y‰yô gMôTM
�"—)‘)ô M
ô`66ˆb�i‰iô 66ôrD8�—‘ô D8ôNb!�2—9‘9ô b!óJ!$óH 9óFô@*C˜Xô *CôZ#�ô #ô#�ô #óôDQg�Hõ Qgr)   