
    ^j&	                         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y)    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y)r   
submodulesnodereturnc                    |j                   t        vry|j                  t        j                  j
                  j                  j                  u ry|j                  t        j                  u ry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   S | d   S )Nvalfake_result )r   s    n/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/torch/fx/passes/backends/cudagraphs.pymeta_fkz4CudaGraphsSupport.is_node_supported.<locals>.meta_fk    s    "'4-4;HT-5HH    tc                 t    t        | t        j                        r| j                  j                  dk7  rdy y y )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\\*qxx}}/F!% 0G*r   )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     77++;;%))..AAIII;;(***	I$sCx. 	IU\\ 	I	&V 	& 	&
 %% 	=A]GAFFO<	= 			(:;
 "!!r   N)__name__
__module____qualname__r   r-   r   nnModulefxNodeboolr4   r   r   r   r   r      s9    "!#uxx"67"?Dxx}}"	"r   gminputsr   c                      t        |       j                  |  t               }t        | |d      }|j	                         }|j                  |      }|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   sZ     !N2  &)%'M -
MK //1J--j9K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      sp     -    6
7 ; H < : ) "  "F&.v&6
XXr   