§
    ŠŠtj<0  ã                   ó¤  — d dl mZ d dlmZ d dlmZmZ  G d„ d¦  «        Z G d„ de¦  «        Z G d„ d	e¦  «        Z	 G d
„ de¦  «        Z
 G d„ de¦  «        Z G d„ de¦  «        Zdedee         defd„Zdedeee	f         de
fd„Zdee         deee	f         deee
f         fd„Zdedededefd„Zdee         deee
f         dedefd„ZdS ) é    )ÚEnum)Ú
NamedTuple)Úmap_argÚNodec                   óV   — e Zd ZdZdeddfd„Zdefd„Zdd„Zde	ddfd	„Z
de	ddfd
„ZdS )Ú	Partitionz—Partition class contains all the information about an individual partition.
    It also provides necessary methods for manipulating the partition.
    Úpartition_idÚreturnNc                 ó°   — t          ¦   «         | _        || _        t          ¦   «         | _        t          ¦   «         | _        d| _        d| _        g | _        d S )Néÿÿÿÿr   )ÚsetÚnodesr	   ÚparentsÚchildrenÚ	bfs_levelÚused_mem_bytesÚlogical_device_ids)Úselfr	   s     úe/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/fx/experimental/partitioner_utils.pyÚ__init__zPartition.__init__   sH   € Ý #¡¤ˆŒ
Ø(ˆÔÝ'*¡u¤uˆŒÝ(+©¬ˆŒØ ˆŒØ#$ˆÔØ-/ˆÔÐÐó    c                 ó*   — t          | j        ¦  «        S ©N)Ústrr	   )r   s    r   Ú__str__zPartition.__str__   s   € Ý�4Ô$Ñ%Ô%Ð%r   c                 ón   — d| _         | j        D ]%}| xj         t          || j        ¦  «        z  c_         Œ&d S )Nr   )r   r   Úget_extra_size_of)r   Únodes     r   Úrecalculate_mem_sizezPartition.recalculate_mem_size   sM   € ØˆÔØ”Jð 	Gð 	GˆDØÐÔÕ#4°T¸4¼:Ñ#FÔ#FÑFÐÔÐð	Gð 	Gr   r   c                 ó  — i }t          |j        |j        ¦  «         t          |j        |j        ¦  «         |D ]%}|j        dv r| j                             |¦  «         Œ&| j                             |¦  «         |                      ¦   «          d S )N>   Úget_attrÚplaceholder)r   ÚargsÚ
setdefaultÚkwargsÚopr   Úaddr   )r   r   Úinput_nodesÚns       r   Úadd_nodezPartition.add_node   s�   € Ø(*ˆÝ�”	˜;Ô1Ñ2Ô2Ð2Ý�”˜[Ô3Ñ4Ô4Ð4àð 	"ð 	"ˆAØŒtÐ2Ð2Ð2Ø”
—’˜qÑ!Ô!Ð!øØŒ
�Š�tÑÔÐØ×!Ò!Ñ#Ô#Ð#Ð#Ð#r   c                 óv  ‡ — |‰ j         v r®‰ j                              |¦  «         i }t          |j        |j        ¦  «         t          |j        |j        ¦  «         |D ]E}t          ˆ fd„|j        D ¦   «         ¦  «        r#|j        dv r‰ j                              |¦  «         ŒF‰  	                    ¦   «          d S d S )Nc              3   ó*   •K  — | ]}|‰j         vV — Œd S r   )r   )Ú.0r)   r   s     €r   ú	<genexpr>z(Partition.remove_node.<locals>.<genexpr>4   s;   øè è € ð ð Ø,-�A˜TœZÐ'ðð ð ð ð ð r   >   r!   r"   )
r   Úremover   r#   r$   r%   ÚallÚusersr&   r   )r   r   r(   Ú
input_nodes   `   r   Úremove_nodezPartition.remove_node(   sæ   ø€ à�4”:ÐÐØŒJ×Ò˜dÑ#Ô#Ð#à,.ˆKÝ�D”I˜{Ô5Ñ6Ô6Ð6Ý�D”K Ô!7Ñ8Ô8Ð8ð *ð 2ð 2�
Ýð ð ð ð Ø1;Ô1Aðñ ô ñ ô ð 2à ”mÐ'BÐBÐBØ”J×%Ò% jÑ1Ô1Ð1øØ×%Ò%Ñ'Ô'Ð'Ð'Ð'ð Ðr   )r
   N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__Úintr   r   r   r   r   r*   r3   © r   r   r   r      s·   € € € € € ðð ð0 Sð 0¨Tð 0ð 0ð 0ð 0ð&˜ð &ð &ð &ð &ðGð Gð Gð Gð
	$˜Tð 	$ dð 	$ð 	$ð 	$ð 	$ð( ð (¨ð (ð (ð (ð (ð (ð (r   r   c                   ó.   — e Zd ZU eed<   eed<   eed<   dS )ÚDeviceÚnameÚavailable_mem_bytesÚ
logical_idN)r4   r5   r6   r   Ú__annotations__r8   r9   r   r   r;   r;   ;   s.   € € € € € € Ø
€I€I�IØÐÐÑØ€O€O�O€O€Or   r;   c                   ó$   — e Zd ZU eed<   eed<   dS )ÚNodeLatencyÚmem_latency_secÚcomputer_latency_secN©r4   r5   r6   Úfloatr?   r9   r   r   rA   rA   A   s*   € € € € € € àÐÐÑàÐÐÑÐÐr   rA   c                   ó.   — e Zd ZU eed<   eed<   eed<   dS )ÚPartitionLatencyrB   rC   Úoverall_latency_secNrD   r9   r   r   rG   rG   H   s6   € € € € € € àÐÐÑàÐÐÑàÐÐÑÐÐr   rG   c                   ó"   — e Zd ZdZdZdZdZdZdS )ÚPartitionModer   é   é   é   é   N)r4   r5   r6   Ú
size_basedÚ	sparse_nnÚ
cost_awareÚkl_basedÚ	aot_basedr9   r   r   rJ   rJ   Q   s'   € € € € € Ø€JØ€IØ€JØ€HØ€I€I€Ir   rJ   c                   óÀ   — e Zd ZU ee         ed<   ej        Zeed<   dZ	e
ed<   i Zeeef         ed<   i Zeeef         ed<   i Zeeee         f         ed<   dZeed	<   d
S )ÚPartitionerConfigÚdevicesÚmodeç        Útransfer_rate_bytes_per_secÚnode_to_latency_mappingÚnode_to_partition_mappingÚ#partition_to_logical_device_mappingFÚsaturate_hostN)r4   r5   r6   Úlistr;   r?   rJ   rO   rW   rY   rE   rZ   Údictr   rA   r[   r8   r\   r]   Úboolr9   r   r   rU   rU   Y   s¨   € € € € € € Ø�&Œ\ÐÐÑØ'Ô2€Dˆ-Ð2Ð2Ñ2Ø),Ð Ð,Ð,Ñ,Ø79Ð˜T $¨Ð"3Ô4Ð9Ð9Ñ9Ø13Ð˜t D¨# IœÐ3Ð3Ñ3Ø@BÐ'¨¨c°4¸´9¨nÔ)=ÐBÐBÑBà€M�4ÐÐÑÐÐr   rU   r   r   r
   c                 ó<  — i }t          | j        |j        ¦  «         t          | j        |j        ¦  «         d}|D ]3}||vr-t	          |dd¦  «        }|r||j        z  }Œ$t          d¦  «        ‚Œ4t	          | dd¦  «        }|r||j        z  }nt          d¦  «        ‚|S )zƒGiven a node and a set of nodes,
    this function return the extra size that needed
    if this node is included in this set.
    r   Ú
size_bytesNznode has no size_bytes attr)r   r#   r$   r%   ÚgetattrÚoutput_sizeÚRuntimeErrorÚ
total_size)r   r   r(   Útotal_size_of_input_nodesr)   rb   s         r   r   r   d   sÓ   € ð %'€KÝˆDŒI�{Ô-Ñ.Ô.Ð.ÝˆDŒK˜Ô/Ñ0Ô0Ð0à !ÐØð Bð Bˆà�Eˆ>ˆ>Ý   L°$Ñ7Ô7ˆJØð BØ)¨ZÔ-CÑCÐ)Ð)å"Ð#@ÑAÔAÐAð õ ˜˜|¨TÑ2Ô2€JØð :Ø! ZÔ%:Ñ:Ð!Ð!åÐ8Ñ9Ô9Ð9Ø$Ð$r   Ú	partitionrZ   c           	      ó   ‡ ‡‡— dt           dt          t                   fd„}dt          dt          dt          fˆˆˆ fd„Š |‰ ¦  «        }t          ddd¬¦  «        }|D ]0} ‰|t          ddd¬¦  «        ¦  «        }|j        |j        k    r|}Œ1|S )	zVGiven a partition and its nodes' latency, return a PartitionLatency for this partitionrh   r
   c                 ó   ‡ — g }‰ j         D ]r}|j        dv rŒi }t          |j        |j        ¦  «         t          |j        |j        ¦  «         t          ˆ fd„|D ¦   «         ¦  «        s|                     |¦  «         Œs|S )z>Given a partition, return a list of nodes on the top bfs level>   r!   r"   c              3   ó<   •K  — | ]}|‰j         v o|j        d vV — ŒdS )>   r!   r"   N)r   r&   )r-   r)   rh   s     €r   r.   zFget_latency_of_one_partition.<locals>.get_top_nodes.<locals>.<genexpr>’   sK   øè è € ð ð àð �Y”_Ð$ÐP¨¬Ð5PÐ)Pðð ð ð ð ð r   )r   r&   r   r#   r$   r%   ÚanyÚappend)rh   Ú	top_nodesr   r(   s   `   r   Úget_top_nodesz3get_latency_of_one_partition.<locals>.get_top_nodes…   s°   ø€ à "ˆ	Ø”Oð 	'ð 	'ˆDàŒwÐ5Ð5Ð5ØØ,.ˆKÝ�D”I˜{Ô5Ñ6Ô6Ð6Ý�D”K Ô!7Ñ8Ô8Ð8õ ð ð ð ð à$ðñ ô ñ ô ð 'ð × Ò  Ñ&Ô&Ð&øØÐr   r   Úpartition_latencyc           	      óž  •— ‰|          }|j         t          |j        |j        ¦  «        z   }|j        |j        z   }|j        |j        z   }t	          | j        ¦  «                             ‰j        ¦  «        }|rFt          ddd¬¦  «        }|D ]/} ‰
|t          |||¦  «        ¦  «        }	|	j         |j         k    r|	}Œ0|S t          |||¦  «        S )zyGiven a top node of a partition, this function returns
        the latency of the critical path in the partition
        rX   ©rB   rC   rH   )	rH   ÚmaxrC   rB   r   r1   Úintersectionr   rG   )r   rp   Únode_latencyrH   rB   rC   r1   Úmax_latencyr)   Únew_partition_latencyÚ
dfs_helperrZ   rh   s             €€€r   rx   z0get_latency_of_one_partition.<locals>.dfs_helper™   s   ø€ ð /¨tÔ4ˆà/ÔCÅcØÔ-¨|Ô/KñG
ô G
ñ 
Ðð
 Ô-°Ô0LÑLð 	ð
 Ô2°\Ô5VÑVð 	õ �D”J‘”×,Ò,¨Y¬_Ñ=Ô=ˆØð 	Ý*Ø #¸#ÐSVðñ ô ˆKð ð 8ð 8�à(2¨
ØÝ$Ø'Ð)=Ð?Rñô ñ)ô )Ð%ð *Ô=Ø!Ô5ò6ð 6ð #8�KøØÐåØÐ1Ð3Fñ
ô 
ð 	
r   rX   rr   )r   r^   r   rG   rH   )rh   rZ   ro   rn   Úcritical_path_latencyr   rp   rx   s   ``     @r   Úget_latency_of_one_partitionrz   €   sÿ   øøø€ ð
¥ð ­tµD¬zð ð ð ð ð((
�ð (
Õ2Bð (
ÕGWð (
ð (
ð (
ð (
ð (
ð (
ð (
ð (
ðX �˜iÑ(Ô(€IÝ,Ø°#È3ðñ ô Ðð ð 6ð 6ˆØ&˜JØÝØ #¸#ÐSVðñ ô ñ
ô 
Ðð Ô1Ø#Ô7ò8ð 8ð %6Ð!øØ Ð r   Ú
partitionsc                 ó>   — i }| D ]}t          ||¦  «        }|||<   Œ|S )zŽGiven all the partitions and node_to_latency_mapping dictionary,
    return a mapping dictionary of each partition to its overall latency
    )rz   )r{   rZ   Úpartition_to_latency_mappingrh   rp   s        r   Ú get_partition_to_latency_mappingr~   Ù   sJ   € ð GIÐ àð Dð Dˆ	Ý8ØÐ.ñ
ô 
Ðð 3DÐ$ YÑ/Ð/Ø'Ð'r   Úparent_partitionÚchild_partitionrY   c                 ó„  — | j         g k    r|j         g k    r| j         |j         k    rdS d}t          ¦   «         }|j        D ]|}i }t          |j        |j        ¦  «         t          |j        |j        ¦  «         |D ]A}|| j        v r6||vr2t          |dd¦  «        }|�
||j        z  }| 	                    |¦  «         ŒBŒ}||z  S )zfGiven two partitions (parent and child),
    calculate the communication latency between the two.
    rX   r   rb   N)
r   r   r   r   r#   r$   r%   rc   rd   r'   )	r   r€   rY   Ú	comm_sizeÚvisited_nodesr   r(   r)   rb   s	            r   Úget_comm_latency_betweenr„   é   sû   € ð 	Ô+¨rÒ1Ð1ØÔ.°"Ò4Ð4ØÔ/°?Ô3UÒUÐUàˆsà€Iå‘E”E€Mð
  Ô%ð 	%ð 	%ˆØ(*ˆÝ�”	˜;Ô1Ñ2Ô2Ð2Ý�”˜[Ô3Ñ4Ô4Ð4Øð 	%ð 	%ˆAØÐ$Ô*Ð*Ð*¨q¸Ð/EÐ/EÝ$ Q¨°dÑ;Ô;�
ØÐ)Ø Ô!7Ñ7�IØ×!Ò! !Ñ$Ô$Ð$øð	%ð Ð2Ñ2Ð2r   r}   c                 óâ   ‡‡‡— dt           dt          dt          fˆˆˆfd„Šdt          t                    dt          t                    fd„} || ¦  «        }d}|D ]} ‰|d¦  «        }||k    r|}Œ|S )zŽGiven all partitions in a graph, find the critical path among all partitions
    and return its latency as the latency of the whole graph
    rh   Úlatency_so_far_secr
   c                 ó¢   •— |‰|          j         z  }| j        r6d}| j        D ]*}t          | |‰¦  «        } ‰|||z   ¦  «        }||k    r|}Œ+|S |S )zJThis function helps to recursively get the latency of a path of partitionsrX   )rH   r   r„   )	rh   r†   Úmax_latency_secÚchildÚcomm_latency_secÚnew_latency_secrx   r}   rY   s	         €€€r   rx   z4get_latency_of_partitioned_graph.<locals>.dfs_helper  s˜   ø€ ð 	Ð:Øô
ä
ñ	Ðð Ôð 	#Ø!ˆOØ"Ô+ð 	6ð 	6�å#;Ø˜uÐ&Añ$ô $Ð ð #- *ØÐ-Ð0@Ñ@ñ#ô #�ð # _Ò4Ð4Ø&5�OøØ"Ð"Ø!Ð!r   r{   c                 ó   — d„ | D ¦   «         }|S )zvThis function is to return all the partitions without parents
        as the starting points of all the paths
        c                 óB   — g | ]}t          |j        ¦  «        d k    ¯|‘ŒS )r   )Úlenr   )r-   rh   s     r   ú
<listcomp>zPget_latency_of_partitioned_graph.<locals>.get_top_partitions.<locals>.<listcomp>1  s4   € ð 
ð 
ð 
Ø#µS¸Ô9JÑ5KÔ5KÈqÒ5PÐ5PˆIÐ5PÐ5PÐ5Pr   r9   )r{   Útop_partitionss     r   Úget_top_partitionsz<get_latency_of_partitioned_graph.<locals>.get_top_partitions,  s(   € ð

ð 
Ø'1ð
ñ 
ô 
ˆð Ðr   rX   )r   rE   r^   )	r{   r}   rY   r‘   r�   Úcritical_path_latency_secrh   Úlatency_secrx   s	    ``     @r   Ú get_latency_of_partitioned_graphr”     s½   øøø€ ð"�ið "½Uð "Åuð "ð "ð "ð "ð "ð "ð "ð "ð,¥t­I¤ð ½4Å	¼?ð ð ð ð ð (Ð'¨
Ñ3Ô3€NØ #ÐØ#ð 4ð 4ˆ	Ø �j ¨CÑ0Ô0ˆØÐ2Ò2Ð2Ø(3Ð%øØ$Ð$r   N)Úenumr   Útypingr   Útorch.fx.noder   r   r   r;   rA   rG   rJ   rU   r   r8   r   r_   rz   r^   r~   rE   r„   r”   r9   r   r   ú<module>r˜      sv  ðØ Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð à 'Ð 'Ð 'Ð 'Ð 'Ð 'Ð 'Ð 'ð1(ð 1(ð 1(ð 1(ð 1(ñ 1(ô 1(ð 1(ðhð ð ð ð ˆZñ ô ð ð ð  ð  ð  ð  �*ñ  ô  ð  ðð ð ð ð �zñ ô ð ðð ð ð ð �Dñ ô ð ð ð  ð  ð  ð  ˜
ñ  ô  ð  ð%˜Dð %¨¨T¬ð %°sð %ð %ð %ð %ð8V!ØðV!Ø37¸¸kÐ8IÔ3JðV!àðV!ð V!ð V!ð V!ðr(Ø�Y”ð(Ø:>¸tÀ[Ð?PÔ:Qð(à	ˆ)Ð%Ð
%Ô&ð(ð (ð (ð (ð !3Øð!3àð!3ð "'ð!3ð ð	!3ð !3ð !3ð !3ðH/%Ø�Y”ð/%à"& yÐ2BÐ'BÔ"Cð/%ð "'ð/%ð ð	/%ð /%ð /%ð /%ð /%ð /%r   