Ë
    g^(h%  ã                   óÞ   — d dl mZmZ d dlmZmZmZ d dlZd dlm	Z	 d dl
mZ d dlmZmZmZ d dlmZmZ d dlmZ d	d
lmZ d	dlmZ ddgZ G d„ de«      Z G d„ de«      Zde	deeef   fd„Zy)é    )ÚABCÚabstractmethod)ÚAnyÚCallableÚUnionN)ÚBackendConfig)Úget_fuser_method_new)Ú_parent_nameÚNodePatternÚPattern)ÚGraphÚNode)Útype_before_parametrizationsé   )ÚFuseCustomConfig)ÚMatchAllNodeÚDefaultFuseHandlerÚFuseHandlerc                   óÔ   — e Zd ZdZedefd„«       Zededee	e
j                  j                  f   dededee   d	ed
edeeee
j                  j(                  ef   f   dedefd„«       Zy)r   z*Base handler class for the fusion patternsÚnodec                  ó   — y ©N© )Úselfr   s     úc/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/ao/quantization/fx/fuse_handler.pyÚ__init__zFuseHandler.__init__   s   € àó    Úload_argÚnamed_modulesÚfused_graphÚ	root_nodeÚextra_inputsÚmatched_node_patternÚfuse_custom_configÚfuser_method_mappingÚis_qatÚreturnc
                  ó   — y r   r   )
r   r   r   r    r!   r"   r#   r$   r%   r&   s
             r   ÚfusezFuseHandler.fuse#   s   € ð 	r   N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r   r   ÚdictÚstrÚtorchÚnnÚModuler   Úlistr   r   r   r   r   Ú
SequentialÚboolr)   r   r   r   r   r      sÈ   „ Ù4àð˜Tò ó ðð ðàðð ˜C §¡§¡Ð0Ñ1ðð ð	ð
 ðð ˜3‘iðð *ðð -ðð # 7¨E°%·(±(×2EÑ2EÀxÐ2OÑ,PÐ#PÑQðð ðð 
òó ñr   c                   óÈ   ‡ — e Zd Zdefˆ fd„Zdedeeej                  j                  f   dededee   ded	ed
eeeej                  j$                  ef   f   dedefd„Zˆ xZS )r   r   c                 ó$   •— t         ‰| �  |«       y r   )Úsuperr   )r   r   Ú	__class__s     €r   r   zDefaultFuseHandler.__init__4   s   ø€ Ü‰Ñ˜Õr   r   r   r    r!   r"   r#   r$   r%   r&   r'   c
                 óà  ‡‡‡‡— |j                   dk(  sJ d«       ‚‰t        |j                  «         Šˆˆˆfd„Š ‰|«      }
ˆfd„Š ‰|
«      }t        |j                  «      \  }}t	        ||«      } ||	g|
¢­Ž }t        ‰|   ||«       |D �cg c]
  } ||«      ‘Œ }}|j                  ||«      }t        |j                  «      }|j                  |«       t        |«      |_        |S c c}w )NÚcall_modulez.Expecting module node to be a call_module Nodec                 ó6  •— t        | t        t        f«      rB| ^}}g }|j                   ‰|«      «       |j	                  ˆfd„|D «       «       t        |«      S | }|j
                  dk(  r‰|j                     S |j
                  dk(  rb|j                  t        j                  j                  j                  k(  r1t        j                  j                  «       }‰j                  |_        |S |j
                  dk(  s|j
                  dk(  r|j                  S t        S )z¿Given a node pattern, extract the corresponding modules
            e.g. input: (relu_node, (bn_node, conv_node))
                 output: (relu_module, (bn_module, conv_module))
            c              3   ó.   •K  — | ]  } ‰|«      –— Œ y ­wr   r   )Ú.0ÚaÚget_moduless     €r   ú	<genexpr>z?DefaultFuseHandler.fuse.<locals>.get_modules.<locals>.<genexpr>Q   s   øè ø€ Ò<°!™{¨1Ÿ~Ñ<ùs   ƒr;   Úcall_functionÚcall_method)Ú
isinstanceÚtupler3   ÚappendÚextendÚopÚtargetr0   r1   Ú
functionalÚreluÚReLUÚtrainingr   )ÚpatternÚnÚargsÚmodulesrK   r@   r   Úroot_modules        €€€r   r@   z,DefaultFuseHandler.fuse.<locals>.get_modulesH   sÛ   ø€ ô
 ˜'¤E¬4 =Ô1Ø"���DØ13�Ø—‘™{¨1›~Ô.Ø—‘Ó<°tÔ<Ô<Ü˜W“~Ð%à�Ø—4‘4˜=Ò(Ø(¨¯©Ñ2Ð2Ø—T‘T˜_Ò,°·±¼U¿X¹X×=PÑ=P×=UÑ=UÒ1UÜ Ÿ8™8Ÿ=™=›?�DØ$/×$8Ñ$8�D”MØ�KØ—T‘T˜_Ò,°·±¸Ò0EØŸ8™8�Oä'Ð'r   c                 ó°   •— t        | t        «      rt        t        ‰| «      «      S t        | t        j                  j
                  «      rt        | «      S | S r   )rD   rE   Úmapr0   r1   r2   r   )ÚmÚget_matched_typess    €r   rV   z2DefaultFuseHandler.fuse.<locals>.get_matched_typesc   sB   ø€ Ü˜!œUÔ#ÜœSÐ!2°AÓ6Ó7Ð7Ü˜!œUŸX™XŸ_™_Ô-Ü3°AÓ6Ð6ØˆHr   )rH   r/   rI   r
   r	   ÚsetattrÚ	node_copyr3   rP   rG   rE   )r   r   r   r    r!   r"   r#   r$   r%   r&   Úmatched_modulesÚmatched_module_typesÚmodule_parent_nameÚmodule_nameÚfuser_methodÚfused_moduleÚinputÚ
extra_argsr   rP   rV   r@   rR   s     `                 @@@r   r)   zDefaultFuseHandler.fuse7   sù   û€ ð �L‰L˜MÒ)ð	<à;ó	<Ø)à#¤C¨	×(8Ñ(8Ó$9Ñ:ˆö	(ñ2 &Ð&:Ó;ˆô	ñ  1°ÓAÐÜ*6°y×7GÑ7GÓ*HÑ'Ð˜KÜ+Ð,@ÐBVÓWˆñ $ FÐ=¨_Ò=ˆÜ�Ð0Ñ1°;ÀÔMØ3?Ö@¨%‘h˜u•oÐ@ˆ
Ð@Ø×$Ñ$ Y°Ó9ˆÜ�D—I‘I‹ˆØ�‰�JÔÜ˜$“KˆŒ	Øˆùò As   ÂC+)r*   r+   r,   r   r   r   r.   r/   r0   r1   r2   r   r3   r   r   r   r   r   r4   r5   r)   Ú__classcell__)r9   s   @r   r   r   3   sª   ø„ ð˜Tõ ð?àð?ð ˜C §¡§¡Ð0Ñ1ð?ð ð	?ð
 ð?ð ˜3‘ið?ð *ð?ð -ð?ð # 7¨E°%·(±(×2EÑ2EÀxÐ2OÑ,PÐ#PÑQð?ð ð?ð 
÷?r   Úbackend_configr'   c                 óz   — i }| j                   j                  «       D ]  \  }}|j                  €Œt        ||<   Œ |S r   )Ú!_pattern_complex_format_to_configÚitemsr]   r   )rb   Úfusion_pattern_to_fuse_handlersrN   Úconfigs       r   Ú'_get_fusion_pattern_to_fuse_handler_clsrh   y   sO   € ð @BÐ#Ø)×KÑK×QÑQÓSò J‰ˆ�Ø×ÑÑ*ä7IÐ+¨GÒ4ðJð +Ð*r   )Úabcr   r   Útypingr   r   r   r0   Ú$torch.ao.quantization.backend_configr   Ú+torch.ao.quantization.fuser_method_mappingsr	   Útorch.ao.quantization.utilsr
   r   r   Útorch.fx.graphr   r   Útorch.nn.utils.parametrizer   Úcustom_configr   Úmatch_utilsr   Ú__all__r   r   r.   rh   r   r   r   ú<module>rs      sr   ðç #ß 'Ñ 'ã Ý >Ý Lß JÑ Jß &Ý Cå +Ý %ð Øð€ô�#ô ô.C˜ô CðL+Ø!ð+à	ˆ'�8Ð
Ñô+r   