§
    ŠŠtj&	  ã                   óÚ   — d dl Z d dlmZmZ d dlmZ d dlZddgZd dlm	Z	 d dl
mZ d dlmZ d d	lmZ d d
lmZ  G d„ de¦  «        Zdej        j        dee         dej        j        fd„ZdS )é    N)ÚMappingÚSequence)ÚAnyÚCudaGraphsSupportÚpartition_cudagraphs)ÚFakeTensorProp)ÚCapabilityBasedPartitioner)ÚOperatorSupport)ÚCALLABLE_NODE_OPS)Ú_pytreec                   óZ   — e Zd Zdeeej        j        f         dej        j	        de
fd„ZdS )r   Ú
submodulesÚnodeÚreturnc                 ó°  ‡— |j         t          vrdS |j        t          j        j        j        j        u rdS |j        t          j	        u rdS dŠdt          t          t          f         dt          j        fd„}dt          dd fˆfd„}|j        D ]%}t!          j        | ||j        ¦  «        ¦  «         Œ&t!          j        | ||j        ¦  «        ¦  «         ‰ S )NFTÚmetar   c                 ó*   — d| v r| d         n| d         S )NÚvalÚfake_result© )r   s    úa/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/fx/passes/backends/cudagraphs.pyÚmeta_fkz4CudaGraphsSupport.is_node_supported.<locals>.meta_fk    s   € Ø"'¨4 - -�4˜”;�;°T¸-Ô5HÐHó    Útc                 óh   •— t          | t          j        ¦  «        r| j        j        dk    rdŠd S d S d S )NÚcudaT)Ú
isinstanceÚtorchÚTensorÚdeviceÚtype)r   Úfound_not_cudas    €r   Úfind_not_cudaz:CudaGraphsSupport.is_node_supported.<locals>.find_not_cuda#   s>   ø€ å˜!�Uœ\Ñ*Ô*ð &¨q¬x¬}ÀÒ/FÐ/FØ!%���ð&ð &Ð/FÐ/Fr   )Úopr   Útargetr   ÚopsÚatenÚembedding_dense_backwardÚdefaultÚoperatorÚgetitemÚdictÚstrr   r   ÚobjectÚall_input_nodesÚpytreeÚ	tree_map_r   )Úselfr   r   r   r#   Únr"   s         @r   Úis_node_supportedz#CudaGraphsSupport.is_node_supported   s
  ø€ ð Œ7Õ+Ð+Ð+Ø�5àŒ;�%œ)œ.ÔAÔIÐIÐIØ�5àŒ;�(Ô*Ð*Ð*Ø�4àˆð	I�$�s¥C˜xœ.ð 	I­U¬\ð 	Ið 	Ið 	Ið 	Ið	&�Vð 	&¨ð 	&ð 	&ð 	&ð 	&ð 	&ð 	&ð
 Ô%ð 	=ð 	=ˆAÝÔ˜]¨G¨G°A´F©O¬OÑ<Ô<Ð<Ð<åÔ˜¨¨°´	Ñ(:Ô(:Ñ;Ô;Ð;ð
 "Ð!Ð!r   N)Ú__name__Ú
__module__Ú__qualname__r   r-   r   ÚnnÚModuleÚfxÚNodeÚboolr4   r   r   r   r   r      sR   € € € € € ð"Ø! # u¤x¤Ð"6Ô7ð"Ø?D¼x¼}ð"à	ð"ð "ð "ð "ð "ð "r   ÚgmÚinputsr   c                 óÆ   —  t          | ¦  «        j        |Ž  t          ¦   «         }t          | |d¬¦  «        }|                     ¦   «         }|                     |¦  «        }|S )zÉ
    Partition an FX graph into sub-GraphModules that can be validly run under
    CUDA graphs.  For a subgraph to be runnable under CUDA, all of the operations
    must involve CUDA tensors only/
    T)Úallows_single_node_partition)r   Ú	propagater   r	   Úpropose_partitionsÚfuse_partitions)r=   r>   Úsupported_opsÚpartitionerÚ
partitionsÚfused_graphs         r   r   r   3   sp   € ð !…N�2ÑÔÔ  &Ð)Ð)Ý%Ñ'Ô'€Mõ -Ø
ˆM¸ðñ ô €Kð ×/Ò/Ñ1Ô1€JØ×-Ò-¨jÑ9Ô9€KØÐr   )r*   Úcollections.abcr   r   Útypingr   r   Ú__all__Ú torch.fx.passes.fake_tensor_propr   Ú!torch.fx.passes.infra.partitionerr	   Ú torch.fx.passes.operator_supportr
   Útorch.fx.passes.tools_commonr   Útorch.utilsr   r0   r   r:   ÚGraphModuler.   r   r   r   r   ú<module>rQ      s  ðØ €€€Ø -Ð -Ð -Ð -Ð -Ð -Ð -Ð -Ø Ð Ð Ð Ð Ð à €€€ð Ð 6Ð
7€Ø ;Ð ;Ð ;Ð ;Ð ;Ð ;Ø HÐ HÐ HÐ HÐ HÐ HØ <Ð <Ð <Ð <Ð <Ð <Ø :Ð :Ð :Ð :Ð :Ð :Ø )Ð )Ð )Ð )Ð )Ð )ð "ð  "ð  "ð  "ð  "˜ñ  "ô  "ð  "ðFØŒÔðØ&.¨vÔ&6ðà
„XÔðð ð ð ð ð r   