
    ^j5                     p   d dl Z d dlZd dlmZmZ d dlmZ d dlm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mZ d d	lmZ g d
Z ej,                  e      xZZ ed       G d d             Zde	j6                  j8                  dedefdZ ed      	 	 	 	 d-ddddede	j6                  j8                  deegef   de eef   dz  de!dz  de!dz  de!dedz  de!defd       Z"deddfdZ#	 d.d e	jH                  jJ                  jL                  d!e eef   dz  defd"Z'ded#e eef   de(e ee	jH                  jJ                  jL                  f   e ee eef   f   e ee eef   f   e)e   f   fd$Z*ded%e ee	jH                  jJ                  jL                  f   d&e ee eef   f   d'e ee eef   f   d(e)e   de ee eef   f   fd)Z+	 d.ded%e ee	jH                  jJ                  jL                  f   d&e ee eef   f   d'e ee eef   f   d(e)e   dedz  defd*Z, ed      dd+ded#e eef   dedz  defd,       Z-y)/    N)defaultdictOrderedDict)Callable)Any)compatibility)_LazyGraphModule_make_graph_module)lazy_format_graph_code)GraphModule)Node)	Partitionsplit_modulesplit_module_simpleT)is_backward_compatiblec                   (    e Zd ZdeddfdZdefdZy)r   namereturnNc                     || _         d| | _        g | _        i | _        i | _        i | _        i | _        t        j                  j                  j                         | _	        i | _        i | _        y )Nsubmod_)r   submod_name
node_namesinputsoutputsdependencies
dependentstorchfxgraphGraphenvironmenttargets)selfr   s     g/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/torch/fx/passes/split_module.py__init__zPartition.__init__   sc    	$TF+%'')(*-/+-+088>>+?+?+A
-/')    c                     d| j                    d| j                   d| j                   d| j                   d| j                   d| j
                   S )Nzname: z
,
 nodes: z,
 inputs: z,
 outputs: z,
 partitions depended on: z,
 partition dependents: )r   r   r   r   r   r   )r"   s    r#   __repr__zPartition.__repr__!   sb    TYYK  ' (} % '((,(9(9': ;&&*oo%68	
r%   )__name__
__module____qualname__strr$   r'    r%   r#   r   r      s!    
*S 
*T 
*
# 
r%   r   modqualnamer   c                     | }|j                  d      D ])  }t        ||      st        d| d      t        ||      }+ |S )N.zNode target z not found!)splithasattrAttributeErrorgetattr)r-   r.   attr_valatoms       r#   _get_attr_from_qualnamer7   ,   sP    Hs# +x& <z!EFF8T*+ Or%   F)partition_affixtuple_returnmroot_msplit_callbackqualname_mapkeep_original_orderkeep_original_node_namekeep_original_input_namer8   r9   c                !   5BCDEFGHIJKL t         j                  dt        d d             dt        dt        t
        t        f   dt        t
        t        j                  j                  j                  f   dt        t        t
        t        f   t        t
        t        j                  j                  j                  f   f   fC fd	}	d
dl}
i Ji Hi Ldt        dt        dz  ddfJLfdKdt        ddfJfd}t        j                  j                  t        j                  j                  t        j                  j                   g}t#               }t#               }i }d}t%               } j&                  j(                  D ]  55j*                  j-                  d      x}st/        |t        j0                  t        j2                  f      rIt/        |j4                  j6                  x}|
j8                        r|Lvr5L|j4                  j6                  <   5j:                  dv r |5       5j:                  dk(  r5j<                  |v r5j<                  t        j                  j                   u rt?        5j@                        dk7  r!tC        dt?        5j@                               t/        5j@                  d
   tD              s$tC        dtG        5j@                  d
                5}t%         5      h      ||<   n/5j<                  t        j                  j                  u rStI        d 5j@                  D              stC        d      |jK                  5       t%         5      h      |5<   d|5<   n5j<                  t        j                  j                  u rt?        5j@                        dk7  r!tC        dt?        5j@                               |5j@                  d
      jK                   5             |jM                  5j@                  d
          5|5j@                  d
   <   |||   jK                   5             |D ]  }||   jK                   5               tI        d |jO                         D              stC        d      |jQ                         D ci c]  \  }}|tS        |       }}}|jQ                         D ci c]  \  }}|tS        |       }}}tT        jW                  tX        jZ                        r,tT        j                  d|       tT        j                  d|       tE        |      xs tE        |      }d} j&                  j(                  D ]  55H5j\                  <   5j:                  dv r"5j:                  dk(  r;t        j                  j&                  j_                  5j@                  d
   Kfd        l|r  5      }||kD  rtC        d!| d"|       |}5j<                  |vst        j                  j&                  j_                  5j@                  5Kfd#       t        j                  j&                  j_                  5j`                  5Kfd$        tc        Jje                               }g }JjQ                         D ],  \  }It?        Ijf                        r|ji                  |       . g }|rw|jk                         }|ji                  |       J|   jl                  D ]A  }J|   jf                  jk                  |       J|   jf                  r1|ji                  |       C |rwt?        |      t?        J      k7  rto        d%      ||fD ]  } | jQ                         D ]  \  5}!t?        |!      d
k(  rtC        d&      5Jt        |!d
            jp                  5<   |!dd D ]  }"Jt        |"         IIj&                  js                  5j:                  5j<                  t        d' 5j@                  D              i 5jF                  (      }#5j*                  ju                         |#_        |#Ijp                  5<      |D ]  }J|   Ii Gd
Dg }$g }%Ijv                  D ]  FHF   }&|&j:                  d)k(  rdt/        |&j<                  t
              rJt/        ty         |&j<                        t        jz                  j|                        r|%ji                  F       {|$ji                  F        |$|%z   D ]  FHF   }&dt        fDFGHIfd*}'|&j:                  d)k(  rt/        |&j<                  t
              s!tC        d+tG        |&j<                               ty         |&j<                        }(t/        |(t        jz                  j|                        r?Ij&                  j                  |&j<                        })|(Ij                  |&j<                  <   n |'       })n |'       })HF   j*                  ju                         |)_        |)Ijp                  HF   <   " GI_;          j&                  j(                  D ]  5t        5d,      sJ5j                     IIjp                  Et        j                  j&                  j_                  5j@                  Efd-      }*t        j                  j&                  j_                  5j`                  Efd.      }+5j:                  d/vr5j<                  },ncty         5j<                        }-5j<                  j                  d0d1      },|-Ij                  |,<   | Ij                   d0|, }.5j<                  ||.<   t/        |*t              stC        d2tG        |*             t/        |+t              stC        d3tG        |+             r5j\                  nd}/Ij&                  js                  5j:                  |,|*|+5jF                  |/4      }#5j*                  ju                         |#_        |#Ijp                  5<    |fD ]  } t        |       D ]  5| 5   }!t?        |!      d
k(  rtC        d5      |!dd D ]  }"Jt        |"         I|5   }0|0tC        d6      Ij&                  js                  |0j:                  |0j<                  Ijp                  5   fi |0jF                  (      }#|0j*                  ju                         |#_           i }1i Bt        j                  j&                  j                         Ci }2|s) j&                  j(                  D ]  5 |	5B|2      \  B}2 n* j&                  j(                  D ]  55|15j\                  <    |s|n|}3t%               }4 j&                  j(                  D 5cg c]  }5|5j:                  d7k(  s|5 }6}5|3D ]  }J|   It        HIfd8Ij                  D              }7t?        |7      dk(  r!|sIj&                  j                  |7d
          nIj&                  j                  |7       |rtIjv                  D 8cg c]  }8|8|6vr|1|8    }9}8|6D ]%  55|4v r |	5B|2      \  B}:|4jK                  5       ' |9D ]%  55|4v r |	5B|2      \  B}2|4jK                  5       ' t        Ij                  Ij&                        |2Ij                  <   Cj                  Ij                  t        Bfd9Ijv                  D                    };t?        Ij                        }<|<dkD  s|<dk(  r\|rZt        j                  j                  j                  |;      }=t        Ij                        D ]  \  }>}?|=|>   j4                  B|?<    |<dk(  s|;Bt        t        Ij                              <    |r*Bs( j&                  j(                  D ]  5 |	5B|2      \  B}2  j&                  j(                  D ][  55j:                  dk(  sCj                  t        j                  j&                  j_                  5j@                  d
   Bfd:             ] t        |2C      }@JjO                         D Ach c]  }A|Aj                   c}A@j*                  d;<   t         j                  dt        d<|@d             |@S c c}}w c c}}w c c}5w c c}8w c c}Aw )=a  
    Creates subgraphs out of main graph

    Args:
        m (GraphModule): Graph module to split
        root_m (torch.nn.Module): root nn module. Not currently used. Included
            because the root nn module is usually transformed via
            torch.fx._symbolic_trace.symbolic_trace (see example below)
        split_callback (Callable[[Node], int]): Callable function
            that maps a given Node instance to a numeric partition identifier.
            split_module will use this function as the policy for which operations
            appear in which partitions in the output Module.
        qualname_map: Optional[Dict[str, str]]: optional output parameter that returns a
            mapping from new target names in the module after split to old target
            names in the original module.
        keep_original_order: Optional[bool]: keep the original order of the GraphModule
            or use the Topological order of the new constructed GraphModule
        keep_original_node_name: Optional[bool]: If the partitioned graphs should
            have the same node names as the original graph.
        keep_original_input_name: bool: If the partitioned graphs should
            have the same input names as the original graph.
        partition_affix: Optional[str]: If specified, the submodules' names will contain
            the affix, e.g. "submod_<affix>_<idx>".
        tuple_return: bool: If True, submodule outputs are always wrapped in a tuple,
            even when there is only a single output value.  This makes all subgraphs
            conform to the convention expected by ``torch._inductor.compile_fx``.

    Returns:
        GraphModule: the module after split.

    Example:

        This is a sample setup:

            import torch
            from torch.fx._symbolic_trace import symbolic_trace
            from torch.fx.graph_module import GraphModule
            from torch.fx.node import Node
            from torch.fx.passes.split_module import split_module

            class MyModule(torch.nn.Module):
                def __init__(self) -> None:
                    super().__init__()
                    self.param = torch.nn.Parameter(torch.rand(3, 4))
                    self.linear = torch.nn.Linear(4, 5)

                def forward(self, x, y):
                    z = self.linear(x + self.param).clamp(min=0.0, max=1.0)
                    w = self.linear(y).clamp(min=0.0, max=1.0)
                    return z + w

            # symbolically trace model
            my_module = MyModule()
            my_module_traced = symbolic_trace(my_module)

            # random mod partitioning
            partition_counter = 0
            NPARTITIONS = 3

            def mod_partition(node: Node):
                global partition_counter
                partition = partition_counter % NPARTITIONS
                partition_counter = (partition_counter + 1) % NPARTITIONS
                return partition

            # split module in module with submodules
            module_with_submodules = split_module(
                my_module_traced, my_module, mod_partition
            )

        Output looks like this. Original graph is broken into partitions

            > print(module_with_submodules)
            GraphModule(
                (submod_0): GraphModule(
                    (linear): Linear(in_features=4, out_features=5, bias=True)
                )
                (submod_1): GraphModule(
                    (linear): Linear(in_features=4, out_features=5, bias=True)
                )
                (submod_2): GraphModule()
            )

            def forward(self, x, y):
                param = self.param
                submod_0 = self.submod_0(x, param, y);  x = param = y = None
                getitem = submod_0[0]
                getitem_1 = submod_0[1];  submod_0 = None
                submod_1 = self.submod_1(getitem, getitem_1);  getitem = getitem_1 = None
                getitem_2 = submod_1[0]
                getitem_3 = submod_1[1];  submod_1 = None
                submod_2 = self.submod_2(getitem_2, getitem_3);  getitem_2 = getitem_3 = None
                return submod_2

        Output of split module is the same as output of input traced module.
        This is an example within a test setting:

            > orig_out = my_module_traced(x, y)
            > submodules_out = module_with_submodules(x, y)
            > self.assertEqual(orig_out, submodules_out)
            True
    z%szpre split_moduleT)colorednodebase_mod_envbase_mod_attrsr   c                    | j                   dk(  r t        | j                        dkD  r| j                  d   nt        j                  j
                  }rX|t        j                  j
                  u rdn|f}j                  d| j                  || j                        || j                  <   n5j                  | j                  | j                  |      || j                  <   | j                  j                         || j                     _        ||fS | j                   dk(  rj                  | j                        || j                  <   | j                  j                         || j                     _        t        | j                  t              s!t!        dt        | j                               t#        | j                        }||| j                  <   ||fS )Nplaceholderr   r,   )args	type_expr)rI   default_valueget_attrExpected str target, got )oplenrH   inspect	Signatureemptycreate_noder   typerG   targetmetacopyrK   
isinstancer+   AssertionErrorr7   )	rC   rD   rE   rJ   rH   r5   base_mod_graphr?   r:   s	         r#   construct_graphz%split_module.<locals>.construct_graph   s   
 77m# #DII 2		!8I8I8O8O  ''7+<+<+B+BBBHX  +9*D*D!II"ii	 +E +TYY' +9*D*DKK"ii"/ +E +TYY'
 ,099>>+;L#( ^++ WW
"&4&=&=dkk&JL#+/99>>+;L#(dkk3/$'@dkkAR@S%TUU.q$++>H*2N4;;'^++r%   r   Ndef_nodeuse_nodec                    ddl m} t        | dd       }t        |dd       }t        j	                  d| j
                  |||j
                  nd|       ||k7  r||G|   }|j                  j                  | j
                         ||j                  j                  |       |/|   }|j                  j                  | j
                         | j                  j                  d      x}t         ||      t              D ]  }|   }	|j                  j                  |	j
                         |   j                  dk7  s@t        |	dd       }
|
P|
   }|j                  j                  |	j
                         |j                  j                  |       |j                  j                  |
        ||j                  j                  |       y y y y )	Nr   free_symbols_fx_partitionz*record_cross_partition_use %s (%s) %s (%s)-example_valuekeyrG   )%torch.fx.experimental.symbolic_shapesr_   r4   logdebugr   r   
setdefaultr   r   rU   getsortedr+   rM   r   )r[   r\   r_   defineduseddef_partitionuse_partitiondef_valss_node	s_defineds_def_partition
partitionssymbol_to_nodes               r#   record_cross_partition_usez0split_module.<locals>.record_cross_partition_use   s   F(OT:x$7		8MM%1HMMs	
 d?" *7 3%%00?#!,,77= *4 0$$//>  (}}00AAGN#L$9sC Q!/!2%,,77D)!,//=@ )0(NI(42<Y2G / 7 7 B B6;; O / : : E Ed K - : : E Ei P#Q$ &!..99'B '3   r%   c                 6    |       }t        |      }dj                  |g      }t        j                  d| j                  |       j                  |      }|t        |      x|<   }|j                  j                  | j                         || _	        y )N_z*instantiate_node_partition_mapping %s (%s))
r+   joinrf   rg   r   ri   r   r   appendr`   )rC   partition_idxpartition_name	partitionr8   rt   r<   s       r#   "instantiate_node_partition_mappingz8split_module.<locals>.instantiate_node_partition_mapping	  s    &t,]+& !XX&GHN		8$))^	

 NN>2	5>~5NNJ~&##DII.+r%   rb   )rG   rK   outputcall_function   z*Expected 1 arg for _set_grad_enabled, got zExpected bool arg, got c              3   >   K   | ]  }t        |t                 y wN)rW   r   .0args     r#   	<genexpr>zsplit_module.<locals>.<genexpr>Y  s     Jz#t44Js   z3Expected all args to be python constants, not Nodesz'Expected 1 arg for _exit_autocast, got c              3   $   K   | ]  }|d u 
 y wr   r,   )r   vs     r#   r   zsplit_module.<locals>.<genexpr>o  s     >q}>s   zautocast must exitzautocast_regions: %szgrad_regions: %s)rG   rK   r   c                      | d       S r   r,   )nrv   s    r#   <lambda>zsplit_module.<locals>.<lambda>  s    (B1d(K r%   zSautocast or set_grad_enabled require monotonically increasing partitions: highest: z, this node's: c                      |       S r   r,   r[   rC   rv   s    r#   r   zsplit_module.<locals>.<lambda>  s    ,FxQU,V r%   c                      |       S r   r,   r   s    r#   r   zsplit_module.<locals>.<lambda>  s    .HSW.X r%   z cycle exists between partitions!z%Expected at least one region for nodec              3       K   | ]  }|  y wr   r,   r   s     r#   r   zsplit_module.<locals>.<genexpr>  s     8ss8s   rM   rT   rH   kwargsrI   rK   c                      r} n
d } dz  j                   j                  |    j                        }d <   |S )Narg_r   rI   )r   rG   rS   )r   rG   counterinpr@   
new_inputs
orig_nodesr}   s     r#   add_placeholderz%split_module.<locals>.add_placeholder  s\    +D "'+DqLG'oo99(o22 :  #'
3""r%   rL   r`   c                     |    S r   r,   r   r    s    r#   r   zsplit_module.<locals>.<lambda>  s    TU r%   c                     |    S r   r,   r   s    r#   r   zsplit_module.<locals>.<lambda>  s    {1~ r%   )call_modulerK   r0   rx   z&Expected tuple for gathered_args, got z'Expected dict for gathered_kwargs, got )rM   rT   rH   r   rI   r   zExpected at least one regionzMissing exit noderG   c              3   B   K   | ]  }j                   |        y wr   )r    )r   r   r   r}   s     r#   r   zsplit_module.<locals>.<genexpr>p  s&      
8<I!!*T"23
s   c              3   (   K   | ]	  }|     y wr   r,   )r   r   rD   s     r#   r   zsplit_module.<locals>.<genexpr>  s     B,t$B   c                 "    | j                      S r   r   )r   rD   s    r#   r   zsplit_module.<locals>.<lambda>  s    |AFF?S r%   partition_nameszpost split_module)Prf   rg   r
   r   dictr+   r   r   graph_moduler   tuplesympyamp_enter_autocast_exit_autocast_C_set_grad_enabledr   setr   nodesrU   ri   rW   SymIntSymFloatrC   exprSymbolrM   rT   rN   rH   rX   boolrS   alladdremovevaluesitemsrj   _LOGGERisEnabledForloggingDEBUGr   map_argr   listkeysr   rz   popr   RuntimeErrorr    rR   rV   r   r7   nnModulerK   r!   r2   r`   replacer   reversedr   r   r   r	   r   proxyProxy	enumeratenextiter)Mr:   r;   r<   r=   r>   r?   r@   r8   r9   rZ   r   r~   GLOBAL_STATE_NODESgrad_regionsautocast_regionsautocast_exitsactive_gradactive_autocastsvals0akr   assert_monotonically_increasinghighest_partitionpidoriginal_partition_orderroot_partitionsr|   sorted_partitionsroot_partition	dependentregions_mappingregionsrnew_nodeplaceholder_inputsget_attr_inputs	orig_noder   	orig_attrrG   gathered_argsgathered_kwargsrT   target_attrr.   r   	exit_nodeorig_mod_envrE   construct_order_partitionsalready_constructed_attr_nodesrC   original_orderoutput_valsrd   orig_mod_attr_nodes_based_mod_attrs
output_valnum_outputsoutput_val_proxyioutput_nameretprD   rY   r   r    r   r   r   r}   rt   rv   ru   sM   ` `  ```                                             `            @@@@@@@@@@@r#   r   r   6   s   h II11dC
!,!,39o!, S%(("7"7"C"CCD!, 
tCIS%((*?*?*K*K%K LL	M	!,F ')J"$J/1N/CT /CTD[ /CT /Cb, ,$ ,. 			!!		  "" 1<L 5@M.0NKu 5: IIMM/22S?3u~~ >?.2=.(,0N388==)77;;*4077o%$++9K*K{{ehh888tyy>Q&(DS^DTU  "$))A,5(+B4		RSCUBV)WXX",/1E0F,G[)		 9 99J		JJ(M  !$$T*),nT.B-C)D &'+t$		 8 88tyy>Q&(A#dii.AQR  !1.22>$3GH ''		!5/3tyy|,"%)).*>?! 	:AQ##N4$89	:i5:n >n&;&;&=>>122 2B1G1G1IJA6!9JJ-9-?-?-ABTQAvayLBLBGMM*,.>?(,7&*+;&<&R\@R#   $
499 771177hHHNN""		!K * &C 3&$,,=+>ocUT  !$ ;;00HHNN""		V HHNN""X9@  $JOO$56!#O%/%5%5%7 3!	9))*"">23
 $&
(,,.  0#N3>> 	2Iy!..22>Bi(55&&y1	2  Z0=>> -l; 7,224 	7MD'7|q $%LMM<@Js71:'33D9 QR[ 7&s1v.	$??66ww;;8dii88"ii 7  IINN$  /7	%%d+7	77. , ;&~.	&(

 )+%'## 	/C"3I
*y//5+Ay/?/?@%((//  &&s+"))#.	/ &7 "	AC"3I#T # # ||z)!)"2"2C8(3D9I9I4J3KL  4Ay7G7GH	i9"+//":":9;K;K"LK:CI%%i&6&67"1"3K-/)#3388:K5@I!!*S/2E"	AF &	w;&|  *34)"4#5#56I $//K!HHNN22499>VWM#hhnn445O ww995aE,,S#6,7	!!&)+ #,"7"7!8&BH-1[[L*mU3$<T-=P<QR  ot4$=d?>S=TU  !8499TD 2277"&)) 3 H !IINN,HM*2I!!$'U*3Z -- _- 	D%d+G7|q $%CDDSb\ &s1v.	*40	$()<==$??66 ||$++#//57'nn 7  NN'') 	. %'L$&L+088>>+?+?+ANCENGGMM 	D+:lN,(L.	 GGMM 	+D&*L#	+ "5:R  &)U" ()ww}}Qt=8PdQNQ4 7E~.	  
@I@Q@Q
 
 {q OO"";q>2OO"";/ %++/n, S!/ / ' 9991@,2.. /22489 , 999/>,0,n /22489 1Cy1
y,,-
 $//!!B1A1ABB


 )++,?{a/L$xx~~33J?"+I,=,="> E;,<Q,?,D,D[)EA:DLd9#4#4567o7Ez <GGMM 	D+:lN,(L.	  77h!!&&tyy|5ST ^^
<C ;E:K:K:M"NQ1=="NCHHII2CF JW
 KBj R"/| #Os+   8AC$AC AC!AC!AC&BAC+c                    t        | j                  j                  dd            }|D ]  }t        |j                        dkD  s|j
                  r)|j                  d   }|j                  j                  d      }|Vg }| j                  j                  |      5  t        |j                        D ]  \  }}t        |t        j                        rl| j                  j                  t        j                  j                   j"                  j$                  ||f      }||j                  d<   |j'                  |       t        |t$              s|j'                  |        	 ddd       t        |j(                        D ]L  }	g }
|	j                  D ])  }||u r|
j+                  |       |
j'                  |       + t-        |
      |	_        N | j                  j/                  |        y# 1 sw Y   xY w)	au  Decompose x.size() into per-dim sym_size.int calls.

    torch.Size objects cannot cross split boundaries because aot_autograd
    cannot handle them as submodule outputs. This replaces each size() call
    with individual sym_size.int(x, dim) nodes:
      - Dynamic dims (SymInt) -> new sym_size.int node
      - Static dims (plain int) -> inlined as literal constant
    call_methodsize)rM   rT   r   r   rb   N)rH   )r   r   
find_nodesrN   rH   r   rU   ri   inserting_afterr   shaperW   r   r   r   opsatensym_sizeintrz   usersextendr   
erase_node)r:   
size_nodesrC   tensor_nodeevdimsr   dim_valdnusernew_argsr   s               r#   _decompose_size_nodesr    s    agg((M&(IJJ !tyy>Aiil!!/2:!#WW$$[1 		)'1 )
7gu||4..		//33;:J / B 07BGGO,KKO-KK()		) $ 	(DHyy )$;OOD)OOC(	)
 hDI	( 	
4 ;!		) 		)s   B2G.G..G7	r   modulesc                    t        dt        fi       }t        j                  |      }|j                  }d|d<   t        t        t        f          |d<   t        t        t        f          |d<   ||nt        t        t        f          |d<   t        t                  |d<   t               |d	<   t               |d
<   d|d<   t               |d<   t               |d<   t               |d<   t               |d<   t               |d<   t               |d<   t               |d<   t               |d<   t               |d<   | |d<   t        t        t        f          |d<   || _
        t        j                  |_        |S )a  Construct a lightweight GraphModule that bypasses expensive init overhead.

    Creates a ``_LazyGraphModule`` instance without going through
    ``GraphModule.__new__`` (MRO traversal + per-instance class creation),
    ``nn.Module.__init__`` (17+ container allocations), or
    ``GraphModule.__init__`` (node iteration + ``recompile()`` codegen).
    Per-instance classes are still created for lazy forward dispatch.

    Only suitable for parameterless graph containers (no parameters, buffers,
    or ``get_attr``/``call_module`` targets to copy from a root module).

    The dict entries below mirror nn.Module.__init__ (torch/nn/modules/module.py).
    Update in lockstep if nn.Module adds or removes instance attributes.
    GraphModuleImplFtraining_parameters_buffersN_modules_non_persistent_buffers_set_backward_pre_hooks_backward_hooks_is_full_backward_hook_forward_hooks_forward_hooks_with_kwargs_forward_hooks_always_called_forward_pre_hooks_forward_pre_hooks_with_kwargs_state_dict_hooks_state_dict_pre_hooks_load_state_dict_pre_hooks_load_state_dict_post_hooks_graphrU   )rS   r   object__new____dict__r   r+   r   r   r   owning_module_lazy_forwardforward)r   r  clsinstds        r#   _make_lite_graph_moduler-    s]   $  #3"5r
:C>>#DAAjMCH~'AmcN$AjM&2GS#X8HAjM'*3xzA#$*}A&=A"&A%-A&1mA"#(3A$%)mA*5-A&'(]A!,A&1mA"#'2}A#$AhKS#X AfIE"00CKKr%   node_to_partitionc                 *   ddl }ddlm i }t        t              }t        t              }i }g }t               }|fdt        dt        t        t        t        t        f   f   dt        fd}	| j                  j                  D ]	  }
|
j                  j                  d      }|Wt        |d	      rKt        |j                  d
      r5|j                  j                  }t!        ||j"                        r	||vr|
||<   |
j$                  dk(  r|
j$                  dk(  r8t&        j(                  j                  j+                  |
j,                  d   |	       ||
   }||
_        ||vrM|j1                  |       |j3                  |       t&        j(                  j                  j5                         ||<   ||||fdt        dt        dt        t        t        t        t        f   f   dt        t        t        t        t        f   f   dt        |j"                  t        f   dt        ffd}t&        j(                  j                  j+                  |
j,                  |       t&        j(                  j                  j+                  |
j6                  |        ||||fS )a  Assign nodes to partitions and detect cross-partition dependencies.

    Returns:
        partition_graphs: Maps partition ID to its Graph.
        partition_inputs: Maps partition ID to ordered dict of input names -> source nodes.
        partition_outputs: Maps partition ID to ordered dict of output names -> source nodes.
        seen_partitions: Partition IDs in order of first appearance.
    r   Nr^   r   _outputsr   c                 p    t        | d      r)|| j                     j                  | j                  |        | S )Nr`   )r2   r`   rh   r   )r   r0  s     r#   _record_output_depz0_detect_dependencies.<locals>._record_output_dep6  s0     1o&Q__%00;r%   rb   rC   r   rG   r   r[   use_pid_inputs_symc                    t        | dd       }||k7  r|||   j                  | j                  |        ||   j                  | j                  |        | j                  j	                  d      }||rt         
|      t              D ]t  }|j	                  |      }|||   j                  |j                  |       |j                  dk7  sFt        |dd       }	|	V||	   j                  |j                  |       v | S )Nr`   rb   rc   rG   )r4   rh   r   rU   ri   rj   r+   rM   )r[   r3  r4  r0  r5  def_pidro   rp   rq   s_pidr_   s             r#   _record_cross_depz/_detect_dependencies.<locals>._record_cross_depU  s     h>G'!&W%00I ++HMM8D"--++O<&4#L$9sC P!%!!>$(33FKKH!995$+FOT$JE$0 ( : :6;; OP Or%   )r   re   r_   r   r   r   r   r  r+   r   r   rU   ri   r2   rC   r   rW   r   rM   r   r   r   rH   r`   rz   r   r   r   )r:   r.  r   partition_graphspartition_inputspartition_outputsru   seen_partitionsseen_partitions_setr2  rC   r   r   r   r9  r_   s                  @r#   _detect_dependenciesr?    sB   " B8:3>t3D4?4E/1N!#O$'E 0AsDdO+, 
  2?iimmO,?wsF3&8QB"ell+.0H%)r"77m#77hHHNN""499Q<1CD% ))""3'##C($)HHNN$8$8$:S! 2B3D-;			 #tCI./	 3S$Y/0		
 u||T)*	 	4 	tyy*;<t{{,=>e2?h -/@/QQr%   r:  r;  r<  r=  c                 <   t        t              }|D ]h  }||   }||   ||   j                         D ]F  \  }}	|j                  ||	j                        }
|	j
                  j                         |
_        |
|	<   H j | j                  j                  D ]  }t        |d      s|j                  }||   }||   t        j                  j                  j                  |j                  j                        }t        j                  j                  j                  |j                   j                        }|j#                  |j$                  |j&                  |||j                        }|j
                  j                         |_        ||<    |D ]  }||   }||   ||   }t)        fd|j+                         D              }t-        |      dk(  r|j/                  |d          Wt-        |      dkD  r|j/                  |       w|j/                  d        |S )zCreate placeholders, clone nodes, and set outputs for each partition.

    Returns:
        partition_env: Maps partition ID to {original_node: cloned_node} dict.
    r   r`   r   c              3   (   K   | ]	  }|     y wr   r,   )r   r   envs     r#   r   z/_clone_nodes_into_partitions.<locals>.<genexpr>  s     OyC	NOr   r   r   r,   )r   r   r   rG   rS   rU   rV   r   r   r2   r`   r   r   r   rH   __getitem__r   rR   rM   rT   r   r   rN   r   )r:   r:  r;  r<  r=  partition_envr   ginp_namer   rG   rC   r   r   r   	out_nodesr   rB  s                    @r#   _clone_nodes_into_partitionsrH  u  s    2=T1BM  )S!C #3C#8#>#>#@ 	)Hi--INN-KK(~~224K(C	N	))  t_-  S!C ..tyy#//J((..00cooN==ww;;"ii ! 
 		(D	%*  
S!C %c*	OI<L<L<NOO{q HH[^$!HH[!HHRL
 r%   c                     t         j                  j                  j                         }i i }| j                  j                  D ]g  }|j
                  dk(  s|j                  |j                  |j                        }	|j                  j                         |	_	        |	|j                  <   i |D ]  }
|	d| d|
 }nd|
 }t        ||
         ||<   t        fd||
   D              }|j                  ||      }t        ||
   j!                               }t#        |      dkD  rOt         j                  j$                  j'                  |      }t)        |      D ]  \  }}||   j*                  |<    t#        |      dk(  s||d   <    | j                  j                  D ][  }|j
                  dk(  s|j-                  t         j                  j                  j/                  |j0                  d   fd	             ] t        ||
      S )zDBuild the outer stitching graph that calls each partition submodule.rG   r   r   rx   c              3   (   K   | ]	  }|     y wr   r,   )r   r   base_envs     r#   r   z)_build_stitching_graph.<locals>.<genexpr>  s     MtHTNMr   r   r   r   c                 "    | j                      S r   r   )r   rK  s    r#   r   z(_build_stitching_graph.<locals>.<lambda>  s    x?O r%   )r  )r   r   r   r   r   rM   rG   rT   rS   rU   rV   r   r-  r   r   r   r   rN   r   r   r   rC   r   r   rH   )r:   r:  r;  r<  r=  r8   
base_graphbase_modulesrC   r   r   r   input_nodesr   	out_namesr   r   r   rK  s                     @r#   _build_stitching_graphrQ    s    %%'J "H+-L $77m#&&)) ' A YY^^%AF"#HTYY$  0&#O#4AcU;K#C5/K$;<LS<Q$R[!M7G7LMM++KE
*3/4467	y>A$xx~~33J?$Y/ :4!1!!4!9!9:^q %/HYq\"!0$  77h&&tyy|5OP #:|DDr%   )r8   c                b    t        | |      \  }}}}t        | ||||       t        | |||||      S )u  Lightweight graph splitter for simple partition patterns.

    A faster alternative to :func:`split_module` for inference-only graphs
    from ``torch.compile``/Dynamo. Because these graphs have no autocast/grad
    regions, no ``get_attr`` nodes, and no non-linear partition dependencies,
    we can skip the topological sort, autocast tracking, and ``get_attr``
    special-casing that ``split_module`` performs. More importantly, we
    construct partition submodules as lightweight ``_LazyGraphModule``
    instances that bypass ``nn.Module.__init__`` and defer ``recompile()``
    codegen — this eliminates the dominant cost when creating 70-100+
    partition submodules for large models.

    Args:
        m: Graph module to split.
        node_to_partition: Maps each operational node to a partition ID.
            Placeholders, get_attr, and output nodes should NOT be included.
        partition_affix: If set, submodule names become
            ``submod_{affix}_{idx}`` instead of ``submod_{idx}``.
    )r?  rH  rQ  )r:   r.  r8   r:  r;  r<  r=  s          r#   r   r     sY    6 	Q 12 K&(9? !	-/@/ "	 r%   )NFFTr   ).rO   r   collectionsr   r   collections.abcr   typingr   r   torch.fx._compatibilityr   torch.fx._lazy_graph_moduler   r	   torch.fx._utilsr
   torch.fx.graph_moduler   torch.fx.noder   __all__	getLoggerr(   rf   r   r   r   r   r+   r7   r  r   r   r   r  r   r   r   r-  r   r   r?  rH  rQ  r   r,   r%   r#   <module>r]     s     0 $   1 L 2 -  ?!!!(+ +g d+
 
 ,
0 C C  d+
 +/',+0%)G
 #'G
G
HHOOG
 dVS[)G
 sCx.4'	G

 G
 "D[G
 #G
 4ZG
 G
 G
 ,G
T)![ )!T )!\ .2*88>>*#{"#d** *ZXRXRD#IXR ehhnn""	"#d39o	d39o	IXRv993 4 4459 3S$Y/09 Cc4i01	9
 #Y9 
#tD$J
 9D #'.E.E3 4 445.E 3S$Y/0.E Cc4i01	.E
 #Y.E 4Z.E .Eb e,
 #'	((D#I( 4Z	(
 ( -(r%   