ó
    EñiŸ  ã                   ó¨  • % S SK r S SKrS SKJr  S SKJr  S SKrS SKJs  J	s  J
r  S SKJ	r	  S SKJ	s  Jr  S SKJr  S SKJrJrJr  S SKJrJr  S SKJr  S SKJr  S	S
KJrJrJ r   / r!\"\   \#S'   S\RH                  S\%S\RH                  S\RH                  4S jr&\RN                  \RP                  \RR                  \RT                  0r+ " S S\	RX                  5      r- " S S\	RX                  5      r.S\-S\.S\RH                  4S jr/ " S S\" SSS/5      5      r0\-Rc                  \Rd                  \RR                  S9r3\.Rc                  \Rh                  \RT                  S9r5\0" \3\5S9r6S\74S jr8S\74S jr9S\74S  jr:S!\S\74S" jr;S#\	RX                  S\74S$ jr<S%\S&\S'\=\>\	RX                  4   S\?\S-  \.S-  4   4S( jr@S)\S'\=\>\	RX                  4   S\S-  4S* jrAS!\S'\=\>\	RX                  4   S\-S-  4S+ jrBS!\S'\=\>\	RX                  4   S\RH                  S-  4S, jrCS!\S'\=\>\	RX                  4   SS4S- jrDS!\S'\=\>\	RX                  4   S.\RH                  S/\RH                  S-  SS4
S0 jrES)\S&\S'\=\>\	RX                  4   S.\RH                  S/\RH                  S-  SS4S1 jrFS)\S'\=\>\	RX                  4   SS4S2 jrGS&\S!\S3\4S4 jrHS&\S'\=\>\	RX                  4   S\=\>\.4   4S5 jrIS&\S'\=\>\	RX                  4   S6\=\>\.4   SS4S7 jrJS&\4S8 jrKS9\	RX                  S:\	RX                  S;\RH                  S\=\>\L4   4S< jrMS=\=\>\L4   S>\%S\4S? jrNg)@é    N)Ú
namedtuple)ÚAny)Ú_get_observed_graph_module_attr)Ú
_with_argsÚObserverBaseÚPerChannelMinMaxObserver)Ú_parent_nameÚcheck_min_max_valid)ÚGraphModule)ÚNodeé   )Úget_new_attr_name_with_prefixÚmaybe_get_next_moduleÚnode_arg_is_weightÚCUSTOM_MODULE_SUPP_LISTÚscaleÚaxisÚinputÚreturnc                 ój   • S/UR                   -  nUR                  U5      X1'   U R                  U5      $ )zMReshapes the scale so that we can multiply it to the input by the given axis.r   )ÚndimÚsizeÚview)r   r   r   Ú	new_shapes       Ú_/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/ao/quantization/fx/_equalize.pyÚreshape_scaler      s1   € à��e—j‘jÑ €IØ—j‘j Ó&€I�OØ�:‰:�iÓ Ð ó    c                   ó�   ^ • \ rS rSrSr\R                  \R                  SSS4 S
U 4S jjjrS r	S r
S rS r\" \5      rS	rU =r$ )Ú_InputEqualizationObserveré,   aö  Observer for tracking the running min/max values of input columns, and
computing the quantization parameters for the overall min/max input values.

Args:
    dtype: Quantized data type
    qscheme: Quantization scheme
    quant_min: Minimum quantization value. If unspecified, it will
        follow the 8-bit setup.
    quant_max: Maximum quantization value. If unspecified, it will
        follow the 8-bit setup.

The running minimum/maximum :math:`x_\text{min/max}` are computed in the
same way as :class:`~torch.ao.quantization.observer.PerChannelMinMaxObserver`,
with the difference that the running min/max values are stored per column.
This observer is intended to be used along with a WeightEqualizationObserver
to calculate the equalization scale.
Nc           	      ó  >• [         TU ]  5         U[        R                  [        R                  1;  a  [        S5      eXl        X l        [        U   n[        SUUUUUS9U l
        [        R                  " S5      U l        / U l        g )Nz Input qscheme must be per-tensorr   ©Úch_axisÚdtypeÚqschemeÚ	quant_minÚ	quant_maxÚfactory_kwargs)ÚsuperÚ__init__ÚtorchÚper_tensor_affineÚper_tensor_symmetricÚ	TypeErrorr$   r%   Ú(qsheme_mapping_per_tensor_to_per_channelr   Ú	input_obsÚtensorÚequalization_scaleÚequalization_shape©Úselfr$   r%   r&   r'   r(   Úper_channel_qschemeÚ	__class__s          €r   r*   Ú#_InputEqualizationObserver.__init__?   s‚   ø€ ô 	‰ÑÔàœ5×2Ñ2´E×4NÑ4NÐOÓOÜÐ>Ó?Ð?àŒ
ØŒäFÀwÑOÐÜ1ØØØ'ØØØ)ñ
ˆŒô #(§,¢,¨q£/ˆÔØ-/ˆÕr   c                 óà   • UR                   S:  d  UR                   S:”  a  [        S5      eS/UR                   -  U l        UR                  S5      U R                  S'   U R	                  U5      $ )Né   é   ú>InputEqualizationObserver only supports Linear and Conv layersr   )r   Ú
ValueErrorr3   r   r0   )r5   Úx_origs     r   ÚforwardÚ"_InputEqualizationObserver.forward\   sc   € Ø�;‰;˜‹?˜fŸk™k¨A›oÜØPóð ð
 $% #¨¯©Ñ"3ˆÔØ%+§[¡[°£^ˆ×Ñ Ñ"à�~‰~˜fÓ%Ð%r   c                 óZ   • U R                   R                  U R                   R                  4$ ©N)r0   Úmin_valÚmax_val©r5   s    r   Úget_input_minmaxÚ+_InputEqualizationObserver.get_input_minmaxh   s!   € Ø—‘×&Ñ&¨¯©×(>Ñ(>Ð?Ð?r   c                 ó¬   • UR                  5       S:X  a  U[        R                  " S5      :X  a  g [        R                  " XR                  5      U l        g )Nr   )Únelementr+   r1   Úreshaper3   r2   ©r5   r2   s     r   Úset_equalization_scaleÚ1_InputEqualizationObserver.set_equalization_scalek   sC   € ð ×&Ñ&Ó(¨AÓ-Ð2DÌÏÊÐUVËÓ2WØÜ"'§-¢-Ø× 7Ñ 7ó#
ˆÕr   c                 ó²  • U R                   R                  5       S:X  a:  U R                   [        R                  " S5      :X  a  [        R
                  " SSS9  gU R                  5       u  p[        U R                   SU5      n[        R                  " [        R                  " X5      5      n[        R                  " [        R                  " X#5      5      nXE4$ )z!Returns the scaled min/max inputsr   z}Must call calculate_equalization_scale before calling calculate_scaled_minmax. Will not scale the next quantization observer.r:   ©Ú
stacklevel©NNr   )r2   rI   r+   r1   ÚwarningsÚwarnrF   r   ÚminÚmulÚmax)r5   Ú
min_inputsÚ
max_inputsÚequalization_scale_reshapedÚmin_input_scaledÚmax_input_scaleds         r   Úcalculate_scaled_minmaxÚ2_InputEqualizationObserver.calculate_scaled_minmaxt   s±   € ð ×#Ñ#×,Ñ,Ó.°!Ó3Ø×'Ñ'¬5¯<ª<¸«?Ó:ä�MŠMðCàòð
 ð
 $(×#8Ñ#8Ó#:Ñ ˆÜ&3Ø×#Ñ# Q¨
ó'
Ð#ô !Ÿ9š9¤U§Y¢Y¨zÓ%WÓXÐÜ Ÿ9š9¤U§Y¢Y¨zÓ%WÓXÐàÐ1Ð1r   )r$   r2   r3   r0   r%   ©r   N)Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r+   Úquint8r,   r*   r?   rF   rL   r\   Úclassmethodr   Ú	with_argsÚ__static_attributes__Ú__classcell__©r7   s   @r   r   r   ,   sX   ø† ñð( �l‰lØ×'Ñ'ØØØð0ð 
÷0ð 0ò:
&ò@ò
ò2ñ2 ˜JÓ'†Ir   r   c                   óŠ   ^ • \ rS rSrSr\R                  \R                  SSS4 S	U 4S jjjrS r	S r
S r\" \5      rSrU =r$ )
Ú_WeightEqualizationObserveré�   aC  Observer for tracking the running min/max values of weight columns and
rows, and computing the quantization parameters for the weight rows.

Args:
    dtype: Quantized data type
    qscheme: Quantization scheme
    quant_min: Minimum quantization value. If unspecified, it will
        follow the 8-bit setup.
    quant_max: Maximum quantization value. If unspecified, it will
        follow the 8-bit setup.

This observer is made up of 1 PerChannelMinMaxObserver `weight_col_obs` used
to record the running minimum and maximum of columns of incoming weight
tensors. This observer is intended to be used along with an
InputEqualizationObserver to calculate the equalization scale.

The running minimum/maximum :math:`w_\text{min/max}` are computed in the
same way as :class:`~torch.ao.quantization.observer.PerChannelMinMaxObserver`.
Nc           	      ó  >• [         TU ]  5         Xl        X l        SU l        UnU[
        R                  [
        R                  1;   a	  [        U   n[        SUUUUUS9U l
        [
        R                  " S5      U l        g )Nr   r"   )r)   r*   r$   r%   r#   r+   r,   r-   r/   r   Úweight_col_obsr1   r2   r4   s          €r   r*   Ú$_WeightEqualizationObserver.__init__¥   s|   ø€ ô 	‰ÑÔàŒ
ØŒØˆŒà%ÐØ”u×.Ñ.´×0JÑ0JÐKÓKÜ"JÈ7Ñ"SÐÜ6ØØØ'ØØØ)ñ
ˆÔô #(§,¢,¨q£/ˆÕr   c                 óz   • UR                   S:  d  UR                   S:”  a  [        S5      eU R                  U5      $ )Nr:   r;   r<   )r   r=   rn   )r5   Úw_origs     r   r?   Ú#_WeightEqualizationObserver.forwardÁ   s:   € Ø�;‰;˜‹?˜fŸk™k¨A›oÜØPóð ð ×"Ñ" 6Ó*Ð*r   c                 óZ   • U R                   R                  U R                   R                  4$ rB   )rn   rC   rD   rE   s    r   Úget_weight_col_minmaxÚ1_WeightEqualizationObserver.get_weight_col_minmaxÉ   s%   € Ø×#Ñ#×+Ñ+¨T×-@Ñ-@×-HÑ-HÐIÐIr   c                 ó   • Xl         g rB   )r2   rK   s     r   rL   Ú2_WeightEqualizationObserver.set_equalization_scaleÌ   s   € Ø"4Õr   )r#   r$   r2   r%   rn   r^   )r_   r`   ra   rb   rc   r+   Úqint8r,   r*   r?   rt   rL   re   r   rf   rg   rh   ri   s   @r   rk   rk   �   sS   ø† ñð, �k‰kØ×'Ñ'ØØØð2ð 
÷2ð 2ò8+òJò5ñ ˜JÓ'†Ir   rk   r0   Ú
weight_obsc                 óà  • U R                  5       u  p#UR                  5       u  pE[        X#5      (       a  [        XE5      (       d+  [        R                  " SSS9  [
        R                  " S5      $ UR                  UR                  :w  a)  [        SSUR                   SUR                   S3-   5      e[
        R                  " XT-
  X2-
  -  5      nSXfS	:H  '   [
        R                  " USSSS
9nU$ )zíCalculates the equalization scale and sets the equalization_scale value
in the observers.

Args:
    input_obs: Observer that tracks the ranges for the input columns
    weight_obs: Observer that tracks the ranges for the weight columns
ztMust run observer before calling calculate_equalization_scale. Returning default equalization scale torch.tensor(1).r:   rO   r   z6Input and Weight must have the same column dimension. zFound z and z shapes instead.g        )ÚnanÚposinfÚneginf)rF   rt   r
   rR   rS   r+   r1   Úshaper=   ÚsqrtÚ
nan_to_num)r0   ry   rW   rX   Úmin_weightsÚmax_weightsr2   s          r   Úcalculate_equalization_scalerƒ   Ò   sü   € ð  )×9Ñ9Ó;Ñ€ZØ!+×!AÑ!AÓ!CÑ€[ô 	˜J×3Ñ3Ü ×9Ñ9ä�ŠðFàò	
ô
 �|Š|˜A‹Ðà×Ñ˜;×,Ñ,Ó,ÜØDØ�z×'Ñ'Ð(¨¨k×.?Ñ.?Ð-@Ð@PÐQñRó
ð 	
ô
 ŸšØ	Ñ	" zÑ'>Ñ?óÐð 56Ð¨SÑ0Ñ1Ü×)Ò)Ð*<À!ÈAÐVWÑXÐØÐr   c                   óˆ   ^ • \ rS rSrSrSr\R                  R                  \R                  R                  4U 4S jjr	Sr
U =r$ )ÚEqualizationQConfigéú   a.  
Describes how to quantize a layer or a part of the network specifically for
input-weight equalization by providing settings (observer classes) for
inputs, outputs, and weights.

Note that EqualizationQConfig needs to contain observer **classes** (like
MinMaxObserver) or a callable that returns instances on invocation, not the
concrete observer instances themselves.
Quantization function will instantiate observers multiple times for each of
the layers.

Observer classes have usually reasonable default arguments, but they can be
overwritten with `with_args` method (that behaves like functools.partial):

my_qconfig = EqualizationQConfig(input_activation=_InputEqualizationObserver.with_args(dtype=torch.qint8),
                                weight=_WeightEqualizationObserver.with_args(dtype=torch.qint8))
© c                 óº   >• [        U[        R                  5      (       d  [        U[        R                  5      (       a  [        S5      e[        TU ]  XU5      nU$ )Nz EqualizationQConfig received observer instance, please pass observer class instead. Use MyObserver.with_args(x=1) to override arguments to constructor if needed)Ú
isinstanceÚnnÚModuler=   r)   Ú__new__)ÚclsÚinput_activationÚweightr5   r7   s       €r   rŒ   ÚEqualizationQConfig.__new__  sO   ø€ ÜÐ&¬¯	©	×2Ñ2´jÀÌÏÉ×6SÑ6SÜðaóð ô ‰w‰˜s°fÓ=ˆØˆr   )r_   r`   ra   rb   rc   Ú	__slots__r+   rŠ   ÚIdentityrŒ   rg   rh   ri   s   @r   r…   r…   ú   s2   ø† ñð$ €Ià&+§h¡h×&7Ñ&7ÀÇÁ×@QÑ@Q÷ õ r   r…   rŽ   r�   )r$   r%   )rŽ   r�   c                 ó–   • [        U 5      [        R                  [        R                  [        R                  [        R
                  4;   $ )z/Checks if the fused node supports equalization.)ÚtypeÚnniÚ
LinearReLUÚ
ConvReLU1dÚ
ConvReLU2dÚ
ConvReLU3d©Úmodules    r   Ú"fused_module_supports_equalizationrœ   '  s4   € ä�‹<Ü�‰Ü�‰Ü�‰Ü�‰ð	ñ ð r   c                 ó–   • [        U 5      [        R                  [        R                  [        R                  [        R
                  4;   $ )z2Checks if the torch.nn node supports equalization.)r”   rŠ   ÚLinearÚConv1dÚConv2dÚConv3drš   s    r   Únn_module_supports_equalizationr¢   1  s*   € ä�‹<œBŸI™I¤r§y¡y´"·)±)¼R¿Y¹YÐGÑGÐGr   c                 ó&   • [        U 5      [        ;   $ )z0Checks if the custom node supports equalization.)r”   r   rš   s    r   Ú#custom_module_supports_equalizationr¤   6  s   € ä�‹<Ô2Ñ2Ð2r   Únodec                 ó¼  • U R                   S:X  aq  [        U[        U R                  5         5      =(       dI    [	        U[        U R                  5         5      =(       d!    [        U[        U R                  5         5      $ U R                   S:X  aK  U R                  [        R                  [        R                  [        R                  [        R                  4;   $ g)zxChecks if the current node supports equalization
Currently we only support nn.Linear/F.Linear and nn.Conv/F.conv layers
Úcall_moduleÚcall_functionF)Úopr¢   ÚstrÚtargetrœ   r¤   ÚFÚlinearÚconv1dÚconv2dÚconv3d)r¥   Úmoduless     r   Únode_supports_equalizationr²   ;  sœ   € ð ‡w�w�-Óä+¨G´C¸¿¹Ó4DÑ,EÓF÷ NÜ1°'¼#¸d¿k¹kÓ:JÑ2KÓL÷Nä2°7¼3¸t¿{¹{Ó;KÑ3LÓMð	
ð
 
�‰�OÓ	#Ø�{‰{œqŸx™x¬¯©´1·8±8¼Q¿X¹XÐFÑFÐFØr   Úobserverc                 ó.   • [        U [        [        45      $ rB   )r‰   r   rk   )r³   s    r   Úis_equalization_observerrµ   J  s   € ÜØÔ-Ô/JÐKóð r   Úinput_eq_obs_nodeÚmodelr±   c                 ó|  • SnU R                    H  n[        XB5      (       d  M  Un  O   Uc  [        S5      eUR                  S:X  aœ  [	        US5      nUc  [        S5      eUnUR                  UR                  5      c  [        SUR                   35      eUR                  UR                  5      R                  5       n[        U[        5      (       d  [        S5      eX74$ UR                  S:X  aI  [        X25      nUb;  U[        UR                  5         n[        U[        5      (       d  [        S5      eX74$ g	)
a  Gets the following weight equalization observer. There should always
exist a weight equalization observer after an input equalization observer.

Returns the operation node that follows the input equalization observer node
and the weight equalization observer
Nz@Expected an operation node after the input equalization observerr§   Ú!equalization_node_name_to_qconfigzOExpected 'equalization_node_name_to_qconfig' attribute in observed graph modulez*No equalization qconfig found for op node úIExpected weight equalization observer to be a _WeightEqualizationObserverr¨   rQ   )Úusersr²   ÚAssertionErrorr©   r   ÚgetÚnamer�   r‰   rk   Úmaybe_get_weight_eq_obs_noderª   r«   )	r¶   r·   r±   Úop_nodeÚuserÚ&maybe_equalization_node_name_to_configr¹   Úweight_eq_obsÚweight_nodes	            r   Úget_op_node_and_weight_eq_obsrÅ   U  s^  € ð €GØ!×'Ô'ˆÜ% d×4Ó4ØˆGÙñ (ð
 �ÜØNó
ð 	
ð ‡z�z�]Ó"ô 2QØÐ6ó2
Ð.ð 2Ñ9Ü Øaóð ð 3ð 	*ð -×0Ñ0°·±Ó>ÑFÜ Ø<¸W¿\¹\¸NÐKóð ð :×=Ñ=¸g¿l¹lÓK×RÑRÓTˆä˜-Ô)D×EÑEÜ Ø[óð ð Ð%Ð%à	�‰�Ó	&Ü2°7ÓDˆØÑ"Ø#¤C¨×(:Ñ(:Ó$;Ñ<ˆMÜ˜mÔ-H×IÑIÜ$Ø_óð ð Ð)Ð)àr   rÀ   c                 ó4  • U R                   S:w  a  [        S5      eU R                   Hm  n[        X5      (       d  M  [	        U[
        5      (       d  M,  UR                   S:X  d  M>  [	        U[        UR                  5         [        5      (       d  Mk  Us  $    g)z8Gets the weight equalization observer node if it exists.r¨   z<maybe_get_weight_eq_obs_node expects a call_function op_noder§   N)	r©   r¼   Úargsr   r‰   r   rª   r«   rk   )rÀ   r±   Únode_args      r   r¿   r¿   �  s�   € ð ‡z�z�_Ó$ÜØJó
ð 	
ð —L”LˆÜ˜g×0Ó0ä˜8¤T×*Ó*Ø—K‘K =Õ0ÜØœC §¡Ó0Ñ1Ô3N÷ó ð  ’ñ !ð r   c                 óx  • [        X5      (       d  [        S5      e[        X[        R                  5      nUc  [        X[
        R                  S9nUc  [        X[        5      O[        X![        5      nUc  g[        X1[        5      nUc  gU[        U5         n[        U[        5      (       d  [        S5      eU$ )aÞ  Gets the following input equalization observer if it exists.

For example, in the case of connecting linear layers:
    x -> inp_obs1 -> eq_obs1 -> linear1 -> out_obs1 -> eq_obs2 -> linear2 -> out_obs2
If the node being passed in is the linear1 node, then we want to return eq_obs2,
the following equalization observer for linear2.

However, if there are no connecting layers:
    x -> inp_obs1 -> eq_obs1 -> linear1 -> out_obs1 -> add
Then we want to return None.

In the case of an unfused linear-relu layer with a connecting linear layer:
    linear1 -> relu -> out_obs1 -> eq_obs2 -> linear2 -> out_obs2
Since it is unfused, we want to skip over the relu layer and return eq_obs2,
the following equalization observer for linear2.
z"Node does not support equalizationN)Útarget_functional_typezPExpected the following equalization observer to be an _InputEqualizationObserver)r²   r¼   r   rŠ   ÚReLUr¬   Úrelur   r   rª   r‰   )r¥   r±   Úmaybe_relu_nodeÚmaybe_obs_nodeÚmaybe_eq_obs_nodeÚmaybe_eq_obss         r   Úmaybe_get_next_input_eq_obsrÑ   ¥  sÆ   € ô( & d×4Ñ4ÜÐAÓBÐBô ,¨D¼2¿7¹7ÓC€OØÑÜ/Ø´!·&±&ñ
ˆð Ñ"ô 	˜d¬\Ô:ä" ?¼\ÓJð ð
 ÑØä-ØÔ!;óÐð Ñ Øàœ3Ð0Ó1Ñ2€LÜ�lÔ$>×?Ñ?ÜØ^ó
ð 	
ð Ðr   c                 óÆ   • [        X5      nU(       aO  UR                  R                  5       S:X  a%  UR                  [        R                  " S5      :X  a  gUR                  $ g)aA  If the next next node is an InputEqualizationObserver then we want to
return its equalization scale, else we return 1

This is used in the case where there are two connecting linear layers:
    linear1 -> LinearOutObs -> InputEqObs -> linear2
In this case, the node given is linear1 and we want to locate the InputEqObs.
r   N)rÑ   r2   rI   r+   r1   )r¥   r±   Únext_inp_eq_obss      r   Ú!maybe_get_next_equalization_scalerÔ   Û  sP   € ô 2°$Ó@€Oæà×.Ñ.×7Ñ7Ó9¸QÓ>Ø×2Ñ2´e·l²lÀ1³oÓEàØ×1Ñ1Ð1Ør   c                 óx  • U[        U R                  5         n[        U[        5      (       d  [	        S5      eU R
                  S   n[        U[        5      (       d  [	        S5      eU[        UR                  5         n[        U[        5      (       d  gUR                  5       u  pVUc  Uc  gXTl	        Xdl
        g)z¦Scales the following input quantization observer's min/max values by
updating the values with the scaled min/max values calculated by the input
equalization observer
zFExpected the module at node.target to be an _InputEqualizationObserverr   z:Expected the input quantization observer node to be a NodeN)rª   r«   r‰   r   r¼   rÇ   r   r   r\   rC   rD   )r¥   r±   Úinput_eq_obsÚinput_quant_obs_nodeÚinput_quant_obsrZ   r[   s          r   Úscale_input_observerrÙ   ñ  s¸   € ð
 œ3˜tŸ{™{Ó+Ñ,€LÜ�lÔ$>×?Ñ?ÜØTó
ð 	
ð  Ÿ9™9 Q™<ÐÜÐ*¬D×1Ñ1ÜØHó
ð 	
ð œcÐ"6×"=Ñ"=Ó>Ñ?€OÜ�o¤|×4Ñ4Øà)5×)MÑ)MÓ)OÑ&ÐØÑÐ$4Ñ$<ØØ.ÔØ.Õr   r2   Únext_equalization_scalec                 óœ  • Uc  g[        U[        U R                  5         5      (       a  U[        U R                  5         S   nOU[        U R                  5         n[        U5      (       d  [	        U5      (       d  [        S5      eUR                  n[        U[        R                  5      (       d  [        S5      e[        USU5      n[        R                  " U[        R                  " U5      5      nUc  [        R                  " U5      Ul        g[        USU5      n[        R                  " Xx5      n[        R                  " U5      Ul        UR                  n	U	c  g[        U	[        R                  5      (       d  [        S5      e[        USU	5      n[        R                  " X˜5      n
[        R                  " U
5      Ul        g)a‡  Scale the weights for input-weight equalization by multiplying the
weight by 1/equalization_scale and next_equalization_scale

Args:
    node: Current node whose weights we want to scale
    equalization_scale: Current node's calculated equalization scale
    next_equalization_scale: Next node's calculated equalization scale if
       the following node needs to be equalized, 1 otherwise
Nr   z@Expected operation module to support equalization (nn or custom)z.Expected op_module.weight to be a torch.Tensorr   z,Expected op_module.bias to be a torch.Tensor)rœ   rª   r«   r¢   r¤   r¼   r�   r‰   r+   ÚTensorr   rU   Ú
reciprocalrŠ   Ú	ParameterÚbias)r¥   r±   r2   rÚ   Ú	op_moduler�   rY   Úscaled_weightÚ next_equalization_scale_reshapedrß   Úscaled_biass              r   Úscale_weight_noderä     s|  € ð Ñ!Øä)¨'´#°d·k±kÓ2BÑ*C×DÑDØœC §¡Ó,Ñ-¨aÑ0‰	àœC §¡Ó,Ñ-ˆ	ä'¨	×2Ñ2Ü.¨y×9Ñ9äØNó
ð 	
ð ×Ñ€FÜ�fœeŸl™l×+Ñ+ÜÐMÓNÐNô #0Ð0BÀAÀvÓ"NÐÜ—I’I˜f¤e×&6Ò&6Ð7RÓ&SÓT€MàÑ&ÜŸ<š<¨Ó6ˆ	ÔØô (5Ð5LÈaÐQWÓ'XÐ$Ü—I’I˜mÓN€Mä—|’| MÓ2€IÔð �>‰>€DØ�|ØÜ�dœEŸL™L×)Ñ)ÜÐKÓLÐLô (5Ð5LÈaÐQUÓ'VÐ$Ü—)’)˜DÓC€KÜ—\’\ +Ó.€I…Nr   c                 ó   • Uc  g[        X5      nUc  gUR                  S   nUc  g[        U[        5      (       a+  [        U[	        UR
                  5         [        5      (       d  [        S5      eUR                  S   nUc  g[        U[        5      (       a  UR                  S:X  d  [        S5      e[        UR
                  5      u  p‰[        X(   U	5      n
[        USU
5      n[        R                  " U
[        R                  " U5      5      nUc  [        X(   Xœ5        g[        USU5      n[        R                  " XÍ5      n[        X(   Xœ5        [        R                   " UR#                  [	        UR
                  5      5      U5      (       d  [        S5      eSnU R                   H@  n[        U[        5      (       d  M  UR                  S:X  d  M,  SUR$                  ;   d  M>  Un  O   Uc  g[        UR
                  5      u  nn[        UU   U5      n[        USU5      n[        R                  " UU5      n[        UU   UU5        g)	z-Scales the weight value for functional layersNr   zKExpected weight_quant_obs_node to be a Node whose module is an ObserverBaseÚget_attrz,Expected weight node to be a 'get_attr' Noder   z8Model buffer for weight does not match the scaled weightrß   )r¿   rÇ   r‰   r   rª   r«   r   r¼   r©   r	   Úgetattrr   r+   rU   rÝ   ÚsetattrÚallcloseÚ
get_bufferr¾   )rÀ   r·   r±   r2   rÚ   Úweight_eq_obs_nodeÚweight_quant_obs_noderÄ   Úweight_parent_nameÚweight_namer�   rY   rá   râ   Ú	bias_noder¥   Úbias_parent_nameÚ	bias_namerß   rã   s                       r   Úscale_weight_functionalrò   N  s.  € ð Ñ!Øô 6°gÓGÐØÑ!Øð /×3Ñ3°AÑ6ÐØÑ$ØäÐ(¬$×/Ñ/Ü�wœsÐ#8×#?Ñ#?Ó@ÑAÄ<×PÑPäØYó
ð 	
ð
 (×,Ñ,¨QÑ/€KØÑØÜ�{¤D×)Ñ)¨k¯n©nÀ
Ó.JÜÐKÓLÐLä&2°;×3EÑ3EÓ&FÑ#ÐÜ�WÑ0°+Ó>€Fô
 #0Ð0BÀAÀvÓ"NÐÜ—I’I˜f¤e×&6Ò&6Ð7RÓ&SÓT€MàÑ&Ü�Ñ+¨[ÔHØô (5Ø  Mó(Ð$ô —I’I˜mÓN€MäˆGÑ'¨ÔDÜ�>Š>˜%×*Ñ*¬3¨{×/AÑ/AÓ+BÓCÀ]×SÑSÜÐWÓXÐXð €IØ—”ˆä�dœD×!Ó! d§g¡g°Õ&;ÀÈ$Ï)É)Õ@SØˆIÙñ	 ð
 ÑØä".¨y×/?Ñ/?Ó"@ÑÐ�iÜ�7Ð+Ñ,¨iÓ8€Dô (5Ð5LÈaÐQUÓ'VÐ$Ü—)’)˜DÐ"BÓC€KÜˆGÐ$Ñ% y°+Õ>r   c                 óD  • [        X5      nUc  gUR                  S   nUc  g[        U[        5      (       d  [	        S5      eU[        UR                  5         n[        U[        UR                  5         [        5      (       d  [	        S5      eUR                  5         g)zlGiven the operation node, we want find the corresponding quantization
observer and reset its min/max values
Nr   z+Expected weight_quant_obs_node to be a NodezBExpected the module at weight_quant_obs_node to be an ObserverBase)	r¿   rÇ   r‰   r   r¼   rª   r«   r   Úreset_min_max_vals)rÀ   r±   rë   rì   Úweight_quant_obss        r   Úclear_weight_quant_obs_noderö   ¢  sŸ   € ô 6°gÓGÐØÑ!Øà.×3Ñ3°AÑ6ÐØÑ$ØÜÐ+¬T×2Ñ2ÜÐJÓKÐKàœsÐ#8×#?Ñ#?Ó@ÑAÐÜ�gœcÐ"7×">Ñ">Ó?Ñ@Ä,×OÑOÜØPó
ð 	
ð ×'Ñ'Õ)r   Ú	prev_nodec                 ó´   • [        UR                  R                  5       5      nU H  nUR                  X5        M     U R                  R                  U5        g)zaRemoves the given node from the model by replacing all of its users with
the given previous node
N)Úlistr»   ÚkeysÚreplace_input_withÚgraphÚ
erase_node)r·   r¥   r÷   Ú
orig_usersÚ	user_nodes        r   Úremove_noder   ¸  sE   € ô �d—j‘j—o‘oÓ'Ó(€JÛˆ	Ø×$Ñ$ TÖ5ñ  ð 
‡K�K×Ñ˜4Õ r   c                 óþ  • 0 nU R                   R                   GH_  nUR                  S:X  d  M  [        XR                     [
        5      (       d  M9  XR                     n[        U[
        5      (       d  [        S5      e[        X0U5      u  pVUb  Uc  M}  UR                  S:X  a—  [        U[        UR                  5         5      (       aI  U[        UR                  5         S   n[        U5      (       d  [        S5      eU" UR                  5        O(U" U[        UR                  5         R                  5        [        XF5      nUR                  U5        UR                  U5        XbUR                  '   GMb     U$ )ag  Update all of the observer's equalization scale. For each
InputEqualizationObserver, we will find the location of the next
WeightEqualizationObserver, create it, and calculate the equalization scale
based on the two observers.

We will then return a dictionary mapping operation node names to
the corresponding WeightEqualizationObservers for that operation.
r§   zBExpected module at node.target to be an _InputEqualizationObserverr   z-Expected fused module to support equalization)rü   Únodesr©   r‰   r«   r   r¼   rÅ   rœ   rª   r¢   r�   rƒ   rL   r¾   )	r·   r±   Úweight_eq_obs_dictr¥   rÖ   rÀ   rÃ   r›   r2   s	            r   Úupdate_obs_for_equalizationr  Æ  sY  € ð ÐØ—‘×!Õ!ˆØ�7‰7�mÕ#¬
Ø—K‘KÑ Ô"<÷)
ó )
ð #§;¡;Ñ/ˆLÜ˜lÔ,F×GÑGÜ$ØXóð ô &CÀ4ÐPWÓ%XÑ"ˆGà‰ -Ñ"7Ùà�z‰z˜]Ó*ô 6°g¼cÀ'Ç.Á.Ó>QÑ6R×SÑSØ$¤S¨¯©Ó%8Ñ9¸!Ñ<�FÜ:¸6×BÑBÜ,ØKóð ñ " &§-¡-Õ0á! '¬#¨g¯n©nÓ*=Ñ">×"EÑ"EÔFô ">Øó"Ðð ×/Ñ/Ð0BÔCØ×0Ñ0Ð1CÔDà/<˜wŸ|™|Ô,ñE "ðH Ðr   r  c           	      ót  • U R                   R                   GHù  nUR                  S:X  Gan  [        XR                     [
        5      (       GaL  UR                  S   nUR                  S   n[        XQ5      (       d  SUR                  ;   a  [        XU5        Mƒ  [        X15        U R                   R                  U5         [        UR                  S-   5      nU" U5      n[        XXR                     R                  5        U R                   R                  SU5      nSSS5        U R                   R!                  W5         XX4n	U R                   R                  S["        R$                  U	5      n
SSS5        UR'                  UW
5        [        XU5        GMƒ  UR)                  UR                  5      c  GM¢  UR)                  UR                  5      n[        U[*        5      (       d  [-        S5      eUR                  nUR/                  5       S	:X  a  U["        R0                  " S	5      :X  a  Sn[3        X15      nUR                  S:X  a  [5        UUUU5        GME  UR                  S:X  a~  [7        UU UUU5        [9        X15      nUc    g[        U[;        UR                  5         [*        5      (       d  [-        S5      e[=        X15        UR                  S   n[        XU5        GMÓ  [?        S
SUR                   SUR                   S3-   5      e   g! , (       d  f       GNþ= f! , (       d  f       GN¾= f)a  Converts the equalization operations and updates the other nodes in the
following way:
    - Removes the input equalization observers and inserts a mul operator
      along with an equalization scale node wherever applicable (we do not
      want to insert a mul operator between connecting linear layers).
    - Updates the input quantization observers with the scaled input min/max
      values.
    - Scales the weights by the current and next equalization scales.
    - Removes the weight equalization observer node if it exists.

Before (after prepare):
                                weight values
                                      |
                                WeightQuantObs
                                      |
                                  WeightEqObs
                                      |
    x -> InpQuantObs -> InpEqObs -> linear -> OutQuantObs

After this function:
                                          scaled weight values
                                                  |
   equalization scale                       WeightQuantObs
          |                                       |
    x -> mul -> InpQuantObs (scaled min/max) -> linear -> OutQuantObs

After convert:
   equalization scale                 scaled weight values
          |                                    |
    x -> mul -> quantize_per_tensor -> quantized::linear

Note that although the equalization observer appeared after the quantization
observer after prepare_fx, the mul node appears before the quantization node
after convert_fx. This is because placing the equalization observer after
the quantization observer in prepare_fx would allow us to keep the invariant
that the graph before the current node inserts its observers is not
modified.

Having the equalization observer before the quantization observer would also
cause some inconsistences between the ordering of the quantization and
equalization observers.
For example, a single linear layer would look like:
    x -> InpEqObs1 -> InpQuantObs1 -> linear1 -> OutQuantObs1
But between two connected linear layers, it would look like:
    linear1 -> OutQuantObs1 -> InpEqObs2 -> linear2 -> OutQuantObs2
r§   r   rÌ   Ú_equalization_scaleræ   Nr¨   rº   r   z=Expected operation node to be 'call_module' or 'call_functionzInstead got node z as 'z'.) rü   r  r©   r‰   r«   r   rÇ   r²   r¾   r   rÙ   Úinserting_beforer   rè   r2   Úcreate_nodeÚinserting_afterr+   rU   rû   r½   rk   r¼   rI   r1   rÔ   rä   rò   r¿   rª   rö   r=   )r·   r±   r  r¥   Úinp_quant_obs_noder÷   Úget_new_eq_scale_namer¾   Úeq_scale_nodeÚinputsÚmul_noderÃ   r2   Úmaybe_next_equalization_scalerë   s                  r   Úconvert_eq_obsr  ù  sß  € ðf —‘×!Õ!ˆØ�7‰7�mÔ#¬
Ø—K‘KÑ Ô"<÷)
ò )
ð "&§¡¨1¡ÐØ*×/Ñ/°Ñ2ˆIô +¨9×>Ñ>Ø˜YŸ^™^Ó+ä˜EÐ);Ô<Ùô ! Ô/ð —‘×-Ñ-Ð.@ÕAÜ(EØ—N‘NÐ%:Ñ:ó)Ð%ñ -¨WÓ5�Ü˜ W¯[©[Ñ%9×%LÑ%LÔMØ %§¡× 7Ñ 7¸
ÀDÓ I�÷ Bð —‘×,Ñ,¨]Õ;Ø#Ð3�Ø Ÿ;™;×2Ñ2°?ÄEÇIÁIÈvÓV�÷ <ð ×1Ñ1°)¸XÔFÜ˜Ð%7×8à×#Ñ# D§I¡IÓ.Ô:Ø.×2Ñ2°4·9±9Ó=ˆMÜ˜mÔ-H×IÑIÜ$Ø_óð ð "/×!AÑ!AÐð #×+Ñ+Ó-°Ó2Ø&¬%¯,ª,°q«/Ó9à%)Ð"Ü,MØó-Ð)ð
 �w‰w˜-Ó'Ü!ØØà&Ø1÷ð —‘˜OÓ+Ü'ØØØà&Ø1ôô &BÀ$Ó%PÐ"Ø%Ñ-ÙÜ!ØœCÐ 2× 9Ñ 9Ó:Ñ;Ô=X÷ñ ô )Øcóð ô ,¨DÔ:ð /×3Ñ3°AÑ6�	Ü˜E°y×Aä ØSØ)¨$¯)©)¨°E¸$¿'¹'¸À"ÐEñFóð òK "÷: BÖAú÷ <Ö;ús   Ã ALÅ/L(Ì
L%	Ì(
L7	c                 óŠ   • [        U R                  SS95      n[        X5      n[        XU5        [	        X R
                  5      $ )zbReference function which applies changes needed for equalization, but
does not quantize the nodes
F)Úremove_duplicate)ÚdictÚnamed_modulesr  r  r   rü   )r·   r±   r  s      r   Ú_convert_equalization_refr  —  sC   € ô �5×&Ñ&¸Ð&Ð>Ó?€Gô 5°UÓDÐÜ�5Ð#5Ô6ä�uŸk™kÓ*Ð*r   Úmodel_aÚmodel_bÚxc           	      ó  • SSK Js  Js  Jn  SSKJn  U" 5       nUS   R                  [        R                  5        UR                  SU SUUR                  US9u  pgU" U5        U" U5        UR                  XgUR                  S5      nUR                  USS[        R                  R                  R                  R                  R                  S5        0 n	U H*  n
XŠ   S	   S   S   S
   nXŠ   S	   S   S   S   S   nXÉU'   M,     U	$ )a  Runs the Numeric Suite on model_a and model_b and returns a dictionary
containing the SQNR between layers in model_a and model_b.

Note: In order to support equalized models, this function has a hacky fix in
which we do not match any torch.mul operators. This is because equalized
models contain extra mul operators to scale the input by the equalization
scale, but this edge case has not been resolved yet within the numeric suite code.

Args:
    model_a: A float model
    model_b: A quantized model
    x: Inputs to use during calibration
r   N)Úget_unmatchable_types_mapÚfuns_unmatchableÚfp32Úint8)Úunmatchable_types_mapÚsqnrÚnode_outputÚfqn)Útorch.ao.ns._numeric_suite_fxÚaoÚnsÚ_numeric_suite_fxÚtorch.ao.ns.fx.mappingsr  Úaddr+   rU   Úadd_loggersÚOutputLoggerÚextract_logger_infoÚ%extend_logger_results_with_comparisonÚfxÚutilsÚcompute_sqnr)r  r  r  r$  r  r  Ú
model_a_nsÚ
model_b_nsÚactivation_comparison_dictÚlayer_sqnr_dictÚkeyÚlayerr  s                r   Úget_layer_sqnr_dictr5  ª  s  € ÷  /Ó.ÝAá5Ó7ÐØÐ,Ñ-×1Ñ1´%·)±)Ô<àŸ^™^ØØØØØ
�‰Ø3ð ,ð Ñ€Jñ ˆq„MÙˆq„Mà!#×!7Ñ!7Ø §¡°ó"Ðð ×,Ñ,Ø"ØØÜ�‰�‰�‰×Ñ×)Ñ)Øôð €OÛ)ˆØ*Ñ/°Ñ>¸vÑFÀqÑIÈ%ÑPˆØ)Ñ.¨}Ñ=¸fÑEÀaÑHÈÑPÐQRÑSˆØ!%˜Óñ *ð
 Ðr   r2  Únum_layers_to_equalizec                 ó®   • [        U R                  5       [        R                  " S5      S9nUSU nU Vs/ s H  oDS   [        4PM     nnSU0nU$ s  snf )a§  Given the layer to SQNR dictionary, find the layers with the highest
quantization errors, and return an equalization_qconfig_dict
specifying to only equalize those top layers.

Args:
    layer_sqnr_dict: Dictionary mapping layer names to SQNR values (found
        when comparing an equalized model against a float model)
    num_layers_to_equalize: Number of layers with the highest quantization
       errors to equalize
r   )r3  Nr   Úmodule_name)ÚsortedÚitemsÚoperatorÚ
itemgetterÚdefault_equalization_qconfig)r2  r6  Úlayer_sqnr_sortedÚlayers_to_equalizeÚitemÚmodule_to_qconfig_listÚequalization_qconfig_dicts          r   Úget_equalization_qconfig_dictrC  á  sv   € ô  ˜×4Ñ4Ó6¼H×<OÒ<OÐPQÓ<RÑSÐØ*Ð+BÐ,BÐCÐñ
 =OóÚ<N°Dˆa‰Ô.Ó/Ñ<Nð ð ð "/Ð0FÐ GÐØ$Ð$ùò	s   ¶A)Or;  rR   Úcollectionsr   Útypingr   r+   Útorch.ao.nn.intrinsicr#  rŠ   Ú	intrinsicr•   Útorch.nnÚtorch.nn.functionalÚ
functionalr¬   Ú%torch.ao.quantization.fx.graph_moduler   Útorch.ao.quantization.observerr   r   r   Útorch.ao.quantization.utilsr	   r
   Útorch.fxr   Útorch.fx.graphr   r-  r   r   r   r   rù   Ú__annotations__rÜ   Úintr   r,   Úper_channel_affiner-   Úper_channel_symmetricr/   r‹   r   rk   rƒ   r…   rf   rd   Úinput_equalization_observerrx   Úweight_equalization_observerr=  Úboolrœ   r¢   r¤   r²   rµ   r  rª   ÚtuplerÅ   r¿   rÑ   rÔ   rÙ   rä   rò   rö   r   r  r  r  Úfloatr5  rC  r‡   r   r   Ú<module>rY     s�  ðä Û Ý "Ý ã ß #Ó #Ý ß Ð Ý Q÷ñ ÷
 JÝ  Ý ÷ñ ð &(Ð ˜˜c™Ó 'ð!˜Ÿ™ð !¨Sð !¸¿¹ð !È%Ï,É,ô !ð 
×Ñ˜U×5Ñ5Ø	×Ñ × ;Ñ ;ð,Ð (ôa( §¡ô a(ôH?( "§)¡)ô ?(ðD%Ø)ð%Ø7Rð%à
‡\�\ô%ôPáÐ$Ð'9¸8Ð&DÓEôðD 9×BÑBØ
�,‰, × :Ñ :ð Cð Ð ð  ;×DÑDØ
�+‰+˜u×:Ñ:ð  Eð  Ð ñ  3Ø0Ð9Uñ Ð ð
°$ô ðH¨tô Hð
3°4ô 3ð
 Tð °tô ð r§y¡yð °Tô ð8Øð8Ø$/ð8Ø:>¸sÀBÇIÁI¸~Ñ:Nð8à
ˆ4�$‰;Ð3°dÑ:Ð:Ñ;ô8ðvØðØ   b§i¡i Ñ0ðà	ˆD�[ôð*3Ø
ð3Ø˜c 2§9¡9˜nÑ-ð3à $Ñ&ô3ðlØ
ðØ˜c 2§9¡9˜nÑ-ðà
‡\�\�DÑôð,/˜tð /¨d°3¸¿	¹	°>Ñ.Bð /Àtô /ð8>/Ø
ð>/à�#�r—y‘y�.Ñ!ð>/ð Ÿ™ð>/ð #Ÿ\™\¨DÑ0ð	>/ð
 
ô>/ðBQ?ØðQ?àðQ?ð �#�r—y‘y�.Ñ!ðQ?ð Ÿ™ð	Q?ð
 #Ÿ\™\¨DÑ0ðQ?ð 
ôQ?ðh*¨ð *¸¸SÀ"Ç)Á)¸^Ñ8Lð *ÐQUô *ð,!�{ð !¨$ð !¸4ô !ð0Øð0Ø!% c¨2¯9©9 nÑ!5ð0à	ˆ#Ð*Ð
*Ñ+ô0ðf[Øð[à�#�r—y‘y�.Ñ!ð[ð ˜SÐ"=Ð=Ñ>ð[ð 
ô	[ð|+ [ô +ð&4Ø�Y‰Yð4Ø!#§¡ð4Ø/4¯|©|ð4à	ˆ#ˆuˆ*Ñô4ðn%Ø˜#˜u˜*Ñ%ð%Ø?Bð%àõ%r   