Ë
    g^(ht'  ã                   ó²  — d dl Z d dlmZ d dlmZmZ d dlZd dlZd dl	m
Z
mZmZmZ d dlmZ d dlmZmZ  G d„ de«      Z G d	„ d
e«      Z G d„ de«      Z G d„ de«      Zdej.                  j0                  dedededej2                  defd„Z G d„ de«      Z G d„ d«      Zdedefd„Zdededefd„Z dededefd„Z!d%d e"d!edefd"„Z#d#ede$e"e%f   fd$„Z&y)&é    N)ÚOrderedDict)ÚcastÚ	TypedDict)Ú_MemRefTypeÚ_ModMemStatsÚ	_ModStateÚ
MemTracker)ÚRuntimeEstimator)ÚSACEstimatorÚSACTradeOffStatsc                   óN   — e Zd ZU ee   ed<   ee   ed<   ee   ed<   ee   ed<   y)ÚModOrderÚfw_pre_orderÚbw_pre_orderÚfw_post_orderÚbw_post_orderN)Ú__name__Ú
__module__Ú__qualname__ÚlistÚstrÚ__annotations__© ó    ú`/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/distributed/_tools/ilp_utils.pyr   r      s*   … Ø�s‘)ÓØ�s‘)ÓØ˜‘9ÓØ˜‘9Ôr   r   c                   ó"   — e Zd ZU eed<   eed<   y)Ú
ModRuntimeÚfwÚbwN)r   r   r   Úfloatr   r   r   r   r   r      s   … ØƒIØ„Ir   r   c                   óò   — e Zd ZU eed<   eed<   eed<   eed<   eed<   eed<   eed<   eed<   eed	<   eed
<   eed<   eed<   eed<   eed<   eed<   eed<   ee   ed<   ee   ed<   ee   ed<   e	eef   ed<   y)ÚModStatsÚfqnÚparam_per_moduleÚgrad_per_moduleÚ
grad_totalÚact_fw_per_moduleÚact_bw_per_moduleÚact_grad_per_moduleÚ	act_totalÚinput_per_moduleÚoutput_per_moduleÚfw_runtime_per_moduleÚbw_runtime_per_moduleÚis_leafÚsac_runtimeÚ
sac_memoryÚ
n_segmentsÚslopesÚ
interceptsÚbreakpointsÚtradeoff_curveN)
r   r   r   r   r   Úintr    Úboolr   r   r   r   r   r"   r"      s‹   … Ø	ƒHàÓàÓàƒOàÓàÓàÓð ƒNàÓàÓà Ó à Ó àƒMàÓàƒOàƒOà�‰KÓà�U‘Óà�e‘Óà  u Ñ-Ô-r   r"   c                   ó(   — e Zd ZU eed<   ee   ed<   y)Ú
ModuleInfoÚ	mod_orderÚ	mod_statsN)r   r   r   r   r   r   r"   r   r   r   r:   r:   I   s   … ØÓØ�H‰~Ôr   r:   ÚmodelÚmem_trackerÚruntime_estimatorÚsac_estimatorÚdevÚreturnc           	      óš  — t        t        j                  |j                  «      «      }|j                  j                  «       D ��ci c]  \  }}||d   |d   dœ“Œ }}}t        |j                  «      t        |j                  «      t        |j                  «      t        |j                  «      dœ}	|j                  «        t        j                  |j                  «      }
|	g dœ}| j                  «       D �]Ë  }|j                  |d«      x}sŒ|
j                  |j                  d«      x}rW|j                   }|j"                  }|j$                  }|j&                  }|j(                  }|j*                  }|j,                  }d}ndx}x}}g x}x}}t/        «       }d	}i d
|j                  “d|j0                  “d|j0                  “d|j2                  t4        j6                     d   |   t8        j:                     “dt=        d|j2                  t4        j>                     d   |   t8        j@                     |j2                  t4        jB                     d   |   t8        j@                     z
  |jD                  z
  «      “dt=        d|j2                  t4        jF                     d   |   t8        j@                     «      “d|j2                  t4        jF                     d   |   t8        jH                     |j2                  t4        j6                     d   |   t8        jH                     z
  “d|j2                  t4        j>                     d   |   t8        j@                     “d|jJ                  “d|jD                  “d||j                     d   “d||j                     d   “d|“d|“d|“d|“d|“|||dœ¥}|d   jM                  |«       �ŒÎ |S c c}}w )aþ  
    Collect modulewise stats for a given model, including memory, runtime, and AC tradeoff stats.

    Args:
        model: nn.Module object
        runtime_estimator: RuntimeEstimator object with runtime stats
        mem_tracker: MemTracker object with memory stats
        sac_estimator: SACEstimator object with AC tradeoff stats
        dev: device the model was run on (used to extract memory stats from MemTracker)

    Returns:
        ModuleInfo: A dictionary with module order and module stats.
    r   r   )r   r   )r   r   r   r   )r;   r<   NFr   Tr#   r$   r%   r&   éÿÿÿÿr'   r(   r)   r*   r+   r,   r-   r.   r/   r0   r1   r2   r3   )r4   r5   r6   r<   )'ÚdictÚcopyÚdeepcopyÚmemory_trackingÚmod_runtimesÚitemsr   Úmod_fw_pre_orderÚmod_bw_pre_orderÚmod_fw_post_orderÚmod_bw_post_orderÚpwlf_sac_tradeoff_curveÚsac_mod_tradeoff_statsÚmodulesÚgetÚmod_fqnr0   r1   r2   r3   r4   Ú
fit_breaksr6   r   Úparameter_memÚ	snapshotsr   ÚPRE_BWr   ÚGRADÚmaxÚPOST_FWÚACTÚPRE_FWÚ
output_memÚPEAK_BWÚTEMPÚ	input_memÚappend)r=   r>   r?   r@   rA   Úmod_mem_statsr#   ÚvÚmod_runtime_statsr;   Úmod_sac_tradeoff_statsÚmodule_infoÚmodÚmod_mem_statÚtradeoff_statsr0   r1   r2   r3   r4   r5   r6   r/   Úmod_stats                           r   Úaggregate_statsrk   N   s9  € ô, :>Ü�‰�k×1Ñ1Ó2ó:€Mð (×4Ñ4×:Ñ:Ó<÷0áˆC�ð 	�A�d‘G 1 T¡7Ñ+Ñ+ð0Ðñ 0ô Ð.×?Ñ?Ó@ÜÐ.×?Ñ?Ó@ÜÐ/×AÑAÓBÜÐ/×AÑAÓBñ	€Ið ×)Ñ)Ô+Ü:>¿-¹-Ø×,Ñ,ó;Ðð
 Øñ€Kð
 �}‰}‹ó 76ˆØ(×,Ñ,¨S°$Ó7Ð7ˆ<Ñ7Ø!7×!;Ñ!;¸L×<PÑ<PÐRVÓ!WÐWˆ~ÐWØ,×8Ñ8�Ø+×6Ñ6�
Ø+×6Ñ6�
Ø'×.Ñ.�Ø+×6Ñ6�
Ø,×7Ñ7�Ø!/×!>Ñ!>�Ø‘à89Ð9�Ð9˜j¨:Ø46Ð6�Ð6˜ kÜ<G»M�Ø�ð&"Ø�|×+Ñ+ð&"à" L×$>Ñ$>ð&"ð " <×#=Ñ#=ð&"ð ˜l×4Ñ4´Y×5EÑ5EÑFÀrÑJÈ3ÑOÜ×$Ñ$ñð	&"ð $¤SØØ ×*Ñ*¬9×+<Ñ+<Ñ=¸bÑAÀ#ÑFÄ{ÇÁÑWØ"×,Ñ,¬Y×-=Ñ-=Ñ>¸rÑBÀ3ÑGÌÏÉÑXñYà"×-Ñ-ñ.ó&ð&"ð $¤SØØ ×*Ñ*¬9×+<Ñ+<Ñ=¸bÑAÀ#ÑFÄ{ÇÁÑWó&ð&"ð" &Ø ×*Ñ*¬9×+<Ñ+<Ñ=¸bÑAÀ#ÑFÄ{×GWÑGWÑXØ"×,Ñ,¬Y×-=Ñ-=Ñ>¸rÑBÀ3ÑGÜ#×(Ñ(ññð%&"ð. ˜\×3Ñ3´I×4EÑ4EÑFÀrÑJÈ3ÑOÜ—O‘Oñð/&"ð4 # L×$:Ñ$:ð5&"ð6 $ \×%<Ñ%<ð7&"ð8 (Ð):¸<×;OÑ;OÑ)PÐQUÑ)Vð9&"ð: (Ð):¸<×;OÑ;OÑ)PÐQUÑ)Vð;&"ð< ˜7ð=&"ð> ˜{ð?&"ð@ ˜jðA&"ðB ˜jðC&"ðD ˜&ðE&"ðF )Ø*Ø"0òK&"ˆHðN ˜Ñ$×+Ñ+¨HÖ5ðo76ðr Ðùóc0s   ÁOc                   ó"   — e Zd ZU eed<   eed<   y)ÚNodeÚindexÚpos_fw_post_orderN)r   r   r   r7   r   r   r   r   rm   rm   ½   s   … ØƒJØÔr   rm   c                   ó,   — e Zd Zdeddfd„Zdeddfd„Zy)ÚGraphÚnrB   Nc                 óf   — g | _         i | _        t        j                  ||f«      | _        g | _        y )N)ÚnodesÚ	name2nodeÚnpÚzerosÚ	ad_matrixr   )Úselfrr   s     r   Ú__init__zGraph.__init__Ã   s,   € Ø!#ˆŒ
Ø*,ˆŒÜŸ™ 1 a &Ó)ˆŒØ(*ˆÕr   Únodec                 ó^   — | j                   j                  |«       || j                  |d   <   y ©Nr#   )rt   ra   ru   )ry   r{   s     r   Úadd_nodezGraph.add_nodeÉ   s&   € Ø�
‰
×Ñ˜$ÔØ&*ˆ�‰�t˜E‘{Ò#r   )r   r   r   r7   rz   rm   r~   r   r   r   rq   rq   Â   s(   „ ð+˜#ð + $ó +ð+˜Tð + dô +r   rq   rf   c                 ó6  ‡— | d   }| d   d   Št        |«      t        ‰«      k(  sJ ‚t        |«      }t        |«      }| d   d   |_        t        |ˆfd„¬«      | d<   t	        |«      D ]L  \  }}t        t        |«      }||d<   |j                  j                  |d   «      |d	<   |j                  |«       ŒN t        |«      D ]S  }t        ||«      D ]B  }t        |j                  |   d   |j                  |   d   «      rd
|j                  |   |<   ŒB ŒS ŒU |S )z›
    Parse module info and create a graph (tree) of modules. The graph will be
    used by MILP solver to find optimal SAC and/or FSDP configurations.
    r<   r;   r   r   c                 ó,   •— ‰j                  | d   «      S r}   )rn   )Úxr   s    €r   ú<lambda>z#parse_module_info.<locals>.<lambda>ß   s   ø€  ×!3Ñ!3°A°e±HÓ!=€ r   )Úkeyrn   r#   ro   é   )Úlenrq   r   ÚsortedÚ	enumerater   rm   rn   r~   ÚrangeÚis_self_or_submodulert   rx   )	rf   r<   Ún_nodesÚgÚiÚone_mod_statsr{   Újr   s	           @r   Úparse_module_infor�   Î   s6  ø€ ð
 ˜KÑ(€IØ˜{Ñ+¨NÑ;€Läˆy‹>œS Ó.Ò.Ð.Ð.Ü�)‹n€Gô 	ˆg‹€AØ! +Ñ.¨Ñ?€A„Oô  &ØÓ=ô €K�Ñô & iÓ0ò Ñˆˆ=Üœ$ Ó.ˆØˆˆW‰Ø$%§O¡O×$9Ñ$9¸$¸u¹+Ó$FˆÐ Ñ!Ø	�
‰
�4Õð	ô �7‹^ò ˆÜ�q˜'Ó"ò 	ˆAÜ# A§G¡G¨A¡J¨uÑ$5°q·w±w¸q±zÀ%Ñ7HÔIØ$%�—‘˜A‘˜qÒ!áñ		ðð €Hr   Úname_descendantÚname_ancestorc                 ó   — | |k(  xs |dz   | v S )z[
    check if name_descendant is a submodule of name_ancestor, or if they are the same
    ú.r   ©r�   r‘   s     r   r‰   r‰   ò   s   € ð ˜mÑ+ÒU¨}¸sÑ/BÀoÐ/UÐUr   c                 ó   — |dz   | v S )zN
    if name_descendant is a submodule of name_ancestor, but not the same
    r“   r   r”   s     r   Úis_submoduler–   ù   s   € ð ˜3Ñ /Ð1Ð1r   ÚbÚunitc                 ób   — |dk(  r	| dz  d›d�S |dk(  r	| dz  d›d�S |dk(  r	| d	z  d›d
�S | d›d�S )zN
    return a string that represent the number of bytes in a desired unit
    ÚKiBi   z.2fz KiBÚMiBi   z MiBÚGiBi   @z GiBz bytesr   )r—   r˜   s     r   Údisplay_bytesr�      sa   € ð ˆu‚}Ø�e‘)˜C� Ð%Ð%Øˆu‚}Ø�e‘)˜C� Ð%Ð%Øˆu‚}Ø�e‘)˜C� Ð%Ð%Ø�ˆW�FÐÐr   Úgraphc                 ó\  — | j                   d   d   }t        | j                   «      }d}t        |«      D ]M  }| j                   |   d   }| j                   |   d   }| j                   |   d   }t        |||z   |z   |z   «      }ŒO | j                   d   d   | j                   d   d   z   }||fS )aV  
    Get the baseline peak memory and runtime.
    Baseline here means there is no FSDP or AC.
    Memory includes the parameters, gradients, activations, and activation gradients.
    Memory does not include e.g., optimizer states, embedding tables, etc.

    Returns:
        int: peak memory in bytes
        float: compute time in ms
    r   r$   r&   r)   r*   r-   r.   )rt   r…   rˆ   rY   )	rž   ÚP_1Ú	num_nodesÚpeak_memrŒ   ÚTG_iÚAG_iÚTA_iÚcompute_times	            r   Ú get_peak_memory_runtime_baseliner§     sÊ   € ð �+‰+�a‰.Ð+Ñ
,€CÜ�E—K‘KÓ €IØ€HÜ�9Óò ;ˆØ�{‰{˜1‰~˜lÑ+ˆØ�{‰{˜1‰~Ð3Ñ4ˆØ�{‰{˜1‰~˜kÑ*ˆÜ�x  t¡¨dÑ!2°TÑ!9Ó:‰ð	;ð 	�‰�A‰Ð.Ñ/Ø
�+‰+�a‰.Ð0Ñ
1ñ	2ð ð �lÐ#Ð#r   )r›   )'rF   Úcollectionsr   Útypingr   r   Únumpyrv   ÚtorchÚ$torch.distributed._tools.mem_trackerr   r   r   r	   Ú*torch.distributed._tools.runtime_estimatorr
   Ú&torch.distributed._tools.sac_estimatorr   r   r   r   r"   r:   ÚnnÚModuleÚdevicerk   rm   rq   r�   r   r8   r‰   r–   r7   r�   Útupler    r§   r   r   r   ú<module>r³      sA  ðÛ Ý #ß "ã ã ÷ó õ Hß Qôˆyô ô�ô ô
(.ˆyô (.ôV�ô ð
lØ�8‰8�?‰?ðlàðlð (ðlð  ð	lð
 
�‰ðlð ólô^ˆ8ô ÷
	+ñ 	+ð! :ð !°%ó !ðHV¨#ð V¸cð VÀdó Vð2 #ð 2°cð 2¸dó 2ñ
�Sð 
 ð 
°ó 
ð$¨Eð $°e¸CÀ¸JÑ6Gô $r   