ó
    Eñi`  ã                   óä   • S SK JrJr  S SKJr  S SKJr  S SKrS SK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  SS/r " S S\5      r " S S\5      rS\	S\\\4   4S jrg)é    )ÚABCÚabstractmethod)ÚCallable)ÚAnyN)ÚBackendConfig)Úget_fuser_method_new)Ú_parent_nameÚNodePatternÚPattern)ÚGraphÚNode)Útype_before_parametrizationsé   )ÚFuseCustomConfig)ÚMatchAllNodeÚDefaultFuseHandlerÚFuseHandlerc                   óÜ   • \ rS rSrSr\S\4S j5       r\S\S\	\
\R                  R                  4   S\S\S	\\   S
\S\S\	\\R                  R(                  \-  4   S\S\4S j5       rSrg)r   é   z*Base handler class for the fusion patternsÚnodec                 ó   • g ©N© )Úselfr   s     Úb/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/ao/quantization/fx/fuse_handler.pyÚ__init__Ú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
                 ó   • g r   r   )
r   r   r    r!   r"   r#   r$   r%   r&   r'   s
             r   ÚfuseÚFuseHandler.fuse$   s   € ð 	r   r   N)Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   r   r   r   ÚdictÚstrÚtorchÚnnÚModuler   Úlistr   r
   r   r   Ú
SequentialÚboolr*   Ú__static_attributes__r   r   r   r   r      sÃ   † Ù4àð˜Tó ó ðð ðàðð ˜C §¡§¡Ð0Ñ1ðð ð	ð
 ðð ˜3‘iðð *ðð -ðð # 7¨E¯H©H×,?Ñ,?À(Ñ,JÐ#JÑKðð ðð 
óó ór   c                   óÒ   ^ • \ rS rSrS\4U 4S jjrS\S\\\	R                  R                  4   S\S\S\\   S	\S
\S\\\	R                  R$                  \-  4   S\S\4S jrSrU =r$ )r   é4   r   c                 ó$   >• [         TU ]  U5        g r   )Úsuperr   )r   r   Ú	__class__s     €r   r   ÚDefaultFuseHandler.__init__5   s   ø€ Ü‰Ñ˜Õr   r   r    r!   r"   r#   r$   r%   r&   r'   r(   c
                 óì  ^^^^• UR                   S:w  a  [        S5      eT[        UR                  5         mUUU4S jmT" U5      n
U4S jmT" U
5      n[	        UR                  5      u  pÍ[        X¸5      nU" U	/U
Q76 n[        TU   Xß5        U Vs/ s H  nU" U5      PM     nnUR                  XA5      n[        UR                  5      nUR                  U5        [        U5      Ul	        U$ s  snf )NÚcall_modulez.Expecting module node to be a call_module Nodec                 óH  >• [        U [        [        45      (       aB  U tp/ nUR                  T" U5      5        UR	                  U4S jU 5       5        [        U5      $ U nUR
                  S:X  a  TUR                     $ UR
                  S:X  ab  UR                  [        R                  R                  R                  L a1  [        R                  R                  5       nTR                  Ul        U$ UR
                  S:X  d  UR
                  S:X  a  UR                  $ [        $ )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   ó4   >#   • U  H  nT" U5      v •  M     g 7fr   r   )Ú.0ÚaÚget_moduless     €r   Ú	<genexpr>Ú?DefaultFuseHandler.fuse.<locals>.get_modules.<locals>.<genexpr>Q   s   øé € Ð<²t°!™{¨1Ÿ~˜~²tùs   ƒrA   Úcall_functionÚcall_method)Ú
isinstanceÚtupler6   ÚappendÚextendÚopÚtargetr3   r4   Ú
functionalÚreluÚReLUÚtrainingr   )ÚpatternÚnÚargsÚmodulesrR   rF   r    Úroot_modules        €€€r   rF   Ú,DefaultFuseHandler.fuse.<locals>.get_modulesH   sÜ   ø€ ô
 ˜'¤E¬4 =×1Ñ1Ø"��Ø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                 óÄ   >• [        U [        5      (       a  [        [        TU 5      5      $ [        U [        R                  R
                  5      (       a  [        U 5      $ U $ r   )rK   rL   Úmapr3   r4   r5   r   )ÚmÚget_matched_typess    €r   r^   Ú2DefaultFuseHandler.fuse.<locals>.get_matched_typesc   sH   ø€ Ü˜!œU×#Ñ#ÜœSÐ!2°AÓ6Ó7Ð7Ü˜!œUŸX™XŸ_™_×-Ñ-Ü3°AÓ6Ð6ØˆHr   )rO   ÚAssertionErrorr2   rP   r	   r   ÚsetattrÚ	node_copyr6   rW   rN   rL   )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   rW   r^   rF   rY   s     `                 @@@r   r*   ÚDefaultFuseHandler.fuse8   sí   û€ ð �<‰<˜=Ó(Ü Ð!QÓRÐRØ#¤C¨	×(8Ñ(8Ó$9Ñ:ˆ÷	(ñ2 &Ð&:Ó;ˆõ	ñ  1°ÓAÐÜ*6°y×7GÑ7GÓ*HÑ'ÐÜ+Ð,@ÓWˆñ $ FÐ=¨_Ò=ˆÜ�Ð0Ñ1°;ÔMÙ3?Ó@²<¨%‘h˜u–o±<ˆ
Ð@Ø×$Ñ$ YÓ9ˆÜ�D—I‘I‹ˆØ�‰�JÔÜ˜$“KˆŒ	Øˆùò As   ÂC1r   )r,   r-   r.   r/   r   r   r   r1   r2   r3   r4   r5   r   r6   r   r
   r   r   r7   r8   r*   r9   Ú__classcell__)r>   s   @r   r   r   4   sª   ø† ð˜T÷ ð>àð>ð ˜C §¡§¡Ð0Ñ1ð>ð ð	>ð
 ð>ð ˜3‘ið>ð *ð>ð -ð>ð # 7¨E¯H©H×,?Ñ,?À(Ñ,JÐ#JÑKð>ð ð>ð 
÷>ò >r   Úbackend_configr(   c                 ó~   • 0 nU R                   R                  5        H  u  p#UR                  c  M  [        X'   M     U$ r   )Ú!_pattern_complex_format_to_configÚitemsrg   r   )rm   Úfusion_pattern_to_fuse_handlersrU   Úconfigs       r   Ú'_get_fusion_pattern_to_fuse_handler_clsrs   y   sE   € ð @BÐ#Ø)×KÑK×QÑQÖS‰ˆØ×ÑÓ*ä7IÐ+Ó4ñ Tð +Ð*r   )Úabcr   r   Úcollections.abcr   Útypingr   r3   Ú$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   r1   rs   r   r   r   Ú<module>r      sr   ðç #Ý $Ý ã Ý >Ý Lß JÑ Jß &Ý Cå +Ý %ð Øð€ô�#ô ô.B˜ô BðJ+Ø!ð+à	ˆ'�8Ð
Ñõ+r   