
    ^jF                         d dl Z d dlZd dlZd dlmZ d dlmZmZ d dlm	Z	 d dl
Z
d dlmZmZ d dlmZ ddgZd	ej"                  fd
Z e       Z ed      e G d d                    Z ed       G d d             Zy)    N)defaultdict)	dataclassfield)Any)GraphNode)compatibilitySubgraphMatcherInternalMatchreturnc                  |   t        j                  t              } t        j                  j                  dd      j                         }| j                  |       t        j                         }t        j                  d      }|j                  |       |j                  |       | j                  |       d| _        | S )NPYTORCH_MATCHER_LOGLEVELWARNINGz%(filename)s > %(message)sF)logging	getLogger__name__osenvirongetuppersetLevelStreamHandler	FormattersetFormatter
addHandler	propagate)loggerlevelconsole	formatters       n/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/torch/fx/passes/utils/matcher_utils.py_init_loggerr"      s    x(FJJNN5yAGGIE
OOE##%G!!">?I#U
gFM    F)is_backward_compatiblec                       e Zd ZU ee   ed<    ee      Zeeef   ed<    ee      Z	ee   ed<    ee      Z
ee   ed<    ee      Zeeef   ed<   d	dZy)
r   anchors)default_factory	nodes_mapplaceholder_nodesreturning_nodesname_node_mapc                     t        | j                  | j                  j                         | j                  j                         | j
                  j                               S )N)r&   r(   r)   r*   )r   r&   r(   copyr)   r*   )selfs    r!   __copy__zInternalMatch.__copy__5   sJ    LLnn))+"4499; 00557	
 	
r#   N)r   r   )r   
__module____qualname__listr   __annotations__r   dictr(   r)   r*   r+   strr/    r#   r!   r   r   #   so     $Z"'"=ItD$J= %*$$?tDz? #("=OT$Z= &+4%@M4T	?@
r#   c                       e Zd Z	 	 	 	 ddedededededdfdZd	ed
edefdZdd	ed
ededefdZ	de
eef   defdZdee   dee   fdZd	ed
ededefdZ	 dd	ed
edededef
dZddededee   fdZy)r
   patternmatch_outputmatch_placeholderremove_overlapping_matchesignore_literalsr   Nc                    || _         || _        || _        || _        || _        t        |j                        dk(  rt        d      |j                  D ]F  }|j                  dk7  s|j                         r$t        |j                        dk(  s=t        d       |j                  D cg c]  }|j                  dk(  s| c}| _        t        t        t        |j                                    }|j                   | _        g | _        |r	|g| _        y|j                   D cg c]  }t        |j                        dk(  s| c}| _        yc c}w c c}w )a  
        Args:
            pattern: the targeted matching pattern, represented in fx.Graph.
            match_output: If True, output node in the pattern graph will be treated as a part of the targeted pattern.
                If False, output node is ignored during match.
            match_placeholder: If True, placeholder node in the pattern graph will be treated as a part of
                the targeted pattern. If False, placeholder nodes will be used a wildcard.
            remove_overlapping_matches: If True, in the case of overlapping matches, only the first match
                will be returned.
            ignore_literals: If True, will not check if literals are equal and
                will instead treat them as wildcards.
        r   z;SubgraphMatcher cannot be initialized with an empty patternoutputzDSubgraphMatcher cannot be initialized with an pattern with dead codeplaceholder   N)r8   r9   r:   r;   r<   lennodes
ValueErrorop	is_impureusersAssertionErrorpattern_placeholder_nodesnextiterreversedall_input_nodespattern_returning_nodespattern_anchors)	r.   r8   r9   r:   r;   r<   nodenoutput_nodes	            r!   __init__zSubgraphMatcher.__init__@   s5   * (!2*D'.w}}"M  MM 	Dww("4>>+;tzz?a'(^ 	 }}*
(=A*
& 4 7893>3N3N$+-$/=D 
 '66$#agg,!:K$D *
$s   *E?E#EEpngnc                 &   t        |j                  t              st        d|j                   d      t        |j                  t              st        d|j                   d      |j                  j
                  t        d      |j                  j
                  t        d      t        j                  j                  j                  |j                  j
                  |j                        }t        j                  j                  j                  |j                  j
                  |j                        }t        |      t        |      uryt        |t        j                        rt        |t        j                        S t        d| d      )	Nz
pn.target z must be a string.z
gn.target z'pn.graph.owning_module must not be Nonez'gn.graph.owning_module must not be NoneFzUnsupported type z when matching attributes)
isinstancetargetr5   rG   graphowning_moduletorchfxgraph_module	_get_attrtypeTensorRuntimeError)r.   rS   rT   pn_valuegn_values        r!   _match_attributesz!SubgraphMatcher._match_attributesz   s&   "))S) :bii[8J!KLL"))S) :bii[8J!KLL88!!) !JKK88!!) !JKK88((222883I3I299U88((222883I3I299U>h/ h-h55!28*<UVWWr#   node_name_matchc                 B   | j                   s|j                  dk(  ry|r||j                  v ry|j                  |j                  k(  rY|j                  dk(  s|j                  dk(  ry|j                  dk(  r| j                  ||      S |j                  |j                  k(  S y)Nr?   Tr>   get_attrF)r:   rD   namerc   rW   )r.   rS   rT   rd   s       r!   _nodes_are_equalz SubgraphMatcher._nodes_are_equal   s    %%"%%=*@"''955BEE>uu%():*$--b"5599		))r#   r(   c                     |j                         D ci c]  \  }}|j                  dk7  s|| }}}|j                         D ],  \  }}|| j                  v r|j                  D ]	  }||vs  y . yc c}}w )Nr?   FT)itemsrD   rM   rF   )r.   r(   rS   rT   lookupusers         r!   _is_containedzSubgraphMatcher._is_contained   s     "+!2$
r2bee}6LBF$
 $
 lln 		!FBT111 ! v% 	!		! $
s
   A4A4matchesc                 N   g }t               }|D ]  }d}|j                  j                         D ]  \  }}|j                  dvs||v sd} n |rA|j	                  |       |j                  j                         D ]%  \  }}|j                  dvs|j                  |       '  |S )NF>   r>   r?   T)setr(   rj   rD   appendadd)r.   rn   non_overlapping_matchesnodes_matchedmatchfound_overlaprS   rT   s           r!   _remove_overlapping_matchesz+SubgraphMatcher._remove_overlapping_matches   s     8:#&5 	.E!M////1 B55 99bM>Q$(M
 !'..u5#oo335 .FBuu$==%))"-.	. '&r#   ru   c                    t        |t              rt        |t              rt        d      t        |t              rPt        |t              s@|j                  dk(  r0||j                  v r|j                  |   |k(  S ||j                  |<   yyt        |t              st        |t              ryt        |      t        |      u xr ||k(  S )Nzpn and gn cannot both be Noder?   TF)rV   r   rG   rD   r(   r^   )r.   rS   rT   ru   s       r!   _match_literalszSubgraphMatcher._match_literals   s    b$Jr4$8 !@AAb$
2t(<uu% ( ??2."44&(#B%*R*>8tBx'4B"H4r#   c                 *   
 t         j                  d||       t        |t              rt        |t              st	        d| d|       |j
                  v rj
                  |   |k(  S |j
                  j                         v ry j                  |||      syt        j                        }|j
                  |<   |j                  dk(  ryd}dt        t           t        t        df   z  d	t        t           t        t        df   z  d
t        f
 fd
d }d }t        |j                        t        |j                        k7  sGt        |j                   j#                               t        |j                   j#                               k7  r|j                  dk(  rt        |j$                  t&        j(                  j*                        r|j$                  j,                  j.                  dt        t        df   dt0        t2        t        f   d
t        t           ffd}	 |	|j                  |j                         } |	|j                  |j                         }nt        |j                        t        |j                        k(  rt        |j                   j#                               t        |j                   j#                               k(  rt        |j                        }t        |j                        }|j5                  t        |j                   j                                      |j5                  t        |j                   j                                      nd}|xr |d uxr |d uxr	  
||      }|st        j                  |      yy)Nz  matching %s to %szpn and gn must be Node, pn: z, gn: Fr?   Targs1.args2r   c                 |   t        |       t        |      k7  ryt        | |      D ]  \  }}t        |t              r$t        |t              rj	                  ||      }nWt        |t
        t        f      r t        |t
        t        f      r
 ||      }n!j                  ||      xs j                  }|r y y)NFT)	rA   ziprV   r   _match_nodesr2   tuplery   r<   )r{   r|   a1a2matched_match_argsru   r.   s        r!   r   z1SubgraphMatcher._match_nodes.<locals>._match_args  s     5zSZ'eU+ !Bb$'Jr4,@"//B>GT5M2z"tUm7T)"b1G ,,RU;St?S?S   ! r#   call_function	orig_argsorig_kwargsc                     g }t              D ]|  \  }}|j                  |v r|j                  ||j                            3|j                  s#|t	        |       k  r|j                  | |          b|j                  |j
                         ~ |S )N)	enumeraterg   rq   
kwarg_onlyrA   default_value)r   r   all_argsischemaargs_schemas        r!   get_all_argumentsz7SubgraphMatcher._match_nodes.<locals>.get_all_arguments#  s     !*;!7 >IAv{{k1 FKK(@A#..1s9~3E 	!5 (<(<=>  r#   )r   inforV   r   rG   r(   valuesrh   r-   rD   r2   r   r   boolrA   argskwargskeysrW   rZ   _ops
OpOverload_schema	argumentsr4   r5   extend)r.   rS   rT   ru   rd   saved_matchmatch_foundpn_argsgn_argsr   r   r   s   `  `      @@r!   r   zSubgraphMatcher._match_nodes   s    	)2r22t$B)= #?t6"!NOO  ??2&",, ''))$$R_= ii&  55M! 	9uS#X.	7;Cy5c?7R		, %)$( BGGBGG,		()T"))..2B-CC(299ejj&;&;<))++55K  c? 9=c3h c  (;G';G\S\)d299>>3C.DIINNI
 /
 277mG277mGNN4		 0 0 234NN4		 0 0 234K  .t#.t#. GW-	 	 IIk*Er#   rX   c                     ddl m} t        t              } j                  D ];  }|j
                  D ]*  } j                  ||      s||   j                  |       , = t        |j                               t        j                  d       g dt        dt        ddf fdt         j                  	      }r	 d|       t              }D cg c]   } j                  |j                        s|" c}t              }	||	k7  rt        j                  d
||	z
         g }
D ]V  }|j                  j                         D cg c]  \  }}|j                   dvr| }}} ||      sF|
j                  |       X t        |
      t              k7  r+t        j                  dt              t        |
      z
          j"                  rEt        |
      } j%                  |
      t              }	||	k7  rt        j                  d||	z
         t        j                  d       S c c}w c c}}w )a  
        Returns:
            The matched subgraphs.
            The returned subgraph would be fully self-contained, meaning the nodes (except placeholder
            and nodes returned by output) can only be consumed by nodes within the matched subgraph.

        Subgraph pattern matcher is implemented with the backtracking style in the following steps:

        1. We first identify all the anchor nodes in the pattern graph. The anchor nodes
        are the "sinks" (nodes with no user other than the output node) of the pattern graph.
        One pattern graph could have multiple anchors if it has multiple return values.

        2. In the target graph, we identify the potential candidate nodes that can be matched
        with each anchor. These anchor-candidate pairs are the starting points for
        pairwise per-node matching.

        3. For each anchor-candidate pair, we simultaneously traverse backwards (DFS) in both
        pattern and target graphs. For every pattern nodes along traversal path, we compare it
        against the target nodes. In case any comparison failed, the match for this anchor-candidate
        pair fails. A match is found when DFS completes traversing the graph. See `self._match_nodes`
        for more details.

        4. In the case of multiple anchors, every anchor will need to find a match using step 3.
        In addition, the matches found between anchors need to have a common intersection node
        in order for the match to be valid. This is implemented with backtracking. See `backtracking`
        for more details.

        Notice: graph traversal must be done in the reverse order because a tensor can have multiple
        consumers, but can only have a single producer. Only with reverse order can we jointly
        traverse the pattern and target graph in a deterministic path.

        Warning: In theory, this backtracking algorithm have an **exponential** time complexity. However,
        in practice, it's unlikely to blow up.

        r   )validate_partitionz"Initial match_candidates_list: %s
anchor_indexru   r   Nc                 J   | t        	      k(  rj                  D cg c]  }|j                  |    c}|_        j                  D cg c]  }|j                  |    c}|_        
j                  |       t        j                  d|       y 	|    \  }}t        j                  |      }|D ]h  }t        j                  d||       j                  |||      }|r | dz   |       nt        j                  d||       t        j                  |      }j y c c}w c c}w )NzFound a match: %s
zTrying to match anchor %s to %sr@   z Failed to match anchor %s to %s
)rA   rH   r(   r)   rM   r*   rq   r   r   r-   r   )r   ru   rS   pattern_anchorcandidate_nodesr   rO   r   backtrackingmatch_candidates_listrn   rd   r.   s           r!   r   z+SubgraphMatcher.match.<locals>.backtracking}  s   s#899262P2P+,.EOOB'+' 372N2N),.EOOB')% u%159.CL.Q+NO))E*K' /=~tT"//"D%  !159KK;^T
 		+./+)s   D
D )r&   z<Filtered out %s matches because they are not fully contained>   r>   r?   zfFiltered out %s matches because                           matched subgraph would form a cycle if fusedzAFiltered out %s matches because matched subgraphs are overlappingzMatches returned: %s)!torch.fx.passes.utils.fuser_utilsr   r   r2   rN   rB   rh   rq   rj   r   r   intr   rA   rm   r(   rD   r;   rw   )r.   rX   rd   r   match_candidatesr   rO   ru   beforeaftervalid_matchesrS   rT   matched_compute_nodesr   r   rn   s   ` `           @@@r!   ru   zSubgraphMatcher.matchK  s5   H 	I 4?t3D"22 	BN B((O$^4;;DAB	B !%%5%;%;%= >9;PQ')	/s 	/= 	/T 	/ 	/B d&:&:; E" W&-UU1C1CEOO1T5UGU?KKN .0 	,E $oo335%B55 99 %! %
 ""78$$U+	, }W-KKHGs=11 **'F66}EGLEWUN
 	*G4K V%s    H7?H7H<)FFTF) )r   r0   r1   r   r   rR   r   rc   r5   rh   r4   rm   r2   r   rw   r   ry   r   ru   r6   r#   r!   r
   r
   >   sF   
 #"'+/ %88 8  	8
 %)8 8 
8tD d t 24 T C QU  tD$J'7 D *'M*'	m	'(5# 53 5} 5 5* PRii i)6iILi	iV~5 ~3 ~]@S ~r#   )r-   r   r   collectionsr   dataclassesr   r   typingr   rZ   torch.fxr   r   torch.fx._compatibilityr	   __all__Loggerr"   r   r   r
   r6   r#   r!   <module>r      s      	 # (     1 o
.gnn  
 e,

 
  -
2 e,J J -Jr#   