ó
    Eñi¥  ã                   ó2  • % S SK r S SKrS SKrS SKrS SKJr  S SKJrJr  S SK	J
r
  S SKJrJrJrJr  S SKrS SKJr  S SKJr  S SKJr  S S	KJr  S
SKJr  S
SKJrJr  S
SKJr  S
SK J!r!J"r"  S
SK#J$r$J%r%J&r&J'r'J(r(J)r)  / SQr*S
r+Sr,Sr-Sr.Sr/Sr0Sr1Sr2Sr3\Rh                  Rk                  \2\.5      r6 \Rh                  Rk                  \1S5      r7\S   \8S'    " S S5      r9\" SS9 " S S5      5       r:\" SS9 " S  S!5      5       r;\" SS9 " S" S#5      5       r<\" SS9 " S$ S%\=5      5       r>\" SS9\
 " S& S'5      5       5       r?\" SS9 " S( S)\5      5       r@\" SS9 S2S*\R‚                  R„                  S+\\   S,\\C   S-\DS.\E\C\4   4
S/ jj5       rF " S0 S15      rGg)3é    N)Údefaultdict)ÚIterableÚSequence)Ú	dataclass)ÚAnyÚLiteralÚ
NamedTupleÚOptional)Útrace_structured)Úcompatibility)Úmap_arg)Úget_size_of_nodeé   )ÚFxGraphDrawer)Úget_node_targetÚOperatorSupportBase)Ú	ShapeProp)Ú!move_non_tensor_nodes_on_boundaryÚsplit_by_tags)ÚCALLABLE_NODE_OPSÚFxNetAccFusionsFinderÚis_node_output_tensorÚNodeListÚNodeSetÚTensors)ÚFxNetAccNodesFinderÚFxNetSplitterInternalErrorÚSubgraphÚSplitResultÚgenerate_inputs_for_submodulesÚ	NodeEventÚNodeEventTrackerFÚ_fx_net_trackerz
_nodes.txtz_all.txtÚ FX_NET_ACC_SPLITTER_TRACKER_MODEÚ%FX_NET_ACC_SPLITTER_TRACKER_DUMP_PATHÚ)FX_NET_ACC_SPLITTER_TRACKER_TRACKED_NODESÚ0)r'   Ú1Ú2Ú3ÚTRACKER_MODEc                   ó4   • \ rS rSr\\\SS4S\S\4S jjr	Sr
g)	Ú_SplitterSettingBaseéL   éÿÿÿÿFÚmax_acc_splitsr   c                 óX  • [         R                  " 5       nUR                  SSS[        SS9  UR                  SSS[        SS9  UR                  S	S
SSSS9  UR                  SSSSSS9  UR                  SSSSSS9  UR	                  5       u  pxUR
                  (       a  UR
                  OUU l        UR                  (       a  UR                  OUU l        UR                  (       a  UR                  OUU l        X@l        UR                  (       a  UR                  U l	        g UU l	        g )Nz--min-acc-module-sizez--min_acc_module_sizeFz.Minimum size limit of an accelerator subgraph.)ÚrequiredÚtypeÚhelpz--max-acc-splitsz--max_acc_splitsz,Enforce a maximum number of split subgraphs.z--skip-fusionz--skip_fusionÚ
store_truezËIf true then no fusion groups. Fusion group is used to enforce no non-tensor data flow between submodules. If we don't have this constrain, setting this to false is recommended as it can reduce overhead.)ÚdefaultÚactionr4   z--allow-non-tensorz--allow_non_tensora˜  For some backends non-tensor data flow between cpu and them are not allowed. Therefore, if a node supported by accelerator but it has non-tensor inputs or outputs to a cpu node we would want to consider it as a cpu node during splitting. However, for some backends we might not care about non-tensor data flow and we can set this option to true to disable the functionality that prevent non-tensor data flow.z#--move-non-tensor-nodes-on-boundaryz#--move_non_tensor_nodes_on_boundarya=  AOTI does not support non-tensor nodes on acc->acc, acc->gpu and gpu->acc boundary. For non-tensor nodes on acc->acc boundary and acc->gpu, we move the nodes from upstream to downstream. For non-tensor nodes on gpu->acc boundary, it is handled by the pre-split process. (by method reduce_acc_nodes_non_tensor_input). )r2   r7   r4   )
ÚargparseÚArgumentParserÚadd_argumentÚintÚparse_known_argsÚmin_acc_module_sizeÚskip_fusionÚallow_non_tensorr0   r   )	Úselfr=   r>   r?   r0   r   ÚparserÚargsÚ_unknowns	            ÚZ/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/fx/passes/splitter_base.pyÚ__init__Ú_SplitterSettingBase.__init__M   sg  € ô ×(Ò(Ó*ˆØ×ÑØ#Ø#ØÜØAð 	ñ 	
ð 	×ÑØØØÜØ?ð 	ñ 	
ð 	×ÑØØØØð#ð 	ñ 		
ð 	×ÑØ Ø ØØðVð 	ñ 	
ð 	×ÑØ1Ø1ØØð>ð 	ñ 		
ð  ×0Ñ0Ó2‰ˆð ×'×'ð ×$Ò$à$ð 	Ô ð
 6:×5E×5E ×!1Ò!1È;ˆÔà%)×%:×%:ˆD×!Ò!Ð@Pð 	Ôð $2Ôð ×5×5ð ×2Ñ2ð 	Õ.ð 3ð 	Õ.ó    )r?   r0   r=   r   r>   N)Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__ÚDEFAULT_MIN_ACC_MODULE_SIZEÚDEFAULT_SKIP_FUSIONÚDEFAULT_ALLOW_NON_TENSORr;   ÚboolrE   Ú__static_attributes__© rG   rD   r-   r-   L   s5   † ð 8Ø'Ø1Ø Ø27ñG
ð
 ðG
ð ,0÷G
ð G
rG   r-   )Úis_backward_compatiblec                   ó�   • \ rS rSrSr S
S\R                  R                  S\S\	\R                  R                     4S jjr
S rS	rg)r!   é—   z¬
An event in graph split that happened on a node.
source: Subject of the event
desc: readable description
dep: Optional dependency, usually the node that caused the event.
NÚsourceÚdescÚdepc                 ó(   • Xl         X l        X0l        g ©N)rU   rV   rW   )r@   rU   rV   rW   s       rD   rE   ÚNodeEvent.__init__    s   € ð ŒØŒ	Ø�rG   c                 ó¤   • U R                   R                   SU R                   SU R                  (       a  U R                  R                   3$ S 3$ )Nú: Ú Ú#)rU   ÚnamerV   rW   ©r@   s    rD   Úto_strÚNodeEvent.to_str§   sE   € ð
 —+‘+×"Ñ"Ð# 2 d§i¡i [°À4Ç8Ç8°$·(±(·-±-Ð1UÐVÐVÐQTÐ1UÐVÐVrG   )rW   rV   rU   rY   )rH   rI   rJ   rK   Ú__doc__ÚtorchÚfxÚNodeÚstrr
   rE   ra   rP   rQ   rG   rD   r!   r!   —   sF   † ñð PTñØ—h‘h—m‘mðØ+.ðØ5=¸e¿h¹h¿m¹mÑ5LõõWrG   r!   c                   ó®   • \ rS rSrSrS rSS\R                  R                  S\	S\
\R                  R                     4S jjrSS	 jrS
 rSS jrS rSrg)r"   é¯   z3
Tracks node events during the splitter execution.
c                 óN   • Xl         X l        / U l        0 U l        [        U l        g rY   )Útracker_modeÚdump_prefixÚeventsÚnode_eventsÚprintÚwriter)r@   rk   rl   s      rD   rE   ÚNodeEventTracker.__init__µ   s$   € Ø(ÔØ&ÔàˆŒàˆÔÜˆ�rG   NÚnoderV   rW   c                 ó4  • [        XU5      nU R                  R                  U5        UR                  U R                  ;  a  / U R                  UR                  '   U R                  UR                     R                  [        U R                  5      S-
  5        g)z!
Add a new event to the tracker.
r   N)r!   rm   Úappendr_   rn   Úlen)r@   rr   rV   rW   Úevents        rD   ÚaddÚNodeEventTracker.add¾   ss   € ô ˜$ cÓ*ˆØ�‰×Ñ˜5Ô!Ø�9‰9˜D×,Ñ,Ó,Ø*,ˆD×Ñ˜TŸY™YÑ'Ø×Ñ˜Ÿ™Ñ#×*Ñ*¬3¨t¯{©{Ó+;¸aÑ+?Õ@rG   c                 ó@  • U(       d  U R                   nU R                  R                  U/ 5       Hk  nU R                  U   nU" X6R	                  5       -   5        U(       d  M3  UR
                  c  MB  U R                  UR
                  R                  SSU-   US9  Mm     g)zÚ
Print a node and its events.
@param recursive: if True, print nodes that caused the events on this current node.
@param tab: Indentation for dependencies.
@param writer: function to write to file. If None, use print.
NTz| ©Ú	recursiveÚtabrp   )rp   rn   Úgetrm   ra   rW   Ú
print_noder_   )r@   Ú	node_namer{   r|   rp   Úidxrv   s          rD   r~   ÚNodeEventTracker.print_nodeÈ   s€   € ö Ø—[‘[ˆFØ×#Ñ#×'Ñ'¨	°2Ö6ˆCØ—K‘K Ñ$ˆEÙ�3Ÿ™›Ñ'Ô(ßˆy˜UŸY™YÓ2Ø—‘Ø—I‘I—N‘N¨d¸¸s¹
È6ð  ó ò	 7rG   c                 óÞ   • 0 nU R                    HZ  n/ X'   U R                   R                  U/ 5       H3  nU R                  U   nX   R                  UR	                  5       5        M5     M\     U$ )z!
Create dict dump on all events.
)rn   r}   rm   rt   ra   )r@   Úretr_   r€   rv   s        rD   Úto_dictÚNodeEventTracker.to_dictÙ   sh   € ð %'ˆØ×$Ô$ˆDØˆC‰IØ×'Ñ'×+Ñ+¨D°"Ö5�ØŸ™ CÑ(�Ø‘	× Ñ  §¡£Ö0ó 6ñ %ð
 ˆ
rG   c                 óŒ   • U(       d  U R                   nU R                   H!  nU" SU S35        U R                  USSUS9  M#     g)zZ
Print all nodes in a list.
@param writer: function to write to file. If None, use print.
zNode: Ú:Fz  rz   N)rp   rn   r~   )r@   rp   r_   s      rD   Ú	print_allÚNodeEventTracker.print_allå   sE   € ö
 Ø—[‘[ˆFØ×$Ô$ˆDÙ�V˜D˜6 Ð#Ô$Ø�O‰O˜D¨E°tÀFˆOÓKò %rG   c                 ó\  ^ ^• [        SS U 4S jS9  S mT R                  S:¼  a=  [        T R                  [        -   S5       nT R                  T" U5      5        SSS5        U U4S	 jnT R                  S
:X  d  T R                  S:X  aŒ  T R                  S
:X  a3  [        R                  R                  [        S5      R                  S5      O?T R                  R                  5        VVs/ s H  u  p4[        U5      S:”  d  M  UPM     snnnU" U5        gg! , (       d  f       NÂ= fs  snnf )zm
Function to be invoked at the end of the finder execution to printout tracked events specified by the mode.
Úartifactc                  ó   • SSS.$ )NÚ!fx_net_acc_splitter_finder_eventsÚjson)r_   ÚencodingrQ   rQ   rG   rD   Ú<lambda>Ú'NodeEventTracker.dump.<locals>.<lambda>÷   s   € Ø;Ø"ò!rG   c                  óL   >• [         R                  " T R                  5       5      $ rY   )rŽ   Údumpsr„   r`   s   €rD   r�   r‘   û   s   ø€ œtŸzšz¨$¯,©,«.Ô9rG   )Úmetadata_fnÚ
payload_fnc                 ó   ^ • U 4S jnU$ )Nc                 ó,   >• TR                  U S-   5      $ )NÚ
)Úwrite)ÚxÚfs    €rD   ÚfnÚ2NodeEventTracker.dump.<locals>.writeln.<locals>.fnÿ   s   ø€ Ø—w‘w˜q 4™xÓ(Ð(rG   rQ   )r›   rœ   s   ` rD   ÚwritelnÚ&NodeEventTracker.dump.<locals>.writelnþ   s   ø€ õ)ð ˆIrG   r   ÚwNc           
      óæ   >• [        TR                  [        -   S5       nU  H3  nT" SU S35        TR                  USST" U5      S9  T" SU S35        M5     S S S 5        g ! , (       d  f       g = f)Nr    z===== Tracking node z =====Tz|-rz   z===== End of tracking node )Úopenrl   ÚNODES_SUFFIXr~   )Únodesr›   r   r@   rž   s      €€rD   Údump_selected_nodesÚ2NodeEventTracker.dump.<locals>.dump_selected_nodes
  st   ø€ Ü�d×&Ñ&¬Ñ5°sÔ;¸qÛ!&�IÙÐ2°9°+¸VÐDÔEØ—O‘OØ!¨T°tÁGÈAÃJð $ñ ñ Ð9¸)¸ÀFÐKÖLñ "'÷ <×;Ö;ús   Ÿ:A"Á"
A0é   é   Ú Ú,)r   rk   r¢   rl   Ú
ALL_SUFFIXrˆ   ÚosÚenvironr}   Ú-ENV_FX_NET_ACC_SPLITTER_TRACKER_TRACKED_NODESÚsplitrn   Úitemsru   )r@   r›   r¥   r_   rm   r¤   rž   s   `     @rD   ÚdumpÚNodeEventTracker.dumpð   s  ù€ ô
 	Øñô :ò	
ò	ð ×Ñ Ó!Ü�d×&Ñ&¬Ñ3°SÔ9¸QØ—‘™w q›zÔ*÷ :ö	Mð ×Ñ Ó! T×%6Ñ%6¸!Ó%;ð
 ×$Ñ$¨Ó)ô —
‘
—‘ÔLÈbÓQ×WÑWØôð
 .2×-=Ñ-=×-CÑ-CÔ-EôÚ-E™\˜TÌÈVËÐWXÉ—DÑ-Eòð ñ   Õ&ð &<÷ :Õ9üó(s   ÁDÃ*D(ÄD(Ä
D%)rl   rm   rn   rk   rp   rY   )Fr©   N)rH   rI   rJ   rK   rc   rE   rd   re   rf   rg   r
   rw   r~   r„   rˆ   r±   rP   rQ   rG   rD   r"   r"   ¯   sT   † ñòñA˜Ÿ™Ÿ™ð A¨Sð A°xÀÇÁÇÁÑ7Nõ Aôò"
ô	Lõ/'rG   r"   c                   ó~   • \ rS rSrSrS\R                  R                  S\S\	4S jr
S\4S jrS	 rS
 rS\4S jrSrg)r   i"  a   
Finds a set of nodes that can be supported on ACC, excluding nodes that have non-tensor
input/output to cpu nodes to prevent non-tensor data flow between backends and cpu.

I.e. if we have a chain:

ACC_NODE_1 -> ACC_NODE_2 -> ACC_NODE_3 -> CPU_NODE_1

where every ACC node produces non-tensor output, then they all should be treated as CPU nodes.

This behavior can be turned off by passing allow_non_tensor=True.
ÚmoduleÚoperator_supportr?   c                 óŠ   • Xl         X l        X0l        [        5       U l        [        [        [        5      [        5      U l	        g rY   )
r´   rµ   r?   ÚsetÚ	acc_nodesr"   r;   r+   ÚDUMP_PREFIXÚtracker)r@   r´   rµ   r?   s       rD   rE   ÚFxNetAccNodesFinder.__init__1  s1   € ð ŒØ 0ÔØ 0ÔÜ"%£%ˆŒä'¬¬LÓ(9¼;ÓGˆ�rG   Úcpu_worklistc                 ó~  • U(       a¶  UR                  S5      nUR                   H‹  nX0R                  ;   d  M  U R                  R                  U5        U R                  R                  USU5        [        U5      (       a  M^  U R                  R                  US5        UR                  U5        M�     U(       a  Mµ  gg)zñ
Transitively excludes nodes from ACC supported set.
For every node in the worklist:
- removes its downstream ACC nodes from ACC supported set,
- if any downstream ACC node produces non-tensor output,
  then it gets added into the worklist.
r   zacc_del|user_of_new_cpu_nodeznew_cpu_node|non_tensor_outputN)ÚpopÚusersr¸   Úremoverº   rw   r   rt   )r@   r¼   rr   Úusers       rD   Ú(reduce_acc_nodes_non_tensor_input_helperÚ<FxNetAccNodesFinder.reduce_acc_nodes_non_tensor_input_helper>  s�   € ö Ø×#Ñ# AÓ&ˆDàŸ
œ
�ØŸ>™>Õ)Ø—N‘N×)Ñ)¨$Ô/Ø—L‘L×$Ñ$ TÐ+IÈ4ÔPÜ0°×6Ó6ØŸ™×(Ñ(¨Ð/OÔPØ$×+Ñ+¨DÖ1ñ #÷ ŠlrG   c                 óD  • / nU R                   R                  R                   Hi  nUR                  [        ;  a  M  X R
                  ;   a  M*  [        U5      (       a  M<  U R                  R                  US5        UR                  U5        Mk     U R                  U5        g)zl
Excludes nodes from ACC supported set that have direct
upstream CPU nodes that produce non-tensor outputs.
z&new_cpu_node|callable_non_tensor_inputN)r´   Úgraphr¤   Úopr   r¸   r   rº   rw   rt   rÂ   )r@   Únon_tensor_cpu_nodesrr   s      rD   Ú!reduce_acc_nodes_non_tensor_inputÚ5FxNetAccNodesFinder.reduce_acc_nodes_non_tensor_inputQ  sƒ   € ð
 *,Ðà—K‘K×%Ñ%×+Ô+ˆDØ�w‰wÔ/Ó/ÙØ—~‘~Ó%ÙÜ$ T×*Ñ*ÙØ�L‰L×Ñ˜TÐ#KÔLØ ×'Ñ'¨Ö-ñ ,ð 	×5Ñ5Ð6JÕKrG   c                 óv  •  / nU R                    Hh  n[        U5      (       a  M  UR                   HC  nX0R                   ;  d  M  UR                  U5        U R                  R                  USU5          Mf     Mj     U(       d  gU H  nU R                   R                  U5        M      U R                  U5        M¹  )zf
Excludes nodes from ACC supported set that produce non-tensor
outputs and have downstream CPU nodes.
z'acc_del|non_tensor_output_with_cpu_userN)r¸   r   r¿   rt   rº   rw   rÀ   rÂ   )r@   Únew_cpu_nodesÚacc_noderÁ   Únew_cpu_nodes        rD   Ú"reduce_acc_nodes_non_tensor_outputÚ6FxNetAccNodesFinder.reduce_acc_nodes_non_tensor_outputd  s§   € ð
 Ø&(ˆMà ŸNœN�Ü(¨×2Ñ2ÙØ$ŸNœN�DØ§>¡>Õ1Ø%×,Ñ,¨XÔ6ØŸ™×(Ñ(Ø$Ð&OÐQUôò ó +ñ +ö !Øã -�Ø—‘×%Ñ% lÖ3ñ !.ð ×9Ñ9¸-ÔHñ) rG   Úreturnc                 ó°  • [        U R                  R                  5       5      n[        5       U l        U R                  R
                  R                   Hª  nUR                  [        ;  a  U R                  R                  US5        M5  U R                  R                  X5      (       d  U R                  R                  US5        Ms  U R                  R                  US5        U R                  R                  U5        M¬     U R                  (       d   U R                  5         U R                  5         U R                  R!                  5         U R                  $ )Nzinit_cpu|not_callablezinit_cpu|operator_supportz(init_acc|callable_and_operator_supported)Údictr´   Únamed_modulesr·   r¸   rÅ   r¤   rÆ   r   rº   rw   rµ   Úis_node_supportedr?   rÈ   rÎ   r±   )r@   Ú
submodulesÚns      rD   Ú__call__ÚFxNetAccNodesFinder.__call__  sô   € Ü˜$Ÿ+™+×3Ñ3Ó5Ó6ˆ
Ü›ˆŒØ—‘×"Ñ"×(Ô(ˆAØ�t‰tÔ,Ó,Ø—‘× Ñ  Ð$;Ô<ÙØ×(Ñ(×:Ñ:¸:×IÑIØ—‘× Ñ  Ð$?Ô@Ùà�L‰L×Ñ˜QÐ JÔKØ�N‰N×Ñ˜qÖ!ñ )ð ×$×$Ø×2Ñ2Ô4Ø×3Ñ3Ô5Ø�‰×ÑÔØ�~‰~ÐrG   )r¸   r?   r´   rµ   rº   N)rH   rI   rJ   rK   rc   rd   re   ÚGraphModuler   rO   rE   r   rÂ   rÈ   rÎ   r   r×   rP   rQ   rG   rD   r   r   "  sZ   † ñðHà—‘×$Ñ$ðHð .ðHð ô	Hð2ÀXô 2ò&Lò&Ið6˜'÷ rG   r   c                   ó   • \ rS rSrSrg)r   i”  rQ   N)rH   rI   rJ   rK   rP   rQ   rG   rD   r   r   ”  s   † ârG   r   c                   ó>   • \ rS rSr% \\S'   \\S'   Sr\\	   \S'   Sr
g)r   i™  Úis_accr¤   NÚdevice_ordinalrQ   )rH   rI   rJ   rK   rO   Ú__annotations__r   rÝ   r
   r;   rP   rQ   rG   rD   r   r   ™  s   ‡ ð ƒLØƒOØ$(€N�H˜S‘MÖ(rG   r   c                   ój   • \ rS rSr% Sr\R                  R                  \S'   \	\
\4   \S'   \
\S'   Srg)r   i¡  a&  
Stores the results of the splitter.

Attributes:
    split_module: root module after splitting.
    submodule_inputs: a dict that maps submodule name to its inputs.
    non_acc_submodule_prefix: the prefix for non acc submodules. For
        acc submodule the prefix is always "_run_on_acc_".
Úsplit_moduleÚsubmodule_inputsÚnon_acc_submodule_prefixrQ   N)rH   rI   rJ   rK   rc   rd   re   rÙ   rÞ   rÒ   rg   r   rP   rQ   rG   rD   r   r   ¡  s-   ‡ ñð —(‘(×&Ñ&Ó&Ø˜3 ˜8‘nÓ$Ø!Ö!rG   r   ÚmodelÚinputsÚtarget_submodulesÚdeepcopyrÐ   c                 ó  ^^	^
^• / m	0 m
U R                  5        VVs0 s H  u  pEXT_M	     snnmUU
U4S jnU R                  5        HW  u  pEXB;   d  M  [        U[        R                  R                  5      (       a  M7  T	R                  UR                  U5      5        MY     U	4S jn [        R                  " 5          U " U6   SSS5        U" 5         T
$ s  snnf ! , (       d  f       N= f! [         a  nU" 5         UeSnAff = f)aZ  
Generate inputs for targeting submdoules in the given model. Note that if two submodules refer to the same obj, this
function doesn't work.

Args:
    model: root model.
    inputs: inputs to the root model.
    target_submodules: submodules that we want to generate inputs for.

Returns:
    A dict that maps from submodule name to its inputs.
c                 óP   >• T(       a  [         R                  " U5      OUTTU    '   g rY   )Úcopyræ   )r´   Úmodule_inputsræ   ÚresultsÚsubmodule_to_namess     €€€rD   Úpre_forwardÚ3generate_inputs_for_submodules.<locals>.pre_forwardÊ  s#   ø€ æ,4ŒD�MŠM˜-Ô(¸-ð 	Ð" 6Ñ*Ò+rG   c                  ó8   >• T H  n U R                  5         M     g rY   )rÀ   )ÚhÚhandless    €rD   Úclean_up_handlesÚ8generate_inputs_for_submodules.<locals>.clean_up_handlesÔ  s   ø€ ÛˆAØ�H‰HŽJò rG   N)	rÓ   Ú
isinstancerd   ÚjitÚScriptModulert   Úregister_forward_pre_hookÚno_gradÚ	Exception)rã   rä   rå   ræ   r_   Úmodrí   rò   Úerñ   rë   rì   s      `     @@@rD   r    r    ²  s×   û€ ð( €GØ€GØ5:×5HÑ5HÔ5JÔKÒ5J©	¨˜#š)Ñ5JÒKÐ÷
ð
 ×(Ñ(Ö*‰	ˆØÕ$Ü˜c¤5§9¡9×#9Ñ#9×:Ó:Ø—‘˜s×<Ñ<¸[ÓIÖJñ +õ
ðÜ�]Š]�_Ù�6‰N÷ ñ ÔØ€Nùó1 L÷" �_ûäó ÙÔØˆûðús;   œCÂ(C+ Â=CÃC+ Ã
C(Ã$C+ Ã(C+ Ã+
DÃ5	C>Ã>Dc                   óp  • \ rS rSrSrSr   S+S\R                  R                  S\	\
   S\S\S	\S
\S\\   4S jjrS\\\4   4S jrS\\R                  R(                  \4   4S jrS rS\R                  R                  S\S\R2                  R4                  4S jrS\R                  R                  S\S\4S jrS\R                  R                  S\4S jrS,S\4S jjrS,S\4S jjr  S-S\\!   S\\R                  R(                  \4   4S jjr"S\\R                  R(                  \4   4S jr#S\S\4S jr$S\4S  jr%S\&\\4   4S! jr'S\(\)   4S" jr*S#\(\)   S\(\)   4S$ jr+S#\(\)   4S% jr,S,S&\S\R                  R                  4S' jjr-S\R                  R                  4S( jr.S\/4S) jr0S*r1g).Ú_SplitterBaseiã  a�  
Splits a GraphModule into sub-GraphModules for execution on CPU or the accelerator.
Output is a GraphModule with supported and unsupported operators grouped into as few sub-GraphModules as possible.
Assumes that only "call_module", "call_function" and "call_method" from FX IR can potentially be executed on the accelerator.

Given the following graph:
      ==> b ==>
    //         \
   a             d
    \         //
      ==> c ==>

class SimpleModule(torch.nn.Module):
    def forward(self, a):
        b = torch.sin(a)
        c = torch.cos(a)
        d = b + c
        return d

and providing "operator_support" that indicates that 'b' and 'c' can be executed on the accelerator,
we will get the following split result:

main:
def forward(self, a):
    run_on_acc_0_0 = self._run_on_acc_0_0(a)
    getitem = run_on_acc_0_0[0]
    getitem_1 = run_on_acc_0_0[1]
    run_on_cpu_1_1 = self._run_on_cpu_1_1(getitem, getitem_1)
    return run_on_cpu_1_1

_run_on_acc_0_0:
def forward(self, a):
    sin_1 = torch.sin(a)
    cos_1 = torch.cos(a)
    return (sin_1, cos_1)

_run_on_cpu_1_1:
def forward(self, sin_1, cos_1):
    add_1 = sin_1 + cos_1
    return add_1
l       d Nr´   Úsample_inputrµ   ÚsettingsÚnon_acc_submodule_nameÚreturn_tupleÚnodes_finderc                 ó‚  • [        U[        R                  R                  5      (       d  [	        S[        U5       35      eXl        [        U R                  5      R                  " U6   X@l	        X0l
        X l        Uc5  [        U R                  U R                  U R                  R                  5      nU" 5       U l        U R                  R                  (       a  0 U l        O[#        XR                  5      " 5       U l        U R%                  5       U l        U R)                  5         XPl        0 U l        X`l        / U l        g)a  
Preprocesses graph before splitting:
- finds nodes supported by ACC,
- finds fusion groups for ACC nodes having non-tensor IO,
- builds a graph of direct dependencies,
- builds a map of fused nodes to their fusions.
As a result we get self.acc_nodes, self.deps and self.fusions.
zExpected GraphModule, got N)rô   rd   re   rÙ   ÚAssertionErrorr3   r´   r   Ú	propagaterÿ   rµ   rþ   r   r?   r¸   r>   Úfusionsr   Ú	find_depsÚdepsÚupdate_deps_for_fusionsr   Ú_node_submodule_mapÚ_return_tupleÚtags)r@   r´   rþ   rµ   rÿ   r   r  r  s           rD   rE   Ú_SplitterBase.__init__  só   € ô$ ˜&¤%§(¡(×"6Ñ"6×7Ñ7Ü Ð#=¼dÀ6»l¸^Ð!LÓMÐMàŒÜ�$—+‘+Ó×(Ò(¨,Ñ7à ŒØ 0ÔØ(ÔØÑÜ.Ø—‘˜T×2Ñ2°D·M±M×4RÑ4RóˆLñ &›ˆŒà�=‰=×$×$ØˆD�Lä0°¿¹ÔHÓJˆDŒLð —N‘NÓ$ˆŒ	Ø×$Ñ$Ô&à&<Ô#Ø35ˆÔ Ø)Ôà!ˆ�	rG   rÐ   c                 ó   • U R                   $ )zïReturns a map from node name to submodule name, e.g.
node: main_module_impl_impl_over_arch_unary_multiple_embedding
  _pooling_embedding_pooling_sparse_entity_equivalence_key
  _proxy_embedding_bag
maps to submodule name of: _run_on_acc_1
)r
  r`   s    rD   Úget_node_submodule_mapÚ$_SplitterBase.get_node_submodule_mapE  s   € ð ×'Ñ'Ð'rG   c                 ó  • [        [        5      nU R                  R                  R                   HQ  nUR
                  [        ;  a  M  UR                   H(  nUR
                  S:w  d  M  X   R                  U5        M*     MS     U$ )zá
Builds a graph of node dependencies. Leaf nodes don't have any
dependencies and the "output" node doesn't have nodes depending on it.

Resulting graph has only direct dependencies, i.e. there are no
transitive dependencies.
Úoutput)	r   r·   r´   rÅ   r¤   rÆ   r   r¿   rw   )r@   r  rr   rÁ   s       rD   r  Ú_SplitterBase.find_depsN  sg   € ô .9¼Ó-=ˆØ—K‘K×%Ñ%×+Ô+ˆDØ�w‰wÔ/Ó/ÙàŸ
œ
�Ø—7‘7˜hÕ&Ø‘J—N‘N 4Ö(ó #ñ	 ,ð ˆrG   c                 ó&  • U R                    H�  nU R                   U   nU Hi  nU R                  U   R                  U R                  U   U-
  5        UR                   H(  nXB;  d  M
  U R                  U   R	                  U5        M*     Mk     Mƒ     g)z´
Updates graph of dependencies so that:
- nodes from the same fusion depend on the same set of outer nodes,
- outer nodes depending on a fusion depend on all nodes in that fusion.
N)r  r  Úupdater¿   rw   )r@   rr   ÚfusionÚfused_neighborrÁ   s        rD   r	  Ú%_SplitterBase.update_deps_for_fusions`  sz   € ð —L”LˆDØ—\‘\ $Ñ'ˆFÛ"(�Ø—	‘	˜$‘×&Ñ& t§y¡y°Ñ'@À6Ñ'IÔJà*×0Ô0�DØÕ)ØŸ	™	 $™×+Ñ+¨DÖ1ó 1ó #)ò !rG   rú   rä   c                 ó   • U$ )z
Lower the model to a backend.
rQ   ©r@   rú   rä   s      rD   Ú_lower_model_to_backendÚ%_SplitterBase._lower_model_to_backends  s	   € ð ˆ
rG   c                 ó   • g)zŒ
When an error occurs during lowering or running the lowered mod, we use this
function to find culprits in the `mod` that causes the error.
zMUnable to find a culprit because _find_culprit() function is not implemented.rQ   r  s      rD   Ú_find_culpritÚ_SplitterBase._find_culprit|  s   € ð _rG   Úsupported_nodesc                 óŒ   ^^• SSSS.m " UU4S jS[         5      nU" USSS	9nUR                  5       nUR                  S
5        g )NÚ	AliceBlueÚchartreuse1Úcrimson)r6   Ú	supportedÚunsupportedc                   ó.   >^ • \ rS rSrU UU4S jrSrU =r$ )ÚE_SplitterBase._draw_graph_based_on_node_support.<locals>.CustomDraweri�  c                 ó’   >• [         TU ]  U5      nUT;   a
  TS   US'   U$ UR                  [        ;   a
  TS   US'   U$ TS   US'   U$ )Nr%  Ú	fillcolorr&  r6   )ÚsuperÚ_get_node_stylerÆ   r   )r@   rr   ÚtemplateÚ	__class__Ú	color_mapr   s      €€€rD   r,  ÚU_SplitterBase._draw_graph_based_on_node_support.<locals>.CustomDrawer._get_node_styleŽ  sk   ø€ Ü ™7Ñ2°4Ó8�Ø˜?Ó*Ø,5°kÑ,B�H˜[Ñ)ð  �ð —W‘WÔ 1Ó1Ø,5°mÑ,D�H˜[Ñ)ð  �ð -6°iÑ,@�H˜[Ñ)à�rG   rQ   )rH   rI   rJ   rK   r,  rP   Ú__classcell__)r.  r/  r   s   @€€rD   ÚCustomDrawerr(  �  s   ù† ÷	 õ 	 rG   r2  Únode_supportT©Úignore_getattrznode_support.dot)r   Úget_main_dot_graphÚ	write_raw)r@   rú   r   r2  ÚdrawerÚ	dot_graphr/  s     `   @rD   Ú!_draw_graph_based_on_node_supportÚ/_SplitterBase._draw_graph_based_on_node_support„  sS   ù€ ð #Ø&Ø$ñ
ˆ	÷
	 ð 
	 œ=ô 
	 ñ ˜c >À$ÑGˆØ×-Ñ-Ó/ˆ	à×ÑÐ.Õ/rG   Ú
dump_graphc           
      óÀ  ^• [        U R                  R                  5       5      n/ n[        [        5      n[        [        5      nS mU R                  R
                  R                   GHC  nUR                  [        ;  a  M  [        X&5      nUR                   Vs/ s H6  n[        U[        R                  R                  5      (       a  T" U5      OS PM8     n	n[        U	5      [!        S [#        [%        U	5      5       5       [        U	5      5      -
  n
['        U	S U
 5      n['        U4S jUR(                  R+                  5        5       5      nU R,                  R/                  X&5      (       a(  UR1                  U5        XG   R3                  X¼45        GM/  XW   R3                  X¼45        GMF     U(       a  U R5                  U R                  U5        SnUR+                  5        H&  u  pïU H  u  p¼XÞ SU S[        U5       S3-  nM     M(     US-  nUR+                  5        H&  u  pïU H  u  p¼XÞ SU S[        U5       S3-  nM     M(     [7        U5        U$ s  snf )	Nc                 óR   • U R                   R                  S5      n[        USS 5      $ )NÚtensor_metaÚdtype)Úmetar}   Úgetattr)Úargr?  s     rD   Ú	get_dtypeÚ5_SplitterBase.node_support_preview.<locals>.get_dtype¥  s#   € ØŸ(™(Ÿ,™, }Ó5ˆKÜ˜;¨°Ó6Ð6rG   c              3   ó4   #   • U  H  u  pUc  M
  Uv •  M     g 7frY   rQ   )Ú.0Úir@  s      rD   Ú	<genexpr>Ú5_SplitterBase.node_support_preview.<locals>.<genexpr>·  s   é € ð â$C™˜Ø÷ ‘AÚ$Cùs   ‚	�	c              3   ó’   >#   • U  H<  u  p[        U[        R                  R                  5      (       d  M0  UT" U5      4v •  M>     g 7frY   )rô   rd   re   rf   )rG  ÚkrC  rD  s      €rD   rI  rJ  Á  s7   øé € ð 'â1‘F�AÜ˜c¤5§8¡8§=¡=×1ó $�‘I˜c“NÕ#Ú1ùs
   ƒ/A¶Az$
Supported node types in the model:
z: (z, z)
z&
Unsupported node types in the model:
)rÒ   r´   rÓ   r   r·   rÅ   r¤   rÆ   r   r   rB   rô   rd   re   rf   ru   ÚnextÚ	enumerateÚreversedÚtupleÚkwargsr°   rµ   rÔ   rt   rw   r:  ro   )r@   r<  rÕ   r   Úsupported_node_typesÚunsupported_node_typesrr   ÚtargetrC  Ú
arg_dtypesÚ
last_indexÚarg_dtypes_tupleÚkwarg_dtypes_tupleÚreportsÚtÚdtypesrD  s                   @rD   Únode_support_previewÚ"_SplitterBase.node_support_previewž  sI  ø€ Ü˜$Ÿ+™+×3Ñ3Ó5Ó6ˆ
à$&ˆÜ*¬3Ó/ÐÜ!,¬SÓ!1Ðò	7ð —K‘K×%Ñ%×+Õ+ˆDØ�w‰wÔ/Ó/Ùä$ ZÓ6ˆFð
  Ÿ9š9óâ$�Cô #-¨S´%·(±(·-±-×"@Ñ"@‘	˜#”ÀdÒJÙ$ð ð ô ˜Z›¬4ñä$-¬h°zÓ.BÔ$Cóô
 �J“ó,ñ ˆJô  % Z°°Ð%<Ó=ÐÜ!&ô 'à"Ÿk™k×/Ñ/Ô1ó'ó "Ðð ×$Ñ$×6Ñ6°z×HÑHØ×&Ñ& tÔ,Ø$Ñ,×0Ñ0Ð2BÐ1W×Xà&Ñ.×2Ñ2Ø%Ð:÷ñE ,öL Ø×2Ñ2°4·;±;ÀÔPà:ˆØ-×3Ñ3Ö5‰IˆAÛ8>Ñ4Ð Ø˜S Ð$4Ð#5°R¼Ð=OÓ8PÐ7QÐQTÐUÑU’ó 9?ñ 6ð 	Ð=Ñ=ˆØ/×5Ñ5Ö7‰IˆAÛ8>Ñ4Ð Ø˜S Ð$4Ð#5°R¼Ð=OÓ8PÐ7QÐQTÐUÑU’ó 9?ñ 8ô 	ˆgŒð ˆùò_s   Â=Ic                 óø  ^^^• SmU R                  5       n[        U Vs/ s H  o3R                  (       d  M  UPM     sn5      n[        U5      U-
  nTS[        U5       S3-  mTSU SU S3-  mU R                  U5      n[        U Vs/ s H  o3R                  (       d  M  UPM     sn5      n[        U5      U-
  nTS[        U5       S3-  mTSU SU S3-  m[	        U5       HK  u  pgTUR                  (       a  SU S	3OU R
                   U S	3-  mT[        UR                  5       S
3-  mMM     U R                  U5        U R                  SS9nUR                  5         U(       aH  [        USSS9n	U	R                  5       n
U
R                  5        H  u  p¼UR                  U S35        M     U R                  nSnUR                  R                   GHž  nUR                   S:X  d  M  SUR"                  ;   d  M(  TSUR"                   S3-  m[%        X�R"                  5      mS nU" UTU R&                  5      n[)        T5      R*                  " U6   SnSmTS-  mTR                  R                   H]  nUR                   S:X  a6  [-        U5      (       d  TSUR.                   S3-  mOU[1        TU5      S   -  nUR                   S:X  d  M[  UnM_     TS-  mS[2        R4                  R6                  4UUU4S jjn[9        WR:                  U5        U R                  [=        UT5      -  nTSU ST S 3-  mTS!U S"3-  mUU:  a  UnUR"                  n U R?                  TU5      n U" U6   TS$-  mGM¡     TS&U S 3-  mTS'U S(3-  m[E        T5        T$ s  snf s  snf ! [@         a    TS#-  mTU RC                  TU5      -  m GMô  f = f! [@         a    TS%-  mTU RC                  TU5      -  m GM   f = f))Nr©   z+Before removing small acc subgraphs, total z subgraphs are created:r]   ú acc subgraphs and z cpu subgraphs.
z*After removing small acc subgraphs, total Ú_run_on_acc_r\   z	 node(s)
T)Ú
remove_tagÚpreviewr4  z.dotÚcall_moduleÚaccz
Processing acc submodule r˜   c                 ód   ^• S mU4S jnUR                  U5      nU " U6   UR                  5         T$ )Nc                 ó
   >• Umg rY   rQ   )r@   rä   Ú
sub_inputss     €rD   Ú
get_inputsÚJ_SplitterBase.split_preview.<locals>.get_submod_inputs.<locals>.get_inputs  s   ø€ à%+™
rG   )r÷   rÀ   )Úmain_modÚsubmodÚexample_inputsrh  Úhandlerg  s        @rD   Úget_submod_inputsÚ6_SplitterBase.split_preview.<locals>.get_submod_inputs  s6   ø€ Ø!%�Jõ,ð $×=Ñ=¸jÓI�FÙ˜nÑ-Ø—M‘M”OØ%Ð%rG   r   zChecking inputs...
ÚplaceholderzInput ú= is not a tensor, this might cause problems during lowering!
r  zChecking outputs...
rr   c                 ór   >• [        U 5      (       d  TSU R                   S3-  mg T[        TU 5      S   -  mg )NzOutput rq  r   )r   r_   r   )rr   rY  rk  Útotal_output_bytess    €€€rD   Ú	get_bytesÚ._SplitterBase.split_preview.<locals>.get_bytes)  s@   ø€ ô 1°×6Ñ6Ø W¨T¯Y©Y¨KÐ7uÐ#vÑv™à*Ô.>¸vÀtÓ.LÈQÑ.OÑOÑ*rG   zTotal input size in bytes is z , total output size in bytes is rª   zF theoretical max qps (bounds by PCIe bandwidth) for this submodule is z.
z#Run into an error during lowering!
zLowering and running succeed!
z$Run into an error during inference!
zB
Theoretical max qps (bounds by PCIe bandwidth) for this model is z bottleneck is submodule Ú.)#Úput_nodes_into_subgraphsru   rÜ   Úremove_small_acc_subgraphsrN  r   r¤   Útagr¯   Úevalr   Úget_all_dot_graphsr°   r7  ÚPCIe_BWrÅ   rÆ   rT  rB  rþ   r   r  r   r_   r   rd   re   rf   r   rB   Úmaxr  ÚRuntimeErrorr  ro   )r@   r<  Ú	subgraphsÚgÚacc_subgraphs_numÚcpu_subgraphs_numrH  ÚsubgraphÚ	split_modr8  Ú
dot_graphsr_   r9  Úmax_qpsÚbottleneck_modulerr   rn  Úsubmod_inputsÚtotal_input_bytesrÖ   Úoutput_nodert  ÚqpsÚlowered_submodrY  rk  rs  s                           @@@rD   Úsplit_previewÚ_SplitterBase.split_previewá  ss  ú€ ØˆØ×1Ñ1Ó3ˆ	Ü©IÓ BªI q¿½§©IÑ BÓCÐÜ 	›NÐ->Ñ>ÐØÐ@ÄÀYÃÐ@PÐPgÐhÑhˆØ�QÐ(Ð)Ð)<Ð=NÐ<OÐO`ÐaÑaˆà×3Ñ3°IÓ>ˆ	Ü©IÓ BªI q¿½§©IÑ BÓCÐÜ 	›NÐ->Ñ>ÐØÐ?ÄÀIÃÐ?OÐOfÐgÑgˆØ�QÐ(Ð)Ð)<Ð=NÐ<OÐO`ÐaÑaˆä$ YÖ/‰KˆAØà—?—?ð ˜q˜c Ñ$à×3Ñ3Ð4°Q°C°rÐ:ñˆGð
 œ#˜hŸn™nÓ-Ð.¨jÐ9Ñ9ŠGñ 0ð 	�‰�ÔØ—J‘J¨$�JÐ/ˆ	Ø�‰ÔæÜ" 9¨iÈÑMˆFØ×2Ñ2Ó4ˆJØ#-×#3Ñ#3Ö#5‘�à×#Ñ# t f¨D MÖ2ñ $6ð Ÿ™ˆØÐà—O‘O×)Õ)ˆDØ�w‰w˜-Õ'¨E°T·[±[Õ,@ØÐ8¸¿¹¸ÀRÐHÑH�ä  ¯K©KÓ8�ò
&ñ !2°)¸VÀT×EVÑEVÓ W�Ü˜&Ó!×+Ò+¨]Ñ;à$%Ð!Ø%&Ð"àÐ1Ñ1�ØŸ™×+Ô+�AØ—t‘t˜}Ó,Ü4°Q×7Ñ7Ø#¨°·±¨xÐ7uÐ'vÑv™Gà-Ô1AÀ&È!Ó1LÈQÑ1OÑOÐ-Ø—t‘t˜xÕ'Ø&'šñ ,ð Ð2Ñ2�ðP¤E§H¡H§M¡M÷ Pñ Pô ˜×(Ñ(¨)Ô4Ø—l‘l¤SÐ):Ð<NÓ%OÑO�ØÐ:Ð;LÐ:MÐMmð  oAð  nBð  BCð  Dñ  D�ØÐcÐdgÐchÐhkÐlÑl�à˜“=Ø!�GØ(,¯©Ð%ðØ%)×%AÑ%AÀ&È-Ó%X�NðAÙ" MÑ2ð
 Ð@Ñ@“GñE *ðH 	ÐXÐY`ÐXaÐabÐcÑcˆØÐ.Ð/@Ð.AÀÐCÑCˆÜˆgŒð ˆùòU !Cùò !Cøôd $ó ØÐEÑE�GØ˜t×1Ñ1°&¸-ÓHÑH�GÛðûô $ó IØÐFÑF�GØ˜t×1Ñ1°&¸-ÓHÑH”GðIús:   ŸN·NÂ
NÂ"NÍN$Í,OÎ$$OÏOÏ$O9Ï8O9Útag_idc                 óv  • [        [        5      nU R                  R                  R                   H…  nUR
                  [        ;  a  M  UR                   H\  nUR
                  [        ;  a  M  Ub-  [        UR                  R                  S5      S   5      U:  d  MI  X#   R                  U5        M^     M‡     U$ )z“
Builds reversed topological node dependencies, if tag_id is specified,
we ignore nodes that are in later subgraph i.e. nodes have greater tag_id.
Ú_r/   )r   r·   r´   rÅ   r¤   rÆ   r   r¿   r;   ry  r¯   rw   )r@   r�  Úresultrr   rÁ   s        rD   Úfind_reverse_depsÚ_SplitterBase.find_reverse_depsT  s�   € ô 0;¼3Ó/?ˆà—K‘K×%Ñ%×+Ô+ˆDØ�w‰wÔ/Ó/ÙàŸ
œ
�Ø—7‘7Ô"3Ó3Ùà‘>¤c¨$¯(©(¯.©.¸Ó*=¸bÑ*AÓ&BÀVÕ&KØ‘L×$Ñ$ TÖ*ó #ñ	 ,ð ˆrG   r  c                 óp  • [        5       nU R                  R                  5        HŽ  u  p4X2;   a  M  [        5       nU H  nUR                  X   5        M     UR	                  U5        U HE  nXQU'   UR
                   H  nXt;  d  M
  X   R                  U5        M     UR                  U5        MG     M�     g rY   )r·   r  r°   r  Údifference_updateÚall_input_nodesrw   )r@   r  Úprocessed_noderr   r  Únew_deprÖ   rC  s           rD   Úupdate_reverse_deps_for_fusionsÚ-_SplitterBase.update_reverse_deps_for_fusionsj  s¥   € Ü›ˆà ŸL™L×.Ñ.Ö0‰LˆDØÓ%Ùä“eˆGó �Ø—‘˜t™wÖ'ñ ð ×%Ñ% fÔ-ó �Ø!�Q‘à×,Ô,�CØÕ(Ø™	×(Ñ(¨Ö0ñ -ð ×"Ñ" 1Ö%ó ò 1rG   ry  c                 óP  • [        5       nU R                  R                  R                   Hw  nUR                  [
        ;   d  M  UR                  U:X  d  M+  UR                   H<  nUR                  [
        ;   d  M  UR                  U:w  d  M+  UR                  U5        M>     My     U$ )zÏ
Finds parent nodes of the `tag` subgraph.

Traverse the inputs of nodes in the subgraph, if input doesn't belong to the subgraph
and is not a placeholder, we consider it as the parent node of the subgraph.
)	r·   r´   rÅ   r¤   rÆ   r   ry  r—  rw   )r@   ry  Úparent_nodesrr   rC  s        rD   Úfind_parent_nodes_of_subgraphÚ+_SplitterBase.find_parent_nodes_of_subgraph…  sy   € ô “uˆà—K‘K×%Ñ%×+Ô+ˆDØ�w‰wÔ+Õ+°·±¸CµØ×/Ô/�CØ—v‘vÔ!2Õ2°s·w±wÀ#µ~Ø$×(Ñ(¨Ö-ó 0ñ ,ð ÐrG   c           	      ót  • U R                  [        UR                  SSS9S   5      S9nU R                  U5        U R	                  U5      n[        5       nU(       aÜ  SnU H   nX&   U::  d  M  X`R                  ;   d  M  Un  O   Uc  gXl        UR                  U5        UR                  U5        XPR                  ;   a.  U R                  U    H  nXt;  d  M
  UR                  U5        M     UR                   H1  nUR                  [        ;   d  M  X„;  d  M   UR                  U5        M3     U(       a  MÛ  gg)zN
Extend the acc subgraph with `tag` going the reversed topological direction.
r‘  r   )Úmaxsplitr/   )r�  N)r“  r;   Úrsplitrš  rž  r·   r¸   ry  rÀ   rw   r  r—  rÆ   r   )	r@   ry  r  r�  Úvisited_nodesrr   rÖ   Úfusion_noderC  s	            rD   Úextend_acc_subgraphÚ!_SplitterBase.extend_acc_subgraph–  s  € ð ×%Ñ%¬S°·±¸CÈ!°Ð1LÈRÑ1PÓ-QÐ%ÐRˆØ×,Ñ,¨TÔ2ð ×9Ñ9¸#Ó>ˆä!$£ˆæØˆDó "�Ø‘7˜mÕ+°·^±^Õ0CØ�DÙñ "ð
 ‰|Øð ŒHØ×Ñ Ô%Ø×Ñ˜dÔ#ð —|‘|Ó#Ø#'§<¡<°Ô#5�KØ"Õ7Ø$×(Ñ(¨Ö5ñ $6ð
 ×+Ô+�Ø—6‘6Ô.Õ.°3Õ3KØ ×$Ñ$ SÖ)ñ ,÷1 ŠlrG   c                 óä  • [        5       n[        5       nU R                  R                  R                   H¶  nUR                  S:X  aK  [        UR                  5      S:X  a2  X0R                  ;   a  UR                  U5        OUR                  U5        UR                  S;  a  Mp  UR                   H6  nX@R                  ;   a  UR                  U5        M%  UR                  U5        M8     M¸     X4$ )z;
Finds nodes that consume module inputs or get_attr nodes.
Úcall_functionr   >   Úget_attrrp  )
r·   r´   rÅ   r¤   rÆ   ru   r—  r¸   rw   r¿   )r@   Ústarter_cpu_nodesÚstarter_acc_nodesrr   rÁ   s        rD   Ústarter_nodesÚ_SplitterBase.starter_nodesÄ  sÁ   € ô &)£UÐÜ%(£UÐØ—K‘K×%Ñ%×+Ô+ˆDà�w‰w˜/Ó)¬c°$×2FÑ2FÓ.GÈ1Ó.LØŸ>™>Ó)Ø%×)Ñ)¨$Õ/à%×)Ñ)¨$Ô/à�w‰wÐ9Ó9ÙàŸ
œ
�ØŸ>™>Ó)Ø%×)Ñ)¨$Ö/à%×)Ñ)¨$Ö/ó	 #ñ ,ð" !Ð3Ð3rG   c                 óÀ  ^ ^	• T R                  5       u  p[        5       m	[        U 4S jU 5       5      (       + n/ n/ nU(       d  U(       Gaa  U(       a  UOUn[        U U	4S jU 5       S 5      nUc5  U(       d  [	        S5      eUR                  [        X4S95        U(       + n/ nMi  UR                  U5        T	R                  U5        UR                  U5        UT R                  ;   aS  UT R                  ;   a"  UR                  T R                  U   T	-
  5        O!UR                  T R                  U   T	-
  5        UR                   HM  nUR                  [        ;  a  M  UT R                  ;   a  UR                  U5        M<  UR                  U5        MO     U(       a  GMW  U(       a  GMa  U(       a  UR                  [        X4S95        U(       d  [	        S5      eU$ )Nc              3   óZ   >#   • U  H   n[        TR                  U   5      S :H  v •  M"     g7f)r   N)ru   r  )rG  rÖ   r@   s     €rD   rI  Ú9_SplitterBase.put_nodes_into_subgraphs.<locals>.<genexpr>ä  s%   øé € Ð$WÒEVÀ¤S¨¯©°1©Ó%6¸!Ö%;ÒEVùs   ƒ(+c              3   óR   >#   • U  H  nTR                   U   T::  d  M  Uv •  M     g 7frY   )r  )rG  rÖ   r@   r£  s     €€rD   rI  r°  î  s"   øé € ÐKšM�q¨T¯Y©Y°q©\¸]Ñ-J—‘šMùs   ƒ'ž	'zSubgraph can't be empty)rÜ   r¤   zCouldn't create subgraphs)r¬  r·   ÚanyrM  r   rt   r   rÀ   rw   r  r¸   r  r¿   rÆ   r   )
r@   Úcurrent_cpu_nodesÚcurrent_acc_nodesÚacc_subgraphÚcurrent_subgraph_nodesr  Úcurrent_nodesrr   rÁ   r£  s
   `        @rD   rw  Ú&_SplitterBase.put_nodes_into_subgraphsÝ  s¬  ù€ à/3×/AÑ/AÓ/CÑ,ÐÜ!$£ˆô "%Ô$WÑEVÓ$WÓ!WÔWˆà+-Ðð %'ˆ	Þ×#4æ1=Ñ-ÐCTˆMÜÝK™MÓKØóˆDð ‰|Þ-Ü4Ð5NÓOÐOà× Ñ Ü LÑOôð $0Ô/�Ø)+Ð&Ùà× Ñ  Ô&Ø×Ñ˜dÔ#Ø"×)Ñ)¨$Ô/ð �t—|‘|Ó#Ø˜4Ÿ>™>Ó)Ø%×,Ñ,¨T¯\©\¸$Ñ-?À-Ñ-OÕPà%×,Ñ,¨T¯\©\¸$Ñ-?À-Ñ-OÔPð Ÿ
œ
�Ø—7‘7Ô"3Ó3Ùð ˜4Ÿ>™>Ó)Ø%×)Ñ)¨$Ö/à%×)Ñ)¨$Ö/ñ #÷A  Ñ×#4Ñ#4öV "Ø×ÑÜ ÑKôö Ü,Ð-HÓIÐIàÐrG   r  c                 óv  • / nU GH/  nUR                   (       aÃ  [        UR                  5      U R                  R                  :¼  a  UR                  U5        MU  [        S[        UR                  5       SU R                  R                   35        U(       a*  US   R                  R                  UR                  5        M¾  SUl         UR                  U5        MØ  U(       a?  US   R                   (       d+  US   R                  R                  UR                  5        GM  UR                  U5        GM2     U$ )zl
This pass finds ACC submodules with less than specified size and merges
them with adjacent CPU submodules.
zBEliminating acc subgraph because it's smaller than the threshold: z < r/   F)rÜ   ru   r¤   rÿ   r=   rt   ro   Úextend)r@   r  r’  rƒ  s       rD   rx  Ú(_SplitterBase.remove_small_acc_subgraphs  sä   € ð
 "$ˆÜ!ˆHØ��Ü�x—~‘~Ó&¨$¯-©-×*KÑ*KÓKØ—M‘M (Ö+äØ\Ü˜xŸ~™~Ó.Ð/¨s°4·=±=×3TÑ3TÐ2UðWôö Ø˜r™
×(Ñ(×/Ñ/°·±Ö?à*/˜œØŸ™ hÖ/æ &¨¡*×"3×"3Ø˜2‘J×$Ñ$×+Ñ+¨H¯N©N×;à—M‘M (×+ñ% "ð& ˆrG   c                 ó”  • / U l         U H»  nUR                  (       a  S[        U R                   5       3O"U R                   [        U R                   5       3nU R                   R	                  U5        UR
                   HA  n[        US5      (       a  [        SU S35      eX4l        X0R                  UR                  '   MC     M½     g )Nr`  ry  zNode z was already tagged)r  rÜ   ru   r   rt   r¤   Úhasattrr   ry  r
  r_   )r@   r  rƒ  ry  rr   s        rD   ry  Ú_SplitterBase.tag:  s­   € ØˆŒ	Û!ˆHð —?—?ð œs 4§9¡9›~Ð.Ñ/à×3Ñ3Ð4´S¸¿¹³^Ð4DÐEð ð
 �I‰I×Ñ˜SÔ!Ø Ÿœ�Ü˜4 ×'Ñ'Ü4°u¸T¸FÐBUÐ5VÓWÐWà”Ø69×(Ñ(¨¯©Ó3ó 'ò "rG   ra  c                 óÞ   • [        U R                  U R                  U R                  S9nU(       a<  U R                  R                  R
                   H  n[        US5      (       d  M  U?M     U$ )N)r  ry  )r   r´   r  r  rÅ   r¤   r½  ry  )r@   ra  rà   rr   s       rD   r¯   Ú_SplitterBase.splitJ  sZ   € Ü$Ø�K‰K˜Ÿ™°×1CÑ1Cñ
ˆö ØŸ™×)Ñ)×/Ô/�Ü˜4 ×'Ó'Øšñ 0ð ÐrG   c                 óx  • U R                  5       nU R                  R                  (       a  [        U5        U R                  U5      n[	        U Vs/ s H  o"R
                  (       d  M  UPM     sn5      n[	        U5      U-
  n[        SU SU S35        U R                  U5        U R                  5       $ s  snf )NzGot r_  z non-acc subgraphs)	rw  rÿ   r   rx  ru   rÜ   ro   ry  r¯   )r@   r  ÚsÚacc_subgraphs_countÚnon_acc_subgraphs_counts        rD   r×   Ú_SplitterBase.__call__T  s£   € Ø×1Ñ1Ó3ˆ	Ø�=‰=×:×:Ü-¨iÔ8Ø×3Ñ3°IÓ>ˆ	Ü!©iÓ"Dªi¨¿8½8§1©iÑ"DÓEÐÜ"% i£.Ð3FÑ"FÐÜØÐ&Ð'Ð':Ð;RÐ:SÐSeÐfô	
ð 	�‰�ÔØ�z‰z‹|Ðùò #Es   ÁB7Á)B7c                 óP  • U " 5       n/ nUR                  5        H  u  p4UR                  U5        M     U R                  R                  S:”  a.  [	        U5      U R                  R                  :”  a  [        S5      e[        XR                  U5      n[        XU R                  5      $ )Nr   ziCannot fulfill max_acc_splits limit. This may cause split fragmentation and result in performance issues.)
Únamed_childrenrt   rÿ   r0   ru   Ú
ValueErrorr    rþ   r   r   )r@   rà   Úsubmodule_namesr_   Ú_modrá   s         rD   Úgenerate_split_resultsÚ$_SplitterBase.generate_split_resultsa  s›   € Ù“vˆØˆØ&×5Ñ5Ö7‰JˆDØ×"Ñ" 4Ö(ñ 8ð �M‰M×(Ñ(¨1Ó,Ü�OÓ$ t§}¡}×'CÑ'CÓCäð0óð ô :Ø×+Ñ+¨_ó
Ðô ˜<¸4×;VÑ;VÓWÐWrG   )r
  r  r¸   r  r  r´   r   rµ   rþ   rÿ   r  )Ú_run_on_cpu_FN©FrY   )2rH   rI   rJ   rK   rc   r|  rd   re   rÙ   r   r   r   r-   rg   rO   r
   r   rE   rÒ   r  rf   r   r  r	  r   ÚnnÚModuler  r  r   r:  r\  r�  r;   r“  rš  rž  r¥  rP  r¬  Úlistr   rw  rx  ry  r¯   r×   r   rË  rP   rQ   rG   rD   rý   rý   ã  sF  † ñ(ðV €Gð '5Ø"Ø6:ñ."à—‘×$Ñ$ð."ð ˜s‘mð."ð .ð	."ð
 'ð."ð !$ð."ð ð."ð Ð2Ñ3õ."ðh(¨¨S°#¨X©ô (ð˜4 §¡§¡¨wÐ 6Ñ7ô ò$2ð&Ø—8‘8×'Ñ'ðØ18ðà	�‰�‰ôð_ §¡×!5Ñ!5ð _¸wð _È3ô _ð0Ø—8‘8×'Ñ'ð0Ø:Bô0ñ4A¨tõ AñFm¨õ mðh '+ñØ˜s‘mðà	ˆe�h‰h�m‰m˜WÐ$Ñ	%õð,&°D¸¿¹¿¹ÈÐ9OÑ4Pô &ð6°ð ¸ô ð"(* sô (*ð\4˜u W¨gÐ%5Ñ6ô 4ð2@¨$¨x©.ô @ðD°D¸±Nð ÀtÈHÁ~ô ð6:˜T (™^ô :ñ  ð °·±×1EÑ1Eõ ð˜%Ÿ(™(×.Ñ.ô ðX¨÷ XrG   rý   rÎ  )Hr8   ré   rŽ   r¬   Úcollectionsr   Úcollections.abcr   r   Údataclassesr   Útypingr   r   r	   r
   rd   Útorch._loggingr   Útorch.fx._compatibilityr   Útorch.fx.noder   Ú"torch.fx.passes.graph_manipulationr   Úgraph_drawerr   rµ   r   r   Ú
shape_propr   Úsplit_utilsr   r   Útools_commonr   r   r   r   r   r   Ú__all__rL   rM   rN   ÚTRACKER_DUMP_PATHr£   r«   Ú$ENV_FX_NET_ACC_SPLITTER_TRACKER_MODEÚ)ENV_FX_NET_ACC_SPLITTER_TRACKER_DUMP_PATHr®   r­   r}   r¹   r+   rÞ   r-   r!   r"   r   rù   r   r   r   rÏ  rÐ  rg   rO   rÒ   r    rý   rQ   rG   rD   Ú<module>râ     s  ðä Û Û Û 	Ý #ß .Ý !ß 5Ó 5ã Ý +Ý 1Ý !Ý ?å 'ß BÝ !ß I÷÷ ò€ð  Ð ØÐ Ø Ð ð &Ð Ø€Ø€
à'IÐ $Ø,SÐ )à/ð .ð �j‰j�n‰nØ-Ð/@ó€ðð -/¯J©J¯N©NØ(¨#ó-€ˆgÐ(Ñ)ó ÷
H
ñ H
ñV  eÑ,÷Wð Wó -ðWñ.  eÑ,÷o'ð o'ó -ðo'ñd  eÑ,÷nð nó -ðnñb  eÑ,ô	 ó 	ó -ð	ñ  eÑ,Ø
÷)ð )ó ó -ð)ñ  eÑ,ô"�*ó "ó -ð"ñ   eÑ,ð
 ñ	-Ø�8‰8�?‰?ð-à�S‰Mð-ð   ‘}ð-ð ð	-ð
 
ˆ#ˆsˆ(�^ô-ó -ð-÷`P
Xò P
XrG   