ó
    EñiÇ*  ã                   ó*  • % S SK r S SKJr  S SKJr  S SKJs  Js  Jr	  S SK
Jr  S SKJrJrJr  / SQrS rS rS rS	 rS
 r0 \R*                  \R,                  4\_\R*                  \R,                  \R.                  4\_\R0                  \R2                  4\_\R0                  \R2                  \R.                  4\_\R4                  \R6                  4\_\R4                  \R6                  \R.                  4\_\R*                  \R.                  4\" \	R8                  5      _\R0                  \R.                  4\" \	R:                  5      _\R4                  \R.                  4\" \	R<                  5      _\R>                  \R,                  4\_\R>                  \R.                  4\" \	R@                  5      _\R2                  \R.                  4\" \	RB                  5      _\R6                  \R.                  4\" \	RD                  5      _\RF                  \R,                  4\_\RH                  \R2                  4\_\RJ                  \R6                  4\_r&\'\(\RR                  \-  4   \*S'   SS jr+S r,S r-S r.S\S\'\\RR                  \-  4   4S jr/g)é    N)ÚCallable)ÚAny)Úget_combined_dictÚMatchAllNodeÚPattern)Úfuse_conv_bnÚfuse_conv_bn_reluÚfuse_linear_bnÚfuse_convtranspose_bnÚget_fuser_methodÚget_fuser_method_newc                 ór  • UR                   UR                   :w  a  [        S5      e[        R                  [        R
                  [        R                  [        R                  [        R                  [        R                  0nU (       a‘  UR                  UR                  :w  a  [        S5      eUR                  (       d  [        S5      eUR                  (       d  [        S5      eUR                  [        U5      5      nUb  U" X5      $ [!        SX4 35      e[        R"                  R%                  X5      $ )aì  Return the fused the conv and bn modules.
Given the conv and bn modules, fuses them and returns the fused module

Args:
    is_qat: a flag for whether we are using quantization aware training fusion
    or post training quantization fusion
    conv: Module instance of type conv2d/conv3d
    bn: Spatial BN instance that needs to be fused with the conv

Examples::

    >>> m1 = nn.Conv2d(10, 20, 3)
    >>> b1 = nn.BatchNorm2d(20)
    >>> # xdoctest: +SKIP
    >>> m2 = fuse_conv_bn(m1, b1)
ú:Conv and BN both must be in the same mode (train or eval).z@Output channel of Conv2d must match num_features of BatchNorm2d.ú7Only support fusing BatchNorm2d with affine set to TrueúGOnly support fusing BatchNorm2d with tracking_running_stats set to TrueúCannot fuse train modules: )ÚtrainingÚAssertionErrorÚnnÚConv1dÚnniÚConvBn1dÚConv2dÚConvBn2dÚConv3dÚConvBn3dÚnum_featuresÚout_channelsÚaffineÚtrack_running_statsÚgetÚtypeÚNotImplementedErrorÚutilsÚfuse_conv_bn_eval)Úis_qatÚconvÚbnÚfused_module_class_mapÚfused_module_classs        Úh/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/ao/quantization/fuser_method_mappings.pyr   r      s  € ð" ‡}�}˜Ÿ™Ó#ÜØHó
ð 	
ô
 	�	‰	”3—<‘<Ü
�	‰	”3—<‘<Ü
�	‰	”3—<‘<ðÐö Ø�?‰?˜d×/Ñ/Ó/Ü ØRóð ð �y�yÜ ØIóð ð ×%×%Ü ØYóð ð 4×7Ñ7¼¸T»
ÓCÐØÑ)Ù% dÓ/Ð/ä%Ð(CÀTÀJÀ<Ð&PÓQÐQä�x‰x×)Ñ)¨$Ó3Ð3ó    c                 óÖ  • UR                   UR                   s=:X  a  UR                   :X  d  O  [        S5      eSnU (       aï  [        R                  [        R
                  [        R                  [        R                  [        R                  [        R                  0nUR                  UR                  :w  a  [        S5      eUR                  (       d  [        S5      eUR                  (       d  [        S5      eUR                  [        U5      5      nUb	  U" XU5      $ [!        SXU4 35      e[        R                  [        R"                  [        R                  [        R$                  [        R                  [        R&                  0nUR                  [        U5      5      nUb1  [        R(                  R*                  R-                  X5      nU" Xs5      $ [!        SXU4 35      e)a  Return the fused conv and bv modules.

Given the conv and bn modules, fuses them and returns the fused module

Args:
    is_qat: a flag for whether we are using quantization aware training fusion
    or post training quantization fusion
    conv: Module instance of type conv2d/conv3d
    bn: Spatial BN instance that needs to be fused with the conv

Examples::

    >>> m1 = nn.Conv2d(10, 20, 3)
    >>> b1 = nn.BatchNorm2d(20)
    >>> r1 = nn.ReLU(inplace=False)
    >>> # xdoctest: +SKIP
    >>> m2 = fuse_conv_bn_relu(m1, b1, r1)
r   Nz?Output channel of Conv2d must match num_features of BatchNorm2dr   r   r   zCannot fuse eval modules: )r   r   r   r   r   ÚConvBnReLU1dr   ÚConvBnReLU2dr   ÚConvBnReLU3dr   r   r   r    r!   r"   r#   Ú
ConvReLU1dÚ
ConvReLU2dÚ
ConvReLU3dr$   Úfusionr%   )r&   r'   r(   ÚreluÚfused_moduleÚmap_to_fused_module_trainÚmap_to_fused_module_evalÚ
fused_convs           r+   r	   r	   G   s�  € ð& �M‰M˜RŸ[™[Õ9¨D¯M©MÕ9ÜØHó
ð 	
ð 04€LÞä�I‰I”s×'Ñ'Ü�I‰I”s×'Ñ'Ü�I‰I”s×'Ñ'ð%
Ð!ð
 �?‰?˜d×/Ñ/Ó/Ü ØQóð ð �y�yÜ ØIóð ð ×%×%Ü ØYóð ð 1×4Ñ4´T¸$³ZÓ@ˆØÑ#Ù ¨$Ó/Ð/ä%Ð(CÀTÈtÐDTÐCUÐ&VÓWÐWô �I‰I”s—~‘~Ü�I‰I”s—~‘~Ü�I‰I”s—~‘~ð$
Ð ð
 0×3Ñ3´D¸³JÓ?ˆØÑ#ÜŸ™Ÿ™×:Ñ:¸4ÓDˆJÙ 
Ó1Ð1ä%Ð(BÀDÈdÐCSÐBTÐ&UÓVÐVr,   c                 ó’  • UR                   UR                   :w  a  [        S5      eU (       as  UR                  UR                  :w  a  [        S5      eUR                  (       d  [        S5      eUR
                  (       d  [        S5      e[        R                  " X5      $ [        R                  R                  R                  X5      $ )aï  Return the fused linear and bn modules.
Given the linear and bn modules, fuses them and returns the fused module

Args:
    is_qat: a flag for whether we are using quantization aware training fusion
    or post training quantization fusion
    linear: Module instance of type Linear
    bn: BatchNorm1d instance that needs to be fused with the linear layer

Examples::

    >>> m1 = nn.Linear(20, 10)
    >>> b1 = nn.BatchNorm1d(10)
    >>> # xdoctest: +SKIP
    >>> m2 = fuse_linear_bn(m1, b1)
z<Linear and BN both must be in the same mode (train or eval).z@Output features of Linear must match num_features of BatchNorm1dz7Only support fusing BatchNorm1d with affine set to TruezGOnly support fusing BatchNorm1d with tracking_running_stats set to True)r   r   r   Úout_featuresr   r    r   Ú
LinearBn1dr   r$   r4   Úfuse_linear_bn_eval)r&   Úlinearr(   s      r+   r
   r
   „   s©   € ð" ‡�˜"Ÿ+™+Ó%ÜØJó
ð 	
ö Ø�?‰?˜f×1Ñ1Ó1Ü ØRóð ð �y�yÜ ØIóð ð ×%×%Ü ØYóð ô �~Š~˜fÓ)Ð)ä�x‰x�‰×2Ñ2°6Ó>Ð>r,   c                 óÀ   • UR                   UR                   :w  a  [        S5      eU (       a  [        S5      e[        R                  R
                  R                  XSS9$ )aÓ  Return the fused ConvTranspose and bn modules.
Given ConvTranspose and bn modules, fuses them and returns the fused module

Args:
    convt: Module instance of type ConvTransposeNd
    bn: BatchNormNd instance that needs to be fused with the linear layer.
        batch norm N should match the ConvTranspose N

Examples::

    >>> m1 = nn.ConvTranspose2d(10, 20, 3)
    >>> b1 = nn.BatchNorm2d(20)
    >>> # xdoctest: +SKIP
    >>> m2 = fuse_convtranspose_bn(m1, b1)
zCConvTranspose and BN both must be in the same mode (train or eval).z8Fusing ConvTranspose+BatchNorm not yet supported in QAT.T)Ú	transpose)r   r   Ú	Exceptionr   r$   r4   r%   )r&   Úconvtr(   s      r+   r   r   ¬   sY   € ð  ‡~�~˜Ÿ™Ó$ÜØQó
ð 	
ö ÜØFó
ð 	
ô �x‰x�‰×0Ñ0°ÀdÐ0ÐKÐKr,   c                 ó   ^ • U 4S jnU$ )a  Return a sequential wrapped that for is_qat and two modules.
Given a sequential class for two modules, return a function that takes
is_qat, and then two modules as argument, that ignores the is_qat flag
and always returns the sequential that combines the two input modules
c                 ó   >• T" X5      $ ©N© )r&   Úm1Úm2Ú
sequentials      €r+   Úfuser_methodÚ*_sequential_wrapper2.<locals>.fuser_methodÐ   s   ø€ Ù˜"Ó!Ð!r,   rF   )rI   rJ   s   ` r+   Ú_sequential_wrapper2rL   É   s   ø€ õ"ð Ðr,   Ú _DEFAULT_OP_LIST_TO_FUSER_METHODc                 óx   • Uc  0 n[        [        U5      nUR                  U S5      nUc  [        SU  S35      eU$ )z–Get fuser method for the given list of module types.

Get fuser method for the given list of module types,
return None if fuser method does not exist
Núdid not find fuser method for: Ú )r   rM   r!   r   )Úop_listÚadditional_fuser_method_mappingÚall_mappingsrJ   s       r+   r   r   ê   sU   € ð 'Ñ.Ø*,Ð'Ü$Ü(Ð*Ió€Lð  ×#Ñ# G¨TÓ2€LØÑÜÐ>¸w¸iÀqÐIÓJÐJØÐr,   c                 ó   ^ • U 4S jnU$ )Nc                 ó   >• T" XU5      $ rE   rF   )r&   ÚxÚyÚfs      €r+   ÚreversedÚ_reverse2.<locals>.reversedü   s   ø€ Ù�˜A‹Ðr,   rF   ©rX   rY   s   ` r+   Ú	_reverse2r\   û   s   ø€ õð €Or,   c                 ó   ^ • U 4S jnU$ )Nc                 ó   >• Uu  p4T" XX15      $ rE   rF   )r&   rV   ÚwrW   ÚzrX   s        €r+   rY   Ú_reverse3.<locals>.reversed  s   ø€ Ø‰ˆÙ�˜AÓ!Ð!r,   rF   r[   s   ` r+   Ú	_reverse3rb     s   ø€ õ"ð €Or,   c                 óÈ   • [        U [        [        45      (       a9  U  Vs/ s H  n[        U5      PM     nn[        [        R
                  " U6 5      nU$ U [        /nU$ s  snf )a  Return a list of valid patterns generated from the op_pattern.

Returns a list of valid patterns generated from the op_pattern,
since MatchAllNode can match all types of nodes,
e.g. pattern (torch.nn.Conv2d, torch.add) should also be able to match keys like
(MatchAllNode, torch.add) and (torch.nn.Conv2d, MatchAllNode)

Example Input:
(torch.add, (torch.nn.ReLU, torch.nn.Conv2d))

Example Output:
[(torch.add, (torch.nn.ReLU, torch.nn.Conv2d)),
 (torch.add, (torch.nn.ReLU, MatchAllNode)),
 (torch.add, (MatchAllNode, torch.nn.Conv2d)),
 (torch.add, (MatchAllNode, MatchAllNode)),
 (MatchAllNode, (torch.nn.ReLU, torch.nn.Conv2d)),
 (MatchAllNode, (torch.nn.ReLU, MatchAllNode)),
 (MatchAllNode, (MatchAllNode, torch.nn.Conv2d)),
 (MatchAllNode, (MatchAllNode, MatchAllNode)),
]
)Ú
isinstanceÚtupleÚlistÚ_get_valid_patternsÚ	itertoolsÚproductr   )Ú
op_patternÚsub_patternÚ	sub_combsÚresults       r+   rg   rg   
  sb   € ô. �*œu¤d˜m×,Ñ,ÙISÓTÊ¸+Ô(¨Ö5Éˆ	ÐTÜ”i×'Ò'¨Ð3Ó4ˆð €Mð œlÐ+ˆØ€Mùò	 Us    Arj   Úfuser_method_mappingc                 ó‚   • [        U 5      nSnU H  n UR                  U 5      nUc  M    O   Uc  [        SU  S35      eU$ )zŸGet fuser method.

This will be made default after we deprecate the get_fuser_method
Would like to implement this first and have a separate PR for deprecation
NrO   rP   )rg   r!   r   )rj   rn   Úop_patternsrJ   s       r+   r   r   )  sY   € ô & jÓ1€KØ€LÛ!ˆ
Ø+×/Ñ/°
Ó;ˆØÓ#Ùñ "ð ÑÜÐ>¸z¸lÈ!ÐLÓMÐMØÐr,   rE   )0rh   Úcollections.abcr   Útypingr   Útorch.ao.nn.intrinsicÚaor   Ú	intrinsicr   Útorch.nnÚtorch.ao.quantization.utilsr   r   r   Ú__all__r   r	   r
   r   rL   r   ÚBatchNorm1dÚReLUr   ÚBatchNorm2dr   ÚBatchNorm3dr1   r2   r3   ÚLinearÚ
LinearReLUÚBNReLU2dÚBNReLU3dÚConvTranspose1dÚConvTranspose2dÚConvTranspose3drM   Údictre   Ú
SequentialÚ__annotations__r   r\   rb   rg   r   rF   r,   r+   Ú<module>r‡      s—  ðä Ý $Ý ç #Ó #Ý ß PÑ Pò€ò/4òd:Wòz%?òPLò:
ðKØ‡Y�Y�—‘Ð ðKà‡Y�Y�—‘ §¡Ð(Ð*;ðKð ‡Y�Y�—‘Ð ðKð ‡Y�Y�—‘ §¡Ð(Ð*;ð	Kð
 ‡Y�Y�—‘Ð ðKð ‡Y�Y�—‘ §¡Ð(Ð*;ðKð ‡Y�Y�—‘ÐÑ.¨s¯~©~Ó>ðKð ‡Y�Y�—‘ÐÑ.¨s¯~©~Ó>ðKð ‡Y�Y�—‘ÐÑ.¨s¯~©~Ó>ðKð ‡Y�Y�—‘Ð ðKð ‡Y�Y�—‘ÐÑ.¨s¯~©~Ó>ðKð ‡^�^�R—W‘WÐÑ3°C·L±LÓAðKð ‡^�^�R—W‘WÐÑ3°C·L±LÓAðKð ×Ñ˜Ÿ™Ð(Ð*?ðKð ×Ñ˜Ÿ™Ð(Ð*?ðKð  ×Ñ˜Ÿ™Ð(Ð*?ð!KÐ   $ u¨b¯m©m¸hÑ.FÐ'FÑ"Gó ô(ò"òòð>Øðà˜w¨¯©¸Ñ(@Ð@ÑAõr,   