Ë
    g^(h¬Ž  ã                   óÜ  — d dl Z d dlZd dlZd dlZd dlmZmZmZmZ d dl	Z	d dl
mc mZ d dlmZ d dlmZ d dlmZmZmZmZ d dlmZmZmZ d dlmZmZ dd	lmZm Z m!Z!m"Z"m#Z# erd d
l$m%Z% g Z&de'de'de'de'de(e)ef   f
d„Z*dedefd„Z+dedefd„Z,dedefd„Z-d„ Z.de'de'de'dede'defd„Z/de'de'de'dede'defd„Z0dddedede'fd„Z1dddedede'fd„Z2dede'fd „Z3dede'fd!„Z4d"ede(e)e5eef   f   fd#„Z6d$e(eef   de(eef   fd%„Z7d&ed'efd(„Z8d&ed)efd*„Z9d+ed,e(eef   fd-„Z:d.edefd/„Z;d.eded0e5ed1f   de'def
d2„Z<d.efd3„Z=d.efd4„Z>d&ed)efd5„Z?d.edefd6„Z@d.eded0e5ed1f   de'def
d7„ZAy)8é    N)ÚAnyÚCallableÚOptionalÚTYPE_CHECKING)Úquantized_decomposed_lib)Ú_WrapperModule)ÚDerivedQuantizationSpecÚ
EdgeOrNodeÚQuantizationSpecBaseÚSharedQuantizationSpec)ÚGraphÚGraphModuleÚNode)Úreplace_pattern_with_filtersÚReplacedPatternsé   )Ú"_get_aten_graph_module_for_patternÚ_is_bn_nodeÚ_is_conv_or_conv_transpose_nodeÚ_is_conv_transpose_fnÚfold_bn_weights_into_conv_node)ÚInternalMatchÚis_per_channelÚhas_biasÚbias_is_quantizedÚis_cudaÚreturnc                 ó"  — i }| r¨t        j                  dgt         j                  ¬«      |d<   t        j                  dgt         j                  ¬«      |d<   |rT|rRt        j                  dgt         j                  ¬«      |d<   t        j                  dgt         j                  ¬«      |d<   |rt        j                  d«      |d<   |rF|j                  «       D ]3  \  }}t        |t         j                  «      sŒ!|j                  «       ||<   Œ5 |S )	zu
    Optional example inputs for quantized and folded conv-bn patterns
    used in convert, expressed as kwargs.
    r   ©ÚdtypeÚweight_scaler   Úweight_zero_pointÚ
bias_scaleÚbias_zero_pointÚ	conv_bias)	ÚtorchÚtensorÚfloatÚintÚrandnÚitemsÚ
isinstanceÚTensorÚcuda)r   r   r   r   ÚkwargsÚkÚvs          úb/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/ao/quantization/pt2e/qat_utils.pyÚ,_get_quantized_conv_bn_example_inputs_kwargsr3   $   s×   € ð €Fñ Ü!&§¡¨q¨c¼¿¹Ô!Eˆˆ~ÑÜ&+§l¡l°A°3¼e¿i¹iÔ&HˆÐ"Ñ#ÙÑ)Ü#(§<¡<°°¼5¿;¹;Ô#GˆF�<Ñ Ü(-¯©°a°SÄÇ	Á	Ô(JˆFÐ$Ñ%ÙÜ#Ÿk™k¨!›nˆˆ{ÑÙØ—L‘L“Nò 	%‰DˆAˆqÜ˜!œUŸ\™\Õ*ØŸF™F›H��q’	ð	%ð €Mó    Úconv_fnc                 ó&  ‡ — dt         j                  dt         j                  dt         j                  dt         j                  dt         j                  dt         j                  dt         j                  dt         j                  fˆ fd	„}t        |«      S )
NÚxÚconv_weightr%   Ú	bn_weightÚbn_biasÚbn_running_meanÚbn_running_varr   c                 óR   •—  ‰| ||«      } t        j                  | ||||d¬«      } | S )NT)Útraining)ÚFÚ
batch_norm)r7   r8   r%   r9   r:   r;   r<   r5   s          €r2   Ú_conv_bn_patternz._get_conv_bn_pattern.<locals>._conv_bn_patternA   s5   ø€ ñ �A�{ IÓ.ˆÜ�L‰LØˆ °	¸7ÈTô
ˆð ˆr4   ©r&   r-   r   )r5   rA   s   ` r2   Ú_get_conv_bn_patternrC   @   s‚   ø€ ðÜ�<‰<ðä—\‘\ðô —<‘<ðô —<‘<ð	ô
 —‘ðô Ÿ™ðô Ÿ™ðô 
�‰õô Ð*Ó+Ð+r4   c                 ó&  ‡ — dt         j                  dt         j                  dt         j                  dt         j                  dt         j                  dt         j                  dt         j                  dt         j                  fˆ fd	„}t        |«      S )
Nr7   r8   r%   r9   r:   r;   r<   r   c           	      óâ  •— d}t        j                  ||z   «      }||z  }	dgt        |j                  «      z  }
t	        ‰«      rdnd}d|
|<   dgt        |j                  «      z  }d|d<   ||	j                  |
«      z  }t        j                  || j                  ¬«      } ‰| ||«      } | |	j                  |«      z  } | |j                  |«      z   } t        j                  | ||||d|¬«      } | S )zô
        Approximated method to fuse conv and bn. It requires only one forward pass.
        conv_orig = conv / scale_factor where scale_factor = bn.weight / running_std.
        This is based on `nniqat.ConvBn2d._forward_approximate`.
        çñhãˆµøä>r   r   éÿÿÿÿr   T©r>   Úeps)
r&   ÚsqrtÚlenÚshaper   ÚreshapeÚ
zeros_liker    r?   r@   )r7   r8   r%   r9   r:   r;   r<   Úbn_epsÚrunning_stdÚscale_factorÚweight_shapeÚweight_in_channel_axisÚ
bias_shapeÚscaled_weightÚ	zero_biasr5   s                  €r2   Ú_qat_conv_bn_patternz6_get_qat_conv_bn_pattern.<locals>._qat_conv_bn_patternU   s  ø€ ð ˆÜ—j‘j °&Ñ!8Ó9ˆØ  ;Ñ.ˆØ�sœS ×!2Ñ!2Ó3Ñ3ˆÜ&;¸GÔ&D¡È!ÐØ/1ˆÐ+Ñ,Ø�Sœ3˜{×0Ñ0Ó1Ñ1ˆ
Øˆ
�1‰Ø# l×&:Ñ&:¸<Ó&HÑHˆÜ×$Ñ$ Y°a·g±gÔ>ˆ	Ù�A�} iÓ0ˆØ�×$Ñ$ ZÓ0Ñ0ˆØ�	×!Ñ! *Ó-Ñ-ˆÜ�L‰LØØØØØØØô
ˆð ˆr4   rB   )r5   rW   s   ` r2   Ú_get_qat_conv_bn_patternrX   T   sƒ   ø€ ð%Ü�<‰<ð%ä—\‘\ð%ô —<‘<ð%ô —<‘<ð	%ô
 —‘ð%ô Ÿ™ð%ô Ÿ™ð%ô 
�‰õ%ôN Ð.Ó/Ð/r4   c                 ó&  ‡ — dt         j                  dt         j                  dt         j                  dt         j                  dt         j                  dt         j                  dt         j                  dt         j                  fˆ fd	„}t        |«      S )
Nr7   r8   r%   r9   r:   r;   r<   r   c           	      óx  •— d}t        j                  ||z   «      }||z  }	dgt        |j                  «      z  }
t	        ‰«      rdnd}d|
|<   dgt        |j                  «      z  }d|d<   ||	j                  |
«      z  } ‰| |d«      } | |	j                  |«      z  } t        j                  | ||||d|¬«      } | S )z]
        Same as `_get_qat_conv_bn_pattern`, but handles the case with no conv bias.
        rF   r   r   rG   NTrH   )r&   rJ   rK   rL   r   rM   r?   r@   )r7   r8   r%   r9   r:   r;   r<   rO   rP   rQ   rR   rS   rT   rU   r5   s                 €r2   Ú!_qat_conv_bn_pattern_no_conv_biaszP_get_qat_conv_bn_pattern_no_conv_bias.<locals>._qat_conv_bn_pattern_no_conv_bias€   sÚ   ø€ ð ˆÜ—j‘j °&Ñ!8Ó9ˆØ  ;Ñ.ˆØ�sœS ×!2Ñ!2Ó3Ñ3ˆÜ&;¸GÔ&D¡È!ÐØ/1ˆÐ+Ñ,Ø�Sœ3˜{×0Ñ0Ó1Ñ1ˆ
Øˆ
�1‰Ø# l×&:Ñ&:¸<Ó&HÑHˆÙ�A�} dÓ+ˆØ�×$Ñ$ ZÓ0Ñ0ˆÜ�L‰LØØØØØØØô
ˆð ˆr4   rB   )r5   r[   s   ` r2   Ú%_get_qat_conv_bn_pattern_no_conv_biasr\      sƒ   ø€ ð"Ü�<‰<ð"ä—\‘\ð"ô —<‘<ð	"ô
 —<‘<ð"ô —‘ð"ô Ÿ™ð"ô Ÿ™ð"ô 
�‰õ"ôH Ð;Ó<Ð<r4   c           	      ó^  — d}|rdnd}|rdnd}|r||   nd}|r||   nd}d}	d}
t         j                  }t         j                  j                  }|r0|j	                  | ||||	|
|«      } |j                  | ||||	|
|«      } | S |j                  | |||	|
|«      } |j                  | |||	|
|«      } | S )	a  
    Helper function to append q-dq ops after `x`, using dummy values for the qparams
    and qmin/qmax. We use dummy values here because we match with `ignore_literals=True`
    and will manually replace these values after subgraph rewriting.

    Return the dq node.
    r   r#   r!   r$   r"   g      ð?i�ÿÿÿé   )r&   Úint8ÚopsÚquantized_decomposedÚquantize_per_channelÚdequantize_per_channelÚquantize_per_tensorÚdequantize_per_tensor)r7   r   Úis_biasr/   Úper_channel_axisÚ	scale_keyÚzp_keyÚscaleÚzpÚqminÚqmaxr    Úqds                r2   Ú_append_qdqro   §   sÜ   € ð ÐÙ '‘¨^€IÙ")ÑÐ/B€FÙ!/ˆF�9Ò°S€EÙ)ˆ�Š¨q€BØ€DØ€DÜ�J‰J€Eä	�‰×	'Ñ	'€BÙØ×#Ñ# A u¨bÐ2BÀDÈ$ÐPUÓVˆØ×%Ñ% a¨°Ð4DÀdÈDÐRWÓXˆð €Hð ×"Ñ" 1 e¨R°°t¸UÓCˆØ×$Ñ$ Q¨¨r°4¸¸uÓEˆØ€Hr4   Úbn_is_trainingc                 ó  ‡ ‡‡‡‡‡— dŠdt         j                  dt         j                  dt         j                  dt         j                  dt         j                  dt         j                  dt         j                  fˆˆˆˆˆˆ fd	„}t        |«      S )
a  
    Return the quantized version of QAT conv + BN pattern.
    This is based on `nniqat.ConvBn2d._forward_approximate`,
    used in QAT convert. We first match this pattern and replace
    it with the normal [conv - bn] pattern, then fold the BN
    weights into conv.
    rF   r7   r8   r9   r:   r;   r<   r   c           	      ó*  •— t        j                  |‰z   «      }||z  }dgt        |j                  «      z  }	d|	d<   dgt        |j                  «      z  }
d|
d<   ||j	                  |	«      z  }t        |‰d|¬«      }‰r@t        j                  |d   | j                  ¬«      }‰rt        |‰d|¬«      } ‰| ||«      } n
 ‰| |d «      } | |j	                  |
«      z  } ‰r| |d   j	                  |
«      z   } t        j                  | ||||‰‰¬	«      } | S )
Nr   rG   r   F©rf   r/   r%   r   TrH   )
r&   rJ   rK   rL   rM   ro   rN   r    r?   r@   )r7   r8   r9   r:   r;   r<   r/   rP   rQ   rR   rT   rU   rV   r   rO   rp   r5   r   r   s                €€€€€€r2   Ú_quantized_qat_conv_bn_patternzJ_get_quantized_qat_conv_bn_pattern.<locals>._quantized_qat_conv_bn_patternÔ   s@  ø€ ô —j‘j °&Ñ!8Ó9ˆØ  ;Ñ.ˆØ�sœS ×!2Ñ!2Ó3Ñ3ˆØˆ�Q‰Ø�Sœ3˜{×0Ñ0Ó1Ñ1ˆ
Øˆ
�1‰Ø# l×&:Ñ&:¸<Ó&HÑHˆÜ#ØØØØô	
ˆñ Ü×(Ñ(¨°Ñ)<ÀAÇGÁGÔLˆIÙ Ü'ØØ"Ø Ø!ô	�	ñ ˜˜=¨)Ó4‰Aá˜˜=¨$Ó/ˆAØ�×$Ñ$ ZÓ0Ñ0ˆÙØ�F˜;Ñ'×/Ñ/°
Ó;Ñ;ˆAÜ�L‰LØØØØØØ#Øô
ˆð ˆr4   rB   )r   r   r   r5   rp   rt   rO   s   ````` @r2   Ú"_get_quantized_qat_conv_bn_patternru   Ã   s�   ý€ ð €Fð.Ü�<‰<ð.ä—\‘\ð.ô —<‘<ð.ô —‘ð	.ô
 Ÿ™ð.ô Ÿ™ð.ô 
�‰÷.ò .ô` Ð8Ó9Ð9r4   c                 ó  ‡ ‡‡‡‡‡— dŠdt         j                  dt         j                  dt         j                  dt         j                  dt         j                  dt         j                  dt         j                  fˆˆˆˆˆˆ fd	„}t        |«      S )
zQ
    Quantized QAT conv - bn pattern with bn weights being folded into conv.
    rF   r7   r8   r9   r:   r;   r<   r   c           	      ó¨   •— t        |‰d|¬«      }‰r|d   }‰rt        |‰d|¬«      }nd } ‰| ||«      } t        j                  | ||||‰
‰	¬«      } | S )NFrs   r%   TrH   )ro   r?   r@   )r7   r8   r9   r:   r;   r<   r/   Úbiasr   rO   rp   r5   r   r   s           €€€€€€r2   Ú%_folded_quantized_qat_conv_bn_patternzX_get_folded_quantized_qat_conv_bn_pattern.<locals>._folded_quantized_qat_conv_bn_pattern  s‚   ø€ ô "ØØØØô	
ˆñ Ø˜+Ñ&ˆDÙ Ü"ØØ"Ø Ø!ô	‘ð ˆDÙ�A�{ DÓ)ˆÜ�L‰LØØØØØØ#Øô
ˆð ˆr4   rB   )r   r   r   r5   rp   ry   rO   s   ````` @r2   Ú)_get_folded_quantized_qat_conv_bn_patternrz     s�   ý€ ð €Fð$Ü�<‰<ð$ä—\‘\ð$ô —<‘<ð$ô —‘ð	$ô
 Ÿ™ð$ô Ÿ™ð$ô 
�‰÷$ò $ôL Ð?Ó@Ð@r4   Úmatchr   Úoriginal_graphÚpattern_graphc                 óÆ   — | j                   j                  «       D ]:  }t        |«      sŒt        |j                  «      dkD  xr |j                  d   duc S  t        d«      ‚)zw
    Match filter for the subgraph rewriter that returns True if the conv node in
    the original graph has bias.
    é   Nz5Could not find conv node in matched conv + bn pattern)Ú	nodes_mapÚvaluesr   rK   ÚargsÚ
ValueError)r{   r|   r}   Úns       r2   Ú_has_conv_bias_filterr…   =  sZ   € ð �_‰_×#Ñ#Ó%ò =ˆÜ*¨1Õ-Ü�q—v‘v“; ‘?Ò< q§v¡v¨a¡y¸Ð'<Ò<ð=ô ÐLÓ
MÐMr4   c                 ó   — t        | ||«       S )z�
    Match filter for the subgraph rewriter that returns True if the conv node in
    the original graph does NOT have bias.
    )r…   )r{   r|   r}   s      r2   Ú_no_conv_bias_filterr‡   L  s   € ô % U¨N¸MÓJÐJÐJr4   r„   c                 ó,  — | j                   t        j                  j                  j                  j
                  t        j                  j                  j                  j                  t        j                  j                  j                  j
                  fv S ©N)Útargetr&   r`   ra   rd   Údefaultr'   rb   ©r„   s    r2   Ú_is_quantizer�   X  sc   € Ø�8‰8Ü�	‰	×&Ñ&×:Ñ:×BÑBÜ�	‰	×&Ñ&×:Ñ:×AÑAÜ�	‰	×&Ñ&×;Ñ;×CÑCðð ð r4   c                 ó,  — | j                   t        j                  j                  j                  j
                  t        j                  j                  j                  j                  t        j                  j                  j                  j
                  fv S r‰   )rŠ   r&   r`   ra   re   r‹   r'   rc   rŒ   s    r2   Ú_is_dequantizer�   `  sc   € Ø�8‰8Ü�	‰	×&Ñ&×<Ñ<×DÑDÜ�	‰	×&Ñ&×<Ñ<×CÑCÜ�	‰	×&Ñ&×=Ñ=×EÑEðð ð r4   Úrc                 ón  — dt         t           dt        t        t        t        t           f   fd„}dt        dt        t        t        t        f   fd„}t        t	        | j
                  «      j                  «       «      } ||«      \  }}} || j                  «      \  }}}	|�J ‚|	�J ‚||f||fdœ}
 |t        | j
                  j                  «       «      «      \  }}}|j                  ^}}}|j                  ^}}}t        |t        «      sJ ‚t        |t        «      sJ ‚t        |t        «      sJ ‚t        |t        «      sJ ‚| j
                  |   }| j
                  |   }t        |«      rS ||«      \  }}} ||«      \  }}}| j
                  |   }| j
                  |   }| j
                  |   }||f|
d<   ||f|
d<   ||f|
d	<   ||f|
d
<   t        |j                  «      dkD  rÎt        |j                  «      dkD  r¶|j                  d   }|j                  d   }t        |t        «      sJ ‚t        |t        «      sJ ‚| j
                  |   }t        |«      rS ||«      \  }}} ||«      \  }}}| j
                  |   }| j
                  |   } | j
                  |   }!| |f|
d<   |!|f|
d<   ||f|
d<   |
S )a¯  
    Helper function to extract the nodes in the conv-bn fusion pattern after
    subgraph rewriting, in the form of a map:

        {name: (original_node, replacement_node)}

    The following names must exist in the map:

        "conv", "conv_weight", "conv_input", "bn", "getitem"

    The following names may exist in the map:

        "conv_weight_q", "conv_weight_dq", "conv_bias",
        "conv_bias_q", "conv_bias_dq"
    Únodesr   c                 óâ   — d\  }}}| D ]X  }|j                   dk7  rŒt        |«      r|�J ‚|}t        |«      r|�J ‚|}|j                  t        j
                  k(  sŒS|�J ‚|}ŒZ |€J ‚|€J ‚|||fS )z�
        Return a 3-tuple of (conv_node, bn_node, getitem_node).
        This asserts that the match contains exactly one of each node.
        )NNNÚcall_function)Úopr   r   rŠ   ÚoperatorÚgetitem)r’   Ú	conv_nodeÚbn_nodeÚgetitem_noder„   s        r2   Ú
_get_nodesz._get_conv_bn_pattern_nodes.<locals>._get_nodesy  s§   € ð
 ,<Ñ(ˆ	�7˜LØò 	!ˆAØ�t‰t�Ò&ØÜ.¨qÔ1Ø Ð(Ð(Ð(Ø�	Ü˜1Œ~Ø�Ð&�Ø�Ø�x‰xœ8×+Ñ+Ó+Ø#Ð+Ð+Ð+Ø ‘ð	!ð Ð$Ð$Ð$ØÐ"Ð"Ð"Ø˜7 LÐ1Ð1r4   r„   c                 óÄ   — t        | «      sJ ‚| j                  d   }t        |t        «      sJ ‚t	        |«      sJ ‚|j                  d   }t        |t        «      sJ ‚||| fS )zC
        Return a 3-tuple of (orig_node, q_node, dq_node).
        r   )r�   r‚   r,   r   r�   )r„   Úq_nodeÚ	orig_nodes      r2   Ú_get_q_dq_nodesz3_get_conv_bn_pattern_nodes.<locals>._get_q_dq_nodes�  sg   € ô ˜aÔ Ð Ð Ø—‘˜‘ˆÜ˜&¤$Ô'Ð'Ð'Ü˜FÔ#Ð#Ð#Ø—K‘K ‘Nˆ	Ü˜)¤TÔ*Ð*Ð*Ø˜6 1Ð%Ð%r4   )ÚconvÚbnÚconv_weight_qÚconv_weight_dqÚ
conv_inputr8   r   Úconv_bias_qÚconv_bias_dqr%   )Úlistr   Útupler   Ú_filter_nodes_mapr€   r�   ÚreplacementsÚkeysr‚   r,   r�   rK   )"r�   r›   rŸ   Úoriginal_nodesÚo_convÚo_bnÚ	o_getitemÚr_convÚr_bnÚ	r_getitemÚmappingÚp_convÚ_Úp_conv_inputÚp_conv_weightÚr_conv_inputÚr_conv_weightÚo_conv_inputÚo_conv_weightÚp_conv_weight_qÚp_conv_weight_dqÚr_conv_weight_qÚr_conv_weight_dqÚo_conv_weight_qÚo_conv_weight_dqÚp_conv_biasÚr_conv_biasÚo_conv_biasÚp_conv_bias_qÚp_conv_bias_dqÚr_conv_bias_qÚr_conv_bias_dqÚo_conv_bias_qÚo_conv_bias_dqs"                                     r2   Ú_get_conv_bn_pattern_nodesrË   h  só  € ð"2œ$œt™*ð 2¬¬t´T¼8ÄD¹>Ð/IÑ)Jó 2ð,
&œ4ð 
&¤E¬$´´dÐ*:Ñ$;ó 
&ô Ô+¨A¯K©KÓ8×?Ñ?ÓAÓB€NÙ(¨Ó8Ñ€FˆD�)Ù(¨¯©Ó8Ñ€FˆD�)ð ÐÐÐØÐÐÐà˜Ð Ø�Tˆlñ€Gñ  ¤ Q§[¡[×%5Ñ%5Ó%7Ó 8Ó9�N€VˆQ�Ø(.¯©Ð%€\�= 1Ø(.¯©Ð%€\�= 1Ü�l¤DÔ)Ð)Ð)Ü�m¤TÔ*Ð*Ð*Ü�l¤DÔ)Ð)Ð)Ü�m¤TÔ*Ð*Ð*Ø—;‘;˜|Ñ,€LØ—K‘K Ñ.€Mô �mÔ$Ù;JØó<
Ñ8ˆ�Ð(8ñ <KØó<
Ñ8ˆ�Ð(8ð Ÿ™ MÑ2ˆØŸ+™+ oÑ6ˆØŸ;™;Ð'7Ñ8ÐØ$3°_Ð#Eˆ�Ñ Ø%5Ð7GÐ$HˆÐ Ñ!Ø)¨<Ð8€GˆLÑØ+¨]Ð;€GˆMÑô ˆ6�;‰;Ó˜!Ò¤ F§K¡KÓ 0°1Ò 4Ø—k‘k !‘nˆØ—k‘k !‘nˆÜ˜+¤tÔ,Ð,Ð,Ü˜+¤tÔ,Ð,Ð,Ø—k‘k +Ñ.ˆô ˜+Ô&Ù9HÈÓ9UÑ6ˆK˜¨Ù9HÈÓ9UÑ6ˆK˜¨ØŸ+™+ kÑ2ˆKØŸK™K¨Ñ6ˆMØŸ[™[¨Ñ8ˆNØ&3°]Ð%CˆG�MÑ"Ø'5°~Ð&FˆG�NÑ#Ø +¨[Ð9ˆ�ÑØ€Nr4   r€   c                 ój   — i }| j                  «       D ]  \  }}|€Œ	|j                  dk(  rŒ|||<   Œ |S )zÔ
    Return a filtered `nodes_map` returned from the subgraph rewriter.
    The filtered `nodes_map` will contain only nodes that are actually
    matched in the pattern, excluding None or placeholder nodes.
    Úplaceholder)r+   r•   )r€   Únew_nodes_mapÚpattern_nodeÚ
graph_nodes       r2   r©   r©   Ù  sN   € ð ')€MØ$-§O¡OÓ$5ò 1Ñ ˆ�jàÐØà�?‰?˜mÒ+ØØ&0ˆ�lÒ#ð1ð Ðr4   Úoriginal_nodeÚnew_nodec                 óæ   — t        | «      sJ ‚t        |«      sJ ‚t        |j                  «      }t        |«      dk  r|j	                  d«       t        |dd «      | j                  dd z   |_        y)a=  
    Copy over literal args in conv, such as stride and padding, from the matched node
    in the original graph to its replacement in the new graph.

    This is needed due to the following limitation in the subgraph rewriter when used
    with dynamo export: literal (non-tensor) args are not supported in the match and
    replacement patterns. This is because dynamo export automatically inlines these
    literal args, making them dead placeholder nodes. In the future, we should check
    if dynamo export can optionally disable this inlining, or if subgraph rewriter
    can do the copying for us. See https://github.com/pytorch/pytorch/issues/100419.

    Note: Unlike other tensor args like conv weights and biases, literal args are
    preserved in the original nodes after replacement, so we can access them here.
    é   N)r   r§   r‚   rK   Úappendr¨   )rÑ   rÒ   Únew_argss      r2   Ú_copy_over_literal_conv_argsr×   ì  sj   € ô +¨=Ô9Ð9Ð9Ü*¨8Ô4Ð4Ð4ä�H—M‘MÓ"€HÜ
ˆ8ƒ}�qÒà�‰˜ÔÜ˜( 2 A˜,Ó'¨-×*<Ñ*<¸Q¸RÐ*@Ñ@€H…Mr4   Úreplacement_nodec                 óÂ  — t        | «      sJ ‚t        |«      sJ ‚d| j                  vry| j                  d   j                  }i }t        |j	                  «       «      }|d   d   ||j
                  d   <   |d   d   ||j
                  d   <   t        |j
                  «      dkD  r&t        |«      dkD  r|d   d   ||j
                  d   <   ||j                  d   _        y)a  
    Update the `input_qspec_map` in the annotation after subgraph rewriting.

    The original annotation referred to the nodes in the original graph,
    so the keys in the `input_qspec_map` will need to be updated to reflect
    the corresponding nodes in the replacement graph.
    Úquantization_annotationNr   r   r   )r   ÚmetaÚinput_qspec_mapr§   r+   r‚   rK   )rÑ   rØ   Úoriginal_input_qspec_maprÜ   Úall_configss        r2   Ú._update_conv_input_qspec_map_after_replacementrß     sü   € ô +¨=Ô9Ð9Ð9Ü*Ð+;Ô<Ð<Ð<Ø ¨×(:Ñ(:Ñ:ØØ,×1Ñ1Ø!ñ ç�oð ð €Oô Ð/×5Ñ5Ó7Ó8€Kà0;¸A±¸qÑ0A€OÐ$×)Ñ)¨!Ñ,Ñ-à0;¸A±¸qÑ0A€OÐ$×)Ñ)¨!Ñ,Ñ-ä
Ð× Ñ Ó! AÒ%¬#¨kÓ*:¸QÒ*>Ø4?À±NÀ1Ñ4EˆÐ(×-Ñ-¨aÑ0Ñ1ØGVÐ×ÑÐ3Ñ4ÕDr4   ÚnodeÚoriginal_to_replacement_nodec                 ó  ‡‡— dt         fˆfd„Šdt        fˆfd„}d| j                  vry| j                  d   }|j                  j	                  «       D ]  \  }} ||«      |j                  |<   Œ  ||j
                  «      |_        y)ai  
    Update the `SharedQuantizationSpec`s and `DerivedQuantizationSpec`s
    used in `node`'s quantization annotation after subgraph rewriting.

    The original annotation referred to the nodes in the original graph,
    so the nodes used in these special quantization specs will need to
    be updated to the corresponding nodes in the replacement graph.
    Úedge_or_nodec                 ó(  •— t        | t        «      r| }‰j                  ||«      S t        | t        «      rIt	        | «      dk(  r;t        d„ | D «       «      r)| \  }}‰j                  ||«      ‰j                  ||«      fS t        dt        | «      «      ‚)Nr   c              3   ó<   K  — | ]  }t        |t        «      –— Œ y ­wr‰   )r,   r   )Ú.0r7   s     r2   ú	<genexpr>zZ_update_special_qspecs_after_replacement.<locals>._get_new_edge_or_node.<locals>.<genexpr>9  s   è ø€ Ò>¨A”J˜q¤$×'Ñ>ùs   ‚z"unexpected type for edge_or_node: )r,   r   Úgetr¨   rK   Úallrƒ   Útype)rã   Ú_nodeÚsrcÚdestrá   s       €r2   Ú_get_new_edge_or_nodezG_update_special_qspecs_after_replacement.<locals>._get_new_edge_or_node2  s’   ø€ Ü�l¤DÔ)Ø ˆEØ/×3Ñ3°E¸5ÓAÐAä�|¤UÔ+Ü�LÓ! QÒ&ÜÑ>°Ô>Ô>à$‰IˆC�à,×0Ñ0°°cÓ:Ø,×0Ñ0°°tÓ<ðð ô
 ÐAÄ4ÈÓCUÓVÐVr4   Úqspecc                 óø   •— t        | t        «      r ‰| j                  «      }t        |«      S t        | t        «      r6| j                  D �cg c]
  } ‰|«      ‘Œ }}t        j                  | |¬«      S | S c c}w )N)Úderived_from)r,   r   rã   r	   rñ   ÚdataclassesÚreplace)rï   Únew_edge_or_noder7   Únew_derived_fromrî   s       €r2   Ú_get_new_qspecz@_update_special_qspecs_after_replacement.<locals>._get_new_qspecC  su   ø€ Ü�eÔ3Ô4Ù4°U×5GÑ5GÓHÐÜ)Ð*:Ó;Ð;Ü˜Ô6Ô7ØBG×BTÑBTÖU¸QÑ 5°aÕ 8ÐUÐÐUÜ×&Ñ& uÐ;KÔLÐLàˆLùò  Vs   ÁA7rÚ   N)r
   r   rÛ   rÜ   r+   Úoutput_qspec)rà   rá   rö   Ú
annotationÚ
input_noderï   rî   s    `    @r2   Ú(_update_special_qspecs_after_replacementrú   %  s‹   ù€ ðW¬Jõ Wð"Ô2õ ð !¨¯	©	Ñ1ØØ—‘Ð4Ñ5€JØ'×7Ñ7×=Ñ=Ó?ò GÑˆ
�EÙ1?ÀÓ1Fˆ
×"Ñ" :Ò.ðGá,¨Z×-DÑ-DÓE€JÕr4   Úmc           	      óæ  — t        j                  ddd«      t        j                  ddd«      t        j                  d«      t        j                  d«      t        j                  d«      t        j                  d«      t        j                  d«      f}t        j                  dddd«      t        j                  dddd«      t        j                  d«      t        j                  d«      t        j                  d«      t        j                  d«      t        j                  d«      f}t        d„ | j                  j                  D «       «      }|s| S t         j
                  j                  «       rddgndg}|D ]v  }t        | t        j                  ||¬«      } t        | t        j                  ||¬«      } t        | t        j                  ||¬«      } t        | t        j                  ||¬«      } Œx | S )Nr   rÔ   c              3   ó2   K  — | ]  }t        |«      –— Œ y ­wr‰   ©r   ©ræ   r„   s     r2   rç   z$_fuse_conv_bn_qat.<locals>.<genexpr>l  ó   è ø€ Ò7 A”˜Q—Ñ7ùó   ‚TF©r   )r&   r*   ÚanyÚgraphr’   r.   Úis_availableÚ_fuse_conv_bn_qat_helperr?   Úconv1dÚconv2dÚconv_transpose1dÚconv_transpose2d)rû   Ú_conv1d_bn_example_inputsÚ_conv2d_bn_example_inputsÚhas_bnÚis_cuda_optionsr   s         r2   Ú_fuse_conv_bn_qatr  U  s†  € ô 	�‰�A�q˜!ÓÜ�‰�A�q˜!ÓÜ�‰�A‹Ü�‰�A‹Ü�‰�A‹Ü�‰�A‹Ü�‰�A‹ð!Ðô 	�‰�A�q˜!˜QÓÜ�‰�A�q˜!˜QÓÜ�‰�A‹Ü�‰�A‹Ü�‰�A‹Ü�‰�A‹Ü�‰�A‹ð!Ðô Ñ7¨¯©¯©Ô7Ó7€FÙØˆÜ',§z¡z×'>Ñ'>Ô'@�t˜U‘mÀuÀg€OØ"ò 
ˆÜ$ØŒq�x‰xÐ2¸Gô
ˆô %ØŒq�x‰xÐ2¸Gô
ˆô %ØŒq×!Ñ!Ð#<Àgô
ˆô %ØŒq×!Ñ!Ð#<Àgô
‰ð
ð €Hr4   Úexample_inputs.c                 óV  — | j                   j                  «        | j                  «        t        |«      }t	        |||«      }t        |«      }t	        |||«      }t        | ||t        gd¬«      }| j                  «        t        |«      }	t	        |	||«      }
t        | ||
t        gd¬«      }| j                  «        i }||z   D ]»  }t        |«      }|d   d   j                  j                  dd«      }|j                  «       D ]y  \  }}|\  }}|j                  |_        |dv r2|r0d|j                  vr"t        j                  |«      |j                  d<   t!        |«      rt#        ||«       t%        ||«       |||<   Œ{ Œ½ | j                   j&                  D ]  }t)        ||«       Œ | S )aU  
    Given a graph of decomposed aten ops, replace the (conv + bn) pattern with
    the fused QAT subgraph equivalent. The input graph should already be annotated.
    The annotations in the original nodes will be preserved in the corresponding
    nodes in the new subgraph.

    Note: This also handles the (conv + bn + relu) pattern.
    T)Úmatch_filtersÚignore_literalsr    r   Únn_module_stackN)r¤   r8   )r  Úeliminate_dead_codeÚ	recompilerC   r   rX   r   r…   r\   r‡   rË   rÛ   rè   r+   ÚcopyÚdeepcopyr   r×   rß   r’   rú   )rû   r5   r  r   Úconv_bn_patternÚmatch_patternÚqat_conv_bn_patternÚ"replacement_pattern_with_conv_biasÚreplacements_with_conv_biasÚ qat_conv_bn_pattern_no_conv_biasÚ replacement_pattern_no_conv_biasÚreplacements_no_conv_biasÚ!all_original_to_replacement_nodesr�   Úreplacement_dictÚconv_nn_moduler0   Ú
node_tuplerÑ   rØ   r„   s                        r2   r  r  €  sì  € ð ‡G�G×ÑÔ!Ø‡K�K„Mä*¨7Ó3€OÜ6ØØØó€Mô 3°7Ó;ÐÜ)KØØØó*Ð&ô
 #?Ø	ØØ*Ü,Ð-Øô#Ðð ‡K�K„Mô (MÈWÓ'UÐ$Ü'IØ(ØØó(Ð$ô
 !=Ø	ØØ(Ü+Ð,Øô!Ðð ‡K�K„Mð( )+Ð%Ø(Ð+DÑDò PˆÜ5°aÓ8Ðà)¨&Ñ1°!Ñ4×9Ñ9×=Ñ=Ð>OÐQUÓVˆØ-×3Ñ3Ó5ò 	P‰MˆAˆzØ.8Ñ+ˆMÐ+à$1×$6Ñ$6ÐÔ!ð Ð2Ñ2Ù"Ø%Ð-=×-BÑ-BÑBä;?¿=¹=ÈÓ;XÐ ×%Ñ%Ð&7Ñ8Ü.¨}Ô=ä,¨]Ð<LÔMä>Ø!Ð#3ôð @PÐ-¨mÒ<ñ'	Pð	Pð4 �W‰W�]‰]ò WˆÜ0°Ð4UÕVðWð €Hr4   c           	      ób  — t         j                  j                  j                  }| j                  j
                  D ]Ö  }|j                  dk7  s'|j                  |k7  st        |j                  «      dk(  rŒ:t        |j                  «      D ]j  }| j                  j                  |«      5  | j                  j                  d||j                  |j                  «      }ddd«       |j                  |«       Œl | j                  j!                  |«       ŒØ | j#                  «        y# 1 sw Y   ŒKxY w)a{  
    Helper function to duplicate all dequantize nodes in the graph if the
    node has more than one user. For example:

    Before:
      quantize -> dequantize -> a
                          \--> b
                          \--> c

    After:
      quantize -> dequantize_1 -> a
            \--> dequantize_2 -> b
            \--> dequantize_3 -> c

    This is useful for subgraph rewriting. E.g. if we wish to match the
    pattern [dequantize - a] above, subgraph matching would fail because
    the dequantize node has users outside the matched portion of the graph.
    Instead, we match [dequantize_1 - a], which is safe.
    r”   r   N)r&   r`   ra   re   r  r’   r•   rŠ   rK   Úusersr§   Úinserting_beforeÚcreate_noder‚   r/   Úreplace_input_withÚ
erase_noder  )rû   Údq_opr„   ÚuserrÒ   s        r2   Ú_duplicate_dequantize_noder-  ñ  sì   € ô( �I‰I×*Ñ*×@Ñ@€EØ�W‰W�]‰]ò ˆØ�4‰4�?Ò" a§h¡h°%Ò&7¼3¸q¿w¹w»<È1Ò;LØÜ˜Ÿ™“Mò 	1ˆDØ—‘×)Ñ)¨!Ó,ñ YØŸ7™7×.Ñ.¨ÀÀqÇvÁvÈqÏxÉxÓX�÷Yà×#Ñ# A xÕ0ð	1ð 	
�‰×Ñ˜1Õðð ‡K�K…M÷	Yð Yús   Â(3D%Ä%D.c                 óZ  — t         j                  j                  j                  }| j                  j
                  D ]Í  }|j                  D �cg c]"  }|j                  dk(  r|j                  |k(  r|‘Œ$ }}t        |«      dkD  sŒI| j                  j                  |d   «      5  | j                  j                  d||d   j                  i «      }ddd«       |D ].  }|j                  «       | j                  j                  |«       Œ0 ŒÏ | j                  «        yc c}w # 1 sw Y   ŒTxY w)a  
    Removes duplicate dequant nodes in the graph, for an operator that has
    multiple dequant nodes as a user, replace them with a single dequant node
    that can be shared across all the uses. This should be seen as the "reverse"
    of `_duplicate_dequantize_node`.
    r”   r   r   N)r&   r`   ra   re   r  r’   r&  r•   rŠ   rK   Úinserting_afterr(  r‚   Úreplace_all_uses_withr*  r  )rû   r+  r„   r,  Údq_usersrÒ   Údq_users          r2   Ú_remove_extra_dequantizer3    s  € ô �I‰I×*Ñ*×@Ñ@€EØ�W‰W�]‰]ò ,ˆð Ÿ™ö
àØ�w‰w˜/Ò)¨d¯k©k¸UÒ.Bò ð
ˆð 
ô
 ˆx‹=˜1ÓØ—‘×(Ñ(¨°!©Ó5ñ ØŸ7™7×.Ñ.Ø# U¨H°Q©K×,<Ñ,<¸bó�÷ð $ò ,�Ø×-Ñ-¨hÔ7Ø—‘×"Ñ" 7Õ+ñ,ð,ð ‡K�K…Mùò
÷ð ús   Á'DÂ",D!Ä!D*	c                 ó`  — | j                   |j                   k(  sJ ‚| j                   t        j                  j                  j                  j
                  t        j                  j                  j                  j
                  fv rd}n„| j                   t        j                  j                  j                  j
                  t        j                  j                  j                  j
                  fv rd}nt        d| j                   › d�«      ‚|j                  d| | j                  |d z   |_
        y)z†
    Given a pair of quantize or dequantize nodes, copy over all literal args
    from the original node to the replacement node.
    r   rÔ   z)Expected quantize/dequantize nodes, got 'ú'N)rŠ   r&   r`   ra   rd   r‹   re   rb   rc   rƒ   r‚   )rÑ   rØ   Ústart_copy_arg_indexs      r2   Ú_copy_over_q_dq_argsr7  *  s  € ð ×ÑÐ#3×#:Ñ#:Ò:Ð:Ð:Ø×ÑÜ�	‰	×&Ñ&×:Ñ:×BÑBÜ�	‰	×&Ñ&×<Ñ<×DÑDð ñ ð
  !ÑØ	×	Ñ	Ü�	‰	×&Ñ&×;Ñ;×CÑCÜ�	‰	×&Ñ&×=Ñ=×EÑEð"ñ 
ð
  !ÑäØ7¸×8LÑ8LÐ7MÈQÐOó
ð 	
ð 	×ÑÐ3Ð3Ð4Ø
×
Ñ
Ð1Ð2Ð
3ñ	4ð Õr4   c                 óÖ  — t        j                  ddd«      t        j                  ddd«      t        j                  d«      t        j                  d«      t        j                  d«      t        j                  d«      f}t        j                  dddd«      t        j                  dddd«      t        j                  d«      t        j                  d«      t        j                  d«      t        j                  d«      f}t        d„ | j                  j                  D «       «      }|s| S t         j
                  j                  «       rddgndg}|D ]v  }t        | t        j                  ||¬«      } t        | t        j                  ||¬«      } t        | t        j                  ||¬«      } t        | t        j                  ||¬«      } Œx | j                  j                  D ]Ø  }|j                  t         j                  j                  j                   j"                  k(  sŒ?|j$                  d   j&                  dk(  sŒ\|j$                  d   dk(  sŒot         j(                  j*                  j,                  j.                  |j0                  d	   D �cg c]  }|d   ‘Œ	 c}v sŒ¾| j                  j3                  |«       ŒÚ | j                  j5                  «        | j7                  «        | S c c}w )
Nr   rÔ   c              3   ó2   K  — | ]  }t        |«      –— Œ y ­wr‰   rþ   rÿ   s     r2   rç   z$_fold_conv_bn_qat.<locals>.<genexpr>]  r   r  TFr  r   Úget_attrÚsource_fn_stack)r&   r*   r  r  r’   r.   r  Ú_fold_conv_bn_qat_helperr?   r  r  r	  r
  rŠ   r`   ÚatenÚadd_r-   r‚   r•   ÚnnÚmodulesÚ	batchnormÚBatchNorm2drÛ   r*  r  r  )rû   Ú#_quantized_conv1d_bn_example_inputsÚ#_quantized_conv2d_bn_example_inputsr  r  r   rà   Úvals           r2   Ú_fold_conv_bn_qatrF  H  s=  € ô 	�‰�A�q˜!ÓÜ�‰�A�q˜!ÓÜ�‰�A‹Ü�‰�A‹Ü�‰�A‹Ü�‰�A‹ð+Ð'ô 	�‰�A�q˜!˜QÓÜ�‰�A�q˜!˜QÓÜ�‰�A‹Ü�‰�A‹Ü�‰�A‹Ü�‰�A‹ð+Ð'ô Ñ7¨¯©¯©Ô7Ó7€FÙØˆÜ',§z¡z×'>Ñ'>Ô'@�t˜U‘mÀuÀg€OØ"ò 
ˆÜ$ØŒq�x‰xÐ<Àgô
ˆô %ØŒq�x‰xÐ<Àgô
ˆô %ØŒq×!Ñ!Ð#FÐPWô
ˆô %ØŒq×!Ñ!Ð#FÐPWô
‰ð
ð —‘—‘ò %ˆà�K‰Kœ5Ÿ9™9Ÿ>™>×.Ñ.×5Ñ5Ó5Ø—	‘	˜!‘—‘ :Ó-Ø—	‘	˜!‘ Ó!Ü—‘× Ñ ×*Ñ*×6Ñ6Ø"&§)¡)Ð,=Ñ">Ö?˜3��A“Ò?ò@ð �G‰G×Ñ˜tÕ$ð%ð ‡G�G×ÑÔ!Ø‡K�K„Mà€Hùò @s   ÊK&c           	      óø  — | j                   j                  «        | j                  «        t        | «       g }t	        j
                  ddgddgddgddg«      }|D ]r  \  }}}}	|s|rŒt        ||||«      }
t        |||||	«      }t        |||fi |
¤Ž}t        |||||	«      }t        |||fi |
¤Ž}|j                  t        | ||d¬«      «       Œt | j                  «        t        | «       |D ]á  }t        |«      }|j                  «       D ]  \  }}|j                  |_        Œ t!        |d   Ž  t!        |d   Ž  d|v rd|v sJ ‚t!        |d   Ž  t!        |d   Ž  d}|d	   \  }}|d
   \  }}|d   \  }}d|v r|d   \  }}t#        ||||| «       t%        |j&                  «      j                  «       D ]  }t)        |«      sŒt+        ||«       Œ Œã | j                   j                  «        | j                  «        | S )zn
    Replace the quantized (conv + bn) pattern with conv with bn weights folded into the weights of conv.
    TF)r  r¢   r£   r¥   r¦   Nr    r¡   r8   r%   )r  r  r  r-  Ú	itertoolsÚproductr3   ru   r   rz   Úextendr   r3  rË   r�   rÛ   r7  r   r©   r€   r   r×   )rû   r5   r  r   rª   Úreplacement_optionsr   r   r   rp   r/   r  Úreplacement_patternr�   Únode_maprÑ   rØ   r%   rµ   r˜   r™   r8   s                         r2   r<  r<  €  s†  € ð ‡G�G×ÑÔ!Ø‡K�K„MÜ˜qÔ!ð €LÜ#×+Ñ+Ø	ˆuˆØ	ˆuˆØ	ˆuˆØ	ˆuˆó	Ðð 
ò&
ñ 	ØØØØñ Ñ-ØÜ=Ø˜HÐ&7¸ó
ˆô ;Ø˜HÐ&7¸À.ó
ˆô ;ØØØñ
ð ñ	
ˆô HØ˜HÐ&7¸À.ó
Ðô AØØØñ
ð ñ	
Ðð 	×ÑÜ(ØØØ#Ø $ô	õ	
ð?&
ðN ‡K�K„MÜ˜QÔàò GˆÜ-¨aÓ0ˆð 08¯©Ó/@ò 	7Ñ+ˆMÐ+Ø$1×$6Ñ$6ÐÕ!ð	7ô 	˜h Ñ7Ñ8Ü˜hÐ'7Ñ8Ñ9Ø˜HÑ$Ø! XÑ-Ð-Ð-Ü  (¨=Ñ"9Ñ:Ü  (¨>Ñ":Ñ;ð ˆ	Ø! &Ñ)‰ˆˆIØ ‘~‰ˆˆGØ# MÑ2ÑˆˆKØ˜(Ñ"Ø% kÑ2‰NˆQ�	Ü& y°+¸yÈ'ÐSTÔUô /¨q¯{©{Ó;×BÑBÓDò 	GˆMÜ.¨}Õ=Ü,¨]¸IÕFñ	Gð3Gð: ‡G�G×ÑÔ!Ø‡K�K„MØ€Hr4   )Br  rò   rH  r–   Útypingr   r   r   r   r&   Útorch.nn.functionalr?  Ú
functionalr?   Ú$torch.ao.quantization.fx._decomposedr   Ú'torch.ao.quantization.pt2e.export_utilsr   Útorch.ao.quantization.quantizerr	   r
   r   r   Útorch.fxr   r   r   Útorch.fx.subgraph_rewriterr   r   Úutilsr   r   r   r   r   Ú6torch.fx.passes.utils.matcher_with_name_node_map_utilsr   Ú__all__ÚboolÚdictÚstrr3   rC   rX   r\   ro   ru   rz   r…   r‡   r�   r�   r¨   rË   r©   r×   rß   rú   r  r  r-  r3  r7  rF  r<  © r4   r2   ú<module>r]     sQ  ðã Û Û Û ß 9Ó 9ã ß Ð Ý IÝ B÷ó ÷ .Ñ -ß U÷õ ñ ÝTà
€ðØðàðð ðð ð	ð
 
ˆ#ˆsˆ(�^óð8, (ð ,¨xó ,ð((0 hð (0°8ó (0ðV%=°8ð %=Àó %=òPð8A:ØðA:àðA:ð ðA:ð ð	A:ð
 ðA:ð óA:ðH3AØð3Aàð3Að ð3Að ð	3Að
 ð3Að ó3AðlNØðNàðNð ðNð 
ó	Nð	KØð	Kàð	Kð ð	Kð 
ó		Kð�Dð ˜Tó ð�dð ˜tó ðnÐ"2ð n°t¸CÀÀtÈTÀzÑARÐ<RÑ7Só nðb  d¨D jÑ!1ð °d¸4À¸:Ñ6Fó ð&A°ð AÀó Að2WØðWØ+/óWð@-FØ
ð-Fà"& t¨T zÑ"2ó-Fð`(˜ð (¨ó (ðVnØðnàðnð ˜#˜s˜(‘Oðnð ð	nð
 ónðb +ó ð@ ó ð2¨ð Àó ð<5˜ð 5¨ó 5ðp_Øð_àð_ð ˜#˜s˜(‘Oð_ð ð	_ð
 ô_r4   