
    ^jx                     @   d dl Z d dlZd dlZd dlmZ d dlmZmZmZm	Z	m
Z
mZ d dlmZ d dlmZ d dlmZ d dlmZ d dlZd dlZd dlZd dlmZ d d	lmZ d d
lmZ d dlmZmZm Z m!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, de-e.ej^                  j`                  jb                     edef   f   dejd                  jf                  de4e5ej^                  j`                  jb                  f   de6fdZ7ddddddededee8e8f   de6de6dejd                  jf                  dz  de&e-e8e8f      dz  de6fd Z9dejd                  jf                  d!eejd                  jt                     d"ed#ede	e-ejd                  jt                  e-edf   dz  f      f
d$Z; ed%       G d& d'             Z< G d( d)      Z=d*ej|                  de8fd+Z?dejd                  jf                  de8dz  fd,Z@d-ed.ejd                  jt                  dz  defd/ZA	 d4d-ejd                  jf                  d.ejd                  jt                  dz  de-e6e-edf   e4e5ef   f   fd0ZBdejd                  jf                  de6fd1ZCdejd                  jf                  de8dz  fd2ZDdejd                  jf                  de6fd3ZEy)5    N)defaultdict)Callable
Collection	Container	GeneratorIterableMapping)	dataclass)partial)chain)Any)enable_python_dispatcher)control_deps)FakeTensorMode)compute_unbacked_bindingsrebind_unbackedstatically_known_truesym_eq)_pytree)
OrderedSet)tree_map)flop_registry   )Vpattern.nodemodulesreturnc                    t        |j                        dk(  ryt        |j                  d   t        j                  j
                        r$t        |t        j                  j
                        sy|j                  d   j                  dk7  ryt        |j                  d   j                  t              sy|j                  d   j                  |vryt        ||j                  d   j                           | d   ury|j                  dk7  r|j                  dk7  ry|j                  | d   k7  ryt        |j                  d   j                        dkD  ryy)Nr   Fcall_modulecall_functioncall_methodr   T)lenargs
isinstancetorchfxNodeoptargetstrtypeusers)r   r   r   s      c/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/torch/_inductor/fx_utils.pymatches_module_function_patternr/   *   s
   
 499~diilEHHMM2*ehhmm; yy|-'diil))3/yy|')GDIIaL''();ww/!dgg&>{{gaj 
499Q<"    Tcheck_stridescheck_storager   recursive_idsnewoldexisting_storagesr2   r3   r4   c                   d }t        |       t        |      uryt        | t        j                        s'| |du S t        | t              rv
t               t        |       t        |      f}|v }	j                  |       |	xs; t        |       t        |      k(  xr" t        fdt        | |      D              S t        | t        j                  j                        rr| j                  j                  j                  t!        j"                  | j                  j$                  |j                  j$                              t         j&                  k(  S | |k(  S | j(                  |j(                  k7  s || j*                  |j*                        sy| j,                  |j,                  k7  ryrC| j(                  t        j.                  k(  r& || j1                         |j1                               sysyt3        | j5                         |j5                         k(        rt7        |       t7        |      k7  ryd }
t7        |         dk(  rt7        |       vr	 |
      syy)a  Validate that two FakeTensors (or iterables thereof) are the same, including
    storage locations if enabled.

    check_strides: disabling this flag will remove checks for striding.
    check_storage: disabling this flag will remove checks for storage offset and
    location.  This is useful for subgraph argument and output updating, where storage
    location can change without invalidating the subgraph.
    recursive_ids: This is only for use when recursing through collections, and must not
    be supplied by users of this function.c                 ,    t        t        | |            S N)r   r   )r5   r6   s     r.   is_intlist_samez-_is_fake_tensor_same.<locals>.is_intlist_same]   s    $VC%566r0   FNc           
   3   H   K   | ]  \  }}t        ||         yw)r1   N)_is_fake_tensor_same).0new_iold_ir3   r2   r7   r   r4   s      r.   	<genexpr>z'_is_fake_tensor_same.<locals>.<genexpr>t   s>       %u ))&3&3!&3 s   "Tc           	         t        | j                  d   t        j                        sy| j                  D ]  }t        |j
                  t        j                  j                        s<|j
                  t        j                  j                  j                  j                  u s yt        |j
                  t        j                  j                        rt        |      \  }}}|s yt        j                  5  t!               5  t#        j$                         5 }t        j                  j&                  x}|j)                  |j+                                 |j
                  |i |}d d d        d d d        d d d        t        t        j                        s yt-        |      t-        | j                  d         k(  s y y# 1 sw Y   _xY w# 1 sw Y   cxY w# 1 sw Y   gxY w)NvalTF)r%   metar&   Tensorr-   r*   _opsOperatorBase	_inductor	fx_passes	reinplace_generalized_scatterHigherOrderOperatorget_fake_args_kwargsr   	fake_moder   
contextlib	ExitStack	shape_enventer_contextignore_fresh_unbacked_symbolsget_storage)r   useris_validr$   kwargsstackrQ   new_fake_tensors           r.   any_user_may_aliasz0_is_fake_tensor_same.<locals>.any_user_may_alias   s   $))E*ELL9JJ %	D4;;

(?(?@;;??,,66KKL $++uzz'E'EF  &:$%?"HdF 	?(*	? $$&	? +0 "#!6!66IC''	(O(O(QR"-$++t">v">	? 	? 	? ou||<?+{499U;K/LLK%	N %	? 	? 	? 	? 	? 	?s=   5G GAG	#G+GG
GGGG"	r   )r,   r%   r&   rE   r   r   idaddr#   allziptypespy_sym_typesr   rQ   _maybe_evaluate_staticsympyEqexprtruelayoutshapedevicestridedstrider   storage_offsetrT   )r5   r6   r7   r2   r3   r   r4   r;   id_pairvisitedrZ   s     `````    r.   r=   r=   I   s   (7 CyS	!c5<<(;$;c:&$ *#w3(G.Gg&
  CCH$   ),C   c5;;334""99HHSXX]]CHHMM: :: cz
zzSZZsyy#))'L
zzSZZ 	JJ%--'

cjjl;  2 2 44	S	[-	-,d 	+c*+q0$55"4(r0   valid_subgraphsr$   rW   c           
   /   6  K   dt         j                  j                  dt        t         j                  df   fd}dt         j
                  dt        fd}| j                  t         j                  j                  j                  u r(|d   g d |d	   D        d
 |d	   D        |d   f y| j                  t         j                  j                  j                  u r!t        |d         }|d	   |f |d   |f y| j                  t         j                  j                  j                  u r|d   } ||      }|d   d   }	 ||	      }
t        t         j
                     d t        |d	d |
dd       D              }|d   j
                  |d   j
                  k(  rFt        |      d	k(  r8 ||j!                               r"t#        d t        |dd |
dd       D              sJ d       |g |dd |d   f |d   d   g |
dd |d   f y| j                  t         j                  j                  j$                  u r1|d   } ||      }|d   } ||      }|d   d   }	 ||	      }
t        t         j
                     d t        |d	d |d	d |
dd       D              }|d   j
                  |d   j
                  k(  ri|d   j
                  |d   j
                  k(  rJt        |      d	k(  r< ||j!                               r&t#        d t        |dd |dd |
dd       D              sJ d       |v r|g |dd |d   f |v r|g |dd |d   f |	v r|	g |
dd |d   f yy| j                  t         j                  j                  j&                  t         j                  j                  j(                  t         j                  j                  j*                  fv r|d   t        |d	d       f y| j                  t         j                  j                  j,                  u r|d   t        |dd       f y| j                  t         j                  j                  j.                  u r|d   g d |d	   D        |d   f y| j                  t         j                  j                  j0                  u r!|d   g |d	   d |d   D        |d   f y| j                  t         j                  j                  j2                  t         j                  j                  j4                  fv r g |d   |d   }|d   |f |d	   |f y| j                  t6        u rR|rJ d       |d	   }t        |dd       }|v r||f t9        |t         j                  j                        rt;        j<                  t         j                  j                  fd |      rt?        tA        |jB                  jE                  d!"      |d#$            |jB                  jF                  D cg c]  }|jH                  d%k(  s| }}t        |      d	k(  sJ |d   }d&tJ        dtJ        ffd'}t;        jL                  ||jN                  |jP                  f      \  }}tS        |g|i |E d{    yyytU        jV                  d(| j                   d)       t;        jX                  ||f      \  }}fd*|D        E d{    yc c}w 7 ]7 w)+aA  HOPs that invoke subgraphs take a number of different forms.  This function
    regularizes them, yielding subgraphs from the args and kwargs coupled with
    appropriate args to update those subgraphs.

    If the second yielded value is None, this function was unable to determine what args
    to pass to the subgraph.subgraphr   .c                 `     t         fd j                  j                  d      D              S )Nc              3   6   K   | ]  }t        |        y wr:   )get_fake)r>   nrp   s     r.   rA   zI_extract_subgraphs_and_args.<locals>.get_subgraph_args.<locals>.<genexpr>   s      
&'HQ!
s   placeholderr)   )tuplegraph
find_nodes)rp   s   `r.   get_subgraph_argsz6_extract_subgraphs_and_args.<locals>.get_subgraph_args   s1      
+3>>+D+D+D+V
 
 	
r0   dc                 :    | j                    xr | j                   S r:   )is_floating_point
is_complex)r{   s    r.   
is_integerz/_extract_subgraphs_and_args.<locals>.is_integer   s    &&&;q||+;;r0   r   c              3   &   K   | ]	  }|d      ywr   N r>   as     r.   rA   z._extract_subgraphs_and_args.<locals>.<genexpr>        01!A$0   r   c              3   &   K   | ]	  }|d      ywr   r   r   s     r.   rA   z._extract_subgraphs_and_args.<locals>.<genexpr>   s     3JQAaD3Jr            c              3   4   K   | ]  }|j                     y wr:   dtyper   s     r.   rA   z._extract_subgraphs_and_args.<locals>.<genexpr>  s      5
AGG5
      Nc              3   T   K   | ]   }t        |j                               d k(   " ywr   r#   sizer   s     r.   rA   z._extract_subgraphs_and_args.<locals>.<genexpr>  s'       AFFH"   &(z/flex_attention subgraph arg format has changed!      	   c              3   4   K   | ]  }|j                     y wr:   r   r   s     r.   rA   z._extract_subgraphs_and_args.<locals>.<genexpr>%  s      5
 GG5
r   c              3   T   K   | ]   }t        |j                               d k(   " ywr   r   r   s     r.   rA   z._extract_subgraphs_and_args.<locals>.<genexpr>2  s'       AFFH"r      z8flex_attention_backward subgraph arg format has changed!      c              3   &   K   | ]	  }|d      ywr   r   r   s     r.   rA   z._extract_subgraphs_and_args.<locals>.<genexpr>L  r   r   c              3   &   K   | ]	  }|d      ywr   r   r   s     r.   rA   z._extract_subgraphs_and_args.<locals>.<genexpr>P  s     #:QAaD#:r   zfSubgraph arguments can be renamed, so we cannot consistently handle kwargs at this point in the stack.c                     | v S r:   r   )rp   rn   s    r.   <lambda>z-_extract_subgraphs_and_args.<locals>.<lambda>e  s    X8 r0   ru   rv   Tstrictr!   itemc                 t    t        | t        j                  j                        rj	                  | |       S | S r:   )r%   r&   r'   r(   get)r   placeholder_to_args    r.   replace_placeholderz8_extract_subgraphs_and_args.<locals>.replace_placeholderu  s-    dEHHMM2-11$==r0   z1Please add support for subgraph args to function !c              3   v   K   | ]0  }t        |t        j                  j                        r
|v r|d f 2 y wr:   )r%   r&   r'   GraphModule)r>   srn   s     r.   rA   z._extract_subgraphs_and_args.<locals>.<genexpr>  s8      
!UXX112qO7K I
s   69)-r&   r'   r   rw   rE   r   boolr*   opshigher_orderassociative_scancondflex_attentionr   r   r#   popr]   flex_attention_backwardforeach_mapinvoke_quant_packedinvoke_quantinvoke_subgraphmap_implscan
while_loopwhile_loop_stack_outputr   r%   pytreetree_any_onlydictr^   rx   ry   nodesr)   r   r   r$   rW   _extract_subgraphs_and_argswarningswarntree_flatten)r   rn   r$   rW   rz   r   subgraph_argsscore_subgraphscore_subgraph_argsmask_subgraphmask_subgraph_argsinteger_arg_dtypesfw_subgraphfw_subgraph_argsjoint_subgraphjoint_subgraph_argscontrol_deps_subgraphcontrol_deps_argsrt   wrapped_nodeswrapped_noder   wrapped_argswrapped_kwargsflat_args_kwargs_r   s    `                        @r.   r   r      s    
((&&
	u||S 	!
<ekk <d < {{eii,,=== 1gU0Q0U3J$q'3JUTRSWUUU			..33	3d1g1g}$$1g}$$			..==	= 04Aw/?Q.}='4 5
"#6q#;=OPRQR=ST5
 
  "((DGMM9&'1,-1134 22A68J2A8NO 		= =	= B 3BQ 7B$q'BBB1gbk>/3>d1g>>>			..FF	F1g,[9a/?Q.}='4 5
 1%#Aa("2A&5
 
 Q%%a6#A&,,Q=&'1,-1134 $Ra('+&r* 	F F	F /)A!1"1!5ARAAA_, "G$7$;"Gd2h"GGGO+!E#5bq#9!EDH!EEE ,			**		22		++ 

 1guT!"X&&			..>>	>1guT!"X&&			..77	71g;0Q0;47;;;			..33	3 1gEaE#:$q'#:ET!WEEE			))		66 
 -$q',DG,1g}$$1g}$$		$ 	
8	
z !%Q!$qr(O O3'):::!588#7#7
""HH  8

 "&)//::m:L%" 166<<@WM  }%***(+L# # 
 ,2??#""L$7$78,(L. 3o0<@N  7

B 	?}AN	
 %114.A!
%
 	
 	
5	
s8   X/\3\\A+\7\8A\
\\\)frozenc                       e Zd ZU ej                  j
                  ed<   ej                  j                  j                  ed<   e	e
df   ed<   e	e
df   ed<   y)_FxNodeHashr   r*   .
arg_hasheskwarg_hashesN)__name__
__module____qualname__r&   r'   r(   __annotations__r   Targetrw   intr   r0   r.   r   r     s@    
((--HHMM   c3hS/!r0   r   c                       e Zd ZdZdej
                  j                  ddfdZdej
                  j                  de	fdZ
defdZy)	FakeTensorUpdatera`  
    The main idea here is that it's difficult to maintain accurate fake
    tensors (our primary form of metadata) for each node in our graph as we
    transform it.

    The most reliable way to obtain this information is by rerunning
    faketensor propagation. However, in general, faketensor propagation is
    fairly expensive. So, instead we'd like to only rerun faketensor
    propagation on nodes that have changed.

    In order to detect which nodes have changed, we first hash its node,
    target, and argument lists (which are immutable in FX).

    Then, whenever we call incremental_update, we check which FX nodes have a
    new hash, and recompute the faketensor metadata for that node. Then, we
    continue to recursively compute the faketensors for all users until the
    fake tensors stop changing.

    Since this runs in the context of Inductor, we assume that the input and
    output semantics for the outermost graph are not subject to change after class
    initialization, but we allow striding changes for subgraphs.  Any other changes will
    result in errors or undefined behavior.
    gmr   Nc                 |   t        t                  | _        || _        ddlm}  || j                        D ci c]$  }t        | j                  |      x}t        |      & c}| _        | j                  j                  j                  D ],  }| j                  j                  | j                  |             . y c c}w )Nr   )_get_subgraph_names)r   r   processed_hashesr   torch._inductor.compile_fxr   getattrr   subgraph_updatersrx   r   r\   	hash_node)selfr   r   subgraph_namerp   r   s         r.   __init__zFakeTensorUpdater.__init__  s     *; 7 9 	C "5TWW!=Q
 !-88X;LX;VVQ

 GGMM'' 	<D!!%%dnnT&:;	<Q
s   )B9r   c                     dt         t           dt        t        df   fd}t	        ||j
                   ||j                         |t        j                  |j                  j                                           S )Nn_iterr   .c                 L    dt         dt        fdt        fd| D              S )a  Replace unhashable items from the input with the ids of those items, and
            hash all other items.  This is kludgy, but allows us to attempt to account
            for unhashable classes like torch.fx.Node when hashing args and kwargs for
            other nodes.or   c                 N    	 t        |       S # t        $ r t        |       cY S w xY wr:   )hash	Exceptionr[   )r   s    r.   get_hash_or_idzLFakeTensorUpdater.hash_node.<locals>.get_hash_or_ids.<locals>.get_hash_or_id  s(    !7N  !a5L!s   
 $$c              3   .   K   | ]  } |        y wr:   r   )r>   rt   r   s     r.   rA   zGFakeTensorUpdater.hash_node.<locals>.get_hash_or_ids.<locals>.<genexpr>  s     ;q*;   )objectr   rw   )r   r   s    @r.   get_hash_or_idsz4FakeTensorUpdater.hash_node.<locals>.get_hash_or_ids  s)    !& !S ! ;F;;;r0   )r   r   rw   r   r   r*   r$   r   from_iterablerW   items)r   r   r   s      r.   r   zFakeTensorUpdater.hash_node  sf    	<HSM 	<eCHo 	< KKDII&E//0A0A0CDE	
 	
r0   c           
          t        t              } j                  j                  j                  D ]  }t        |      x}s||xx   dz  cc<    dt        j                  j                  dt        fd}dt        j                  j                  dt        dt        dt        f fd}t        t                  }d}t        t                  }i }	 j                  j                  j                  D ]  }|j                   j                  |      x}
       t        | j                        \  }}}|sC ||g|i |x}s|
 j                   v rt#        |      |vrm ||      sv|rRd	}t%        | j&                  g|i |D ]   \  }}||	v}|rd	|	|<   |t)        |j                  j+                  d      |d      D ]U  \  }}t-        |t/        ||      x}|d	      r"|sJ d       t-        |||d	d	      sJ d       ||j0                  d<   |dz  }W |rt        |j                  j3                         |      \  }}}| j&                  |   j5                         z  }t        |j                  j3                         |      \  }}}t-        |||      sd|	|<   |xs |	|   }# |sd|j0                  v rt6        j8                  5  t;               5   |j<                  |i |}d
d
d
       d
d
d
       d|j0                  v rt-        |j0                  d   ||      r5t?        t6        j8                  j@                  |       ||j0                  d<   |dz  }t6        j8                  j@                  x}rtC        ||      x}r||j0                  d<   t        |      x}r||xx   dz  cc<   |jE                  d |jF                  D                | _        |S # 1 sw Y   xY w# 1 sw Y   xY w)zUpdate FakeTensors on self.graph. We will try to do the minimum amount of work.

        Returns the number of nodes updated, including recursive updates on subgraphs.r   r   r   c                 ^    t        | j                        xr t        | j                  d       S )N_inductor_lowering_function)callabler*   hasattrr   s    r.   should_process_nodezAFakeTensorUpdater.incremental_update.<locals>.should_process_node  s-    % L  -JKKr0   r$   rW   c                 :   | j                   dk(  xr | j                  t        j                  j                  j
                  t        j                  j                  j                  fvxr4 t        j                  t        j                  j                  fd||f      S )Nr!   c                      | j                   v S r:   )r   )r   r   s    r.   r   zUFakeTensorUpdater.incremental_update.<locals>.node_invokes_subgraph.<locals>.<lambda>  s    a4#9#99 r0   )r)   r*   r&   r   r   auto_functionalizedauto_functionalized_v2r   r   r'   r   )r   r$   rW   r   s      r.   node_invokes_subgraphzCFakeTensorUpdater.incremental_update.<locals>.node_invokes_subgraph  s     ?* KK
 II**>>II**AA ((HH((96Nr0   r   FNru   rv   Tr   )r3   z*subgraph args must have consistent values!)r2   r3   ziA subgraph argument other than striding has been modified; FakeTensorUpdater cannot update this argument!rC   r  unbacked_bindingsc              3   2   K   | ]  }t        |        y wr:   )r[   )r>   rU   s     r.   rA   z7FakeTensorUpdater.incremental_update.<locals>.<genexpr>  s     >4bh>s   )$r   r   r   rx   r   get_node_storager&   r'   r(   r   r   r   r   r\   r   rM   r   r[   r   r   r^   ry   r=   rs   rD   output_nodeincremental_updater   rN   r   r*   r   rQ   r   updater-   )r   r7   r   storager  r  current_graph_hashesnodes_updated
to_processsubgraph_updatingsr   rV   r$   rW   invokes_subgraphany_output_updatedrp   r   update_subgraphpr   p_faker   orig_output_argsnew_output_argsrY   rQ   symbol_to_paths   `                           r.   r  z$FakeTensorUpdater.incremental_update  sV    4?s3CGGMM'' 	0D*400w0!'*a/*	0	ehhmm 	 		((--	(+	7:		0  *+68_&
?AGGMM'' |	?D $$T^^D-A%ATB%9$%H"HdF *?t)Ud)Uf)UU%UD111tHJ.&t,%*"/J$000370;A0 C+Hm '/6H&HO&7<*84 %0$'$NN555G)#'% !3DAq $8 !*21h*? ? 1.3	$ (7 !"$P!"
 (<$%$*$52727(" 
!"%0
!" (" 12u - 2C!3F '1E$NN668(2.+Q &)?)?$*,,./ 1E$NN668(1-?A  4+,- 
 <@.x8 +J.@.J 'CCT *etyy.@ ?68 ?"-$++t">v">? ? 		!&:5!13D4' AKK114I.DIIeQM[[222	2";I"WWW 2@		-.*400w0!'*a/*>4::>>y|	?| !57? ? ? ?s$   O(O&O(O%!O((O1	)r   r   r   __doc__r&   r'   r   r   r(   r   r   r   r  r   r0   r.   r   r     sM    0<588// <D <
ehhmm 
 
,pC pr0   r   tc                 6    | j                         j                  S r:   )untyped_storage_cdata)r  s    r.   rT   rT     s    %%%r0   c                     d| j                   vry t        | j                   d   t        j                        sy t        j                  j                  | j                   d         sy t        | j                   d         S )NrC   )rD   r%   r&   rE   _C_has_storagerT   r  s    r.   r
  r
    s]    DIIdii&588  5!12tyy'((r0   xr   c                 n   t        | t        j                  j                        rd| j                  v r| j                  d   S d| j                  v r| j                  d   S | j
                  dk(  rFt        | j                  t              r,t        || j                        rt        || j                        S y| S )zReturn a fake tensor from the meta values of an input FX node.  If the input node
    is a get_attr node, we attempt to resolve it as a member of gm.rC   example_valueget_attrN)
r%   r&   r'   r(   rD   r)   r*   r+   r   r   )r#  r   s     r.   rs   rs     s     !UXX]]#AFF?66%= aff$66/**44:*QXXs";AHH@U2qxx(( Hr0   c                     t        t        t        |      | j                  | j                  f      \  }}t        d t        j                  |i |D              rd||fS d||fS )z{
    First value returns a boolean if any of the input nodes don't have a faketensor and
    weren't resolved from gm.
    )r   c              3   d   K   | ](  }t        |t        j                  j                         * y wr:   )r%   r&   r'   r(   r   s     r.   rA   z'get_fake_args_kwargs.<locals>.<genexpr>  s$      )*
1ehhmm$s   .0FT)r   r   rs   r$   rW   anyr   arg_tree_leaves)r#  r   r$   rW   s       r.   rM   rM     si     GH4qvvqxx6HILD&
 .4.D.Dd.Uf.U  dF""vr0   c                    ddl mm dt        j                  j
                  dt        ffd |       rydt        j                  j
                  dt        ffdt        fd| j                  D              ryy	)
zReturns true if a node is always realized when lowered to inductor IR.

    NOTE: This may return some false negatives. e.g. it doesn't
    handle buffers realized heuristically during lowering, or
    buffers realized indirectly through view ops.
    r   )	fallbacksneeds_realized_inputsr   r   c                     | j                   dk(  r1| j                  t        j                  u r | j                  d         S | j                   dv xs | j                  v S )Nr!   r   )ru   output)r)   r*   operatorgetitemr$   )r   r,  	is_buffers    r.   r2  z#is_node_realized.<locals>.is_buffer  sS    77o%$++9I9I*I TYYq\**ww33Ot{{i7OOr0   Tc                 B    | j                   dk(  xs | j                  v S )Nr/  )r)   r*   )r   r-  s    r.   realizes_inputsz)is_node_realized.<locals>.realizes_inputs  s!    ww("Jdkk5J&JJr0   c              3   .   K   | ]  } |        y wr:   r   )r>   rU   r4  s     r.   rA   z#is_node_realized.<locals>.<genexpr>  s     
8T?4 
8r   F)	torch._inductor.loweringr,  r-  r&   r'   r(   r   r)  r-   )r   r,  r2  r-  r4  s    @@@@r.   is_node_realizedr7    sm     JP P$ P Kehhmm K K 
8TZZ
88 r0   c           	         t        |       rt        | j                  t              ry t	        d      5  t        |       \  }}}|r| j                  t        j                  j                  j                  t        j                  j                  j                  fv rOt        j                  | j                        }|. ||i |d| j                  j                  d      icd d d        S t        j                  j                  j!                  d      5 } | j                  |i | d d d        j#                         }|cd d d        S 	 d d d        y # 1 sw Y   .xY w# 1 sw Y   y xY w)NT)allow_non_fake_inputsout_valrC   F)display)countable_fxr%   r*   r+   r   rM   r&   r   r   r   r   r   r   rD   utilsflop_counterFlopCounterModeget_total_flops)r   successr$   rW   flop_formulaflop_counter_modecounted_flopss          r.   count_flops_fxrE    s=   DKK!=	d	3 ! 4T :v {{		&&55		&&>>   -00=+'VVuAUV! ! ))99 :  -"T,V,-
 .==?M )! ! !* - -!* s*   B,E)*EE&EE	EE(c                     t        | t        j                  j                        sJ t	        | d      sy| j
                  }t	        |d      s|t        v S |j                  }|t        v S )z>
    Whether or not we can count the flops of an FX node.
    r*   Foverloadpacket)r%   r&   r'   r(   r   r*   r   rG  )r   r*   packets      r.   r<  r<    s^     dEHHMM***4"[[F6+,&&""F]""r0   r:   )FrO   r0  r   collectionsr   collections.abcr   r   r   r   r   r	   dataclassesr
   	functoolsr   	itertoolsr   typingr   rb   r&   torch.fxtorch._dispatch.pythonr   .torch._inductor.fx_passes.control_dependenciesr   torch._subclasses.fake_tensorr   %torch.fx.experimental.symbolic_shapesr   r   r   r   torch.utilsr   r   torch.utils._ordered_setr   torch.utils._pytreer   torch.utils.flop_counterr   virtualizedr   rw   r,   nnr   Moduler'   r(   r   r+   r   r/   r   r=   r   r   r   r   rE   rT   r
  rs   rM   r7  rE  r<  r   r0   r.   <module>r[     s      #  "       ; G 8  * / ( 2 
4((//0(382DDE
((-- #uxx''.../ 
	H !%8<S	S	S sCx(S
 S S ((--$
S eCHo.5S 
Sln

((--n
uxx334n
 n
 	n

 uUXX))5c?T+AABCn
b $" " "n nb&5<< &C &)588== )S4Z ) --4  & 9=xx}}((..5
4sCx$sCx.01588== T @ 3: 6#uxx}} # #r0   