Ë
    g^(hú'  ã            	       óL  — d dl Z d dlZd dlmZ d dlmZmZmZ d dlZd dl	m
Z d dlmZ d dlmZmZ d dlmZ d dlmZmZ d dlmZmZmZmZmZ d d	lmZmZmZ d d
l m!Z! d dl"m#Z# d dl$m%Z% d dl&m'Z' d dl(m)Z) g d¢Z* G d„ de'«      Z+ G d„ de«      Z,dee-ej\                  f   dee-ej\                  f   fd„Z/dee-ej\                  f   dee-ej\                  f   fd„Z0e1dk(  �r¿ G d„ de«      Z2 e jf                  «       Z4e4jk                  de-de2D � cg c]  } | jl                  ‘Œ c} e2jn                  ¬«       e4jk                  de-d¬ «       e4jk                  d!e-d"¬ «       e4jq                  «       Z9 e:d#e9jv                  › d$e9jx                  › d%e9jz                  › d&�«       d'e9jv                  › d(�Z>e9jz                  e2jn                  jl                  k(  rLej~                  j�                  e9jv                  «      r e0e9jv                  e9jx                  «       y e:e>«       ye9jz                  e2j‚                  jl                  k(  rLej~                  j…                  e9jv                  «      r e/e9jv                  e9jx                  «       y e:e>«       y eCd)e9jz                  › �«      ‚yc c} w )*é    N)ÚEnum)ÚcastÚOptionalÚUnion)Únarrow_tensor_by_index)ÚFileSystemReaderÚFileSystemWriter)Úflatten_state_dict)Ú_EmptyStateDictLoadPlannerÚDefaultLoadPlanner)ÚMetadataÚSTATE_DICT_TYPEÚSTORAGE_TYPESÚTensorPropertiesÚTensorStorageMetadata)ÚLoadItemTypeÚLoadPlanÚLoadPlanner)Ú_create_chunk_list)Ú_load_state_dict)Ú_save_state_dict)ÚStorageReader)ÚFuture)Údcp_to_torch_saveÚtorch_save_to_dcpÚBroadcastingTorchSaveReaderÚDynamicMetaLoadPlannerc                   ó  — e Zd ZdZ	 	 ddeeeej                  f      de	ddfd„Z
defd„Zded	eded   fd
„Zdededdfd„Zdedefd„Zdee   dee   fd„Zddeeej                  df   ddfd„Zedeeej                  f   defd„«       Zy)r   aI  
    StorageReader for reading a Torch Save file. This reader will read the entire checkpoint
    on the coordinator rank, and then broadcast and shard each tensor to all ranks.

    . N.B. Intended to be used with DynamicMetaLoadPlanner

    .. warning::
        Current implementation only supports loading Tensors.

    >>> # xdoctest: +SKIP("undefined vars")
    >>> sd = {"mode": model}
    >>> dcp.load(
    >>>    sd,
    >>>    storage_reader=BroadcastingTorchSaveReader(),
    >>>    planner=DynamicMetaLoadPlanner(),
    >>>    checkpoint_id="path_to_model.pt"
    >>> )
    NÚcheckpoint_idÚcoordinator_rankÚreturnc                 ó    — || _         || _        y ©N)r   r    )Úselfr   r    s      úg/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/distributed/checkpoint/format_utils.pyÚ__init__z$BroadcastingTorchSaveReader.__init__;   s   € ð
 +ˆÔØ 0ˆÕó    c                 ó   — t        i ¬«      S )zGExtends the default StorageReader to support building the metadata file©Ústate_dict_metadata)r   )r$   s    r%   Úread_metadataz)BroadcastingTorchSaveReader.read_metadataC   s   € ô ¨BÔ/Ð/r'   ÚplanÚplannerc           	      óª  — t        t        |«      }| j                  rK| j                  €J ‚t	        j
                  | j                  dd¬«      }|j                  rt        |«      \  }}nd}|j                  D �]¾  }|j                  t        j                  k(  r9t        d|j                  j                  › dt        | «      j                  › d�«      ‚| j                  rGt        j                   j#                  «       }||j                  j                     j%                  |«      }n6t	        j&                  |j(                  |j                  j                     «      }t        j*                  || j,                  d¬«       t/        ||j0                  |j2                  «      }|j5                  |«      j7                  «       }|j9                  «       |j9                  «       k(  s6J d	|j                  › d
|j9                  «       › d|j9                  «       › �«       ‚|j;                  |«       |j=                  ||«       �ŒÁ t?        «       }	|	jA                  d«       |	S )zê
        Reads torch save data on the coordinator rank, and broadcast afterwards
        this incurrs a communication cost, but avoids having to load
        the entire checkpoint on each rank, hopefully preventing OOM issues
        NÚcpuF)Úmap_locationÚweights_onlyúNon-tensor value identified at ú. At this time ú only supports loading Tensors.)ÚsrcÚasync_opzreq z mismatch sizes, z vs )!r   r   Úis_coordinatorr   ÚtorchÚloadr
   ÚitemsÚtyper   ÚBYTE_IOÚRuntimeErrorÚstorage_indexÚfqnÚ__name__ÚdistÚdistributed_c10dÚ_get_pg_default_deviceÚtoÚ
empty_likeÚ
state_dictÚ	broadcastr    r   Ústorage_offsetsÚlengthsÚresolve_tensorÚdetachÚsizeÚcopy_Úcommit_tensorr   Ú
set_result)
r$   r,   r-   Útorch_state_dictÚ_ÚreqÚ	pg_deviceÚtensorÚtarget_tensorÚfuts
             r%   Ú	read_dataz%BroadcastingTorchSaveReader.read_dataI   s  € ô Ô)¨7Ó3ˆð ×ÒØ×%Ñ%Ð1Ð1Ð1Ü$Ÿz™zØ×"Ñ"°ÀUô Ðð ×)Ò)Ü&8Ð9IÓ&JÑ#Ð ¡!à#Ðà—:‘:ó 	6ˆCØ�x‰xœ<×/Ñ/Ò/Ü"Ø5°c×6GÑ6G×6KÑ6KÐ5Lð M$Ü$(¨£J×$7Ñ$7Ð#8Ð8WðYóð ð ×"Ò"Ü ×1Ñ1×HÑHÓJ�	Ø)¨#×*;Ñ*;×*?Ñ*?Ñ@×CÑCÀIÓN‘ä×)Ñ)¨'×*<Ñ*<¸S×=NÑ=N×=RÑ=RÑ*SÓT�ä�N‰N˜6 t×'<Ñ'<ÀuÕMä+¨F°C×4GÑ4GÈÏÉÓUˆFØ#×2Ñ2°3Ó7×>Ñ>Ó@ˆMØ ×%Ñ%Ó'¨6¯;©;«=Ò8ð Ø�s×(Ñ(Ð)Ð):Ø ×%Ñ%Ó'Ð(¨¨V¯[©[«]¨Oð=óÐ8ð ×Ñ Ô'Ø×!Ñ! # }Ö5ð/	6ô2 “hˆØ�‰�tÔØˆ
r'   Úmetadatar7   c                 óŒ   — || _         | j                   r#t        j                  «       | j                  k(  sJ ‚| j                  €J ‚y©ú*Implementation of the StorageReader methodN)r7   rA   Úget_rankr    r   )r$   rX   r7   s      r%   Úset_up_storage_readerz1BroadcastingTorchSaveReader.set_up_storage_reader|   s?   € à,ˆÔØ×ÒÜ—=‘=“? d×&;Ñ&;Ò;Ð;Ð;à×!Ñ!Ð-Ð-Ñ-r'   c                 ó   — |S ©r[   © )r$   r,   s     r%   Úprepare_local_planz.BroadcastingTorchSaveReader.prepare_local_plan„   s   € àˆr'   Úglobal_planc                 ó   — |S r_   r`   )r$   rb   s     r%   Úprepare_global_planz/BroadcastingTorchSaveReader.prepare_global_planˆ   s   € àÐr'   c                 ó   — || _         yrZ   )r   )r$   r   s     r%   Úresetz!BroadcastingTorchSaveReader.resetŒ   s
   € à*ˆÕr'   c                 ó@   — t         j                  j                  |«      S r_   )ÚosÚpathÚisfile)Úclsr   s     r%   Úvalidate_checkpoint_idz2BroadcastingTorchSaveReader.validate_checkpoint_id�   s   € ô �w‰w�~‰~˜mÓ,Ð,r'   )Nr   r#   )r@   Ú
__module__Ú__qualname__Ú__doc__r   r   Ústrrh   ÚPathLikeÚintr&   r   r+   r   r   r   rW   Úboolr]   ra   Úlistrd   rf   Úclassmethodrl   r`   r'   r%   r   r   '   s  „ ñð* <@Ø !ñ1à  c¨2¯;©;Ð&6Ñ 7Ñ8ð1ð ð1ð 
ó	1ð0˜xó 0ð1˜hð 1°ð 1ÀÈÁó 1ðf.¨hð .Èð .ÐQUó .ð xð °Hó ð¨t°H©~ð À$ÀxÁ.ó ñ+ 5¨¨b¯k©k¸4Ð)?Ñ#@ð +ÈDó +ð ð-°5¸¸b¿k¹kÐ9IÑ3Jð -Ètò -ó ñ-r'   r   c            	       ó@   ‡ — e Zd ZdZ	 	 ddedee   deddfˆ fd„Zˆ xZ	S )	r   aœ  
    Extension of DefaultLoadPlanner, which creates a new Metadata object based on the passed in state dict,
    avoiding the need to read metadata from disk. This is useful when reading formats which don't have a
    metadata file, like Torch Save files.

    . N.B. Intended to be used with BroadcastingTorchSaveReader

    .. warning::
        Current implementation only supports loading Tensors.

    >>> # xdoctest: +SKIP("undefined vars")
    >>> sd = {"mode": model}
    >>> dcp.load(
    >>>    sd,
    >>>    storage_reader=BroadcastingTorchSaveReader(),
    >>>    planner=DynamicMetaLoadPlanner(),
    >>>    checkpoint_id="path_to_model.pt"
    >>> )
    NrF   rX   r7   r!   c           	      ó|  •— t         ‰| �  |||«       i }| j                  j                  «       D ]z  \  }}t	        j
                  |«      s%t        d|› dt        | «      j                  › d�«      ‚t        t        |j                  ¬«      |j                  «       t        |«      «      ||<   Œ| t        |¬«      | _        y)zdSetups of the planner, extnding default behavior by creating the Metadata object from the state dictr2   r3   r4   )Údtyper)   N)ÚsuperÚset_up_plannerrF   r:   r8   Ú	is_tensorr=   r;   r@   r   r   rx   rL   r   r   rX   )r$   rF   rX   r7   r*   ÚkeyrT   Ú	__class__s          €r%   rz   z%DynamicMetaLoadPlanner.set_up_planner«   s¸   ø€ ô 	‰Ñ˜z¨8°^ÔDà8:ÐØŸ?™?×0Ñ0Ó2ò 	‰KˆC�Ü—?‘? 6Ô*Ü"Ø5°c°Uð ;$Ü$(¨£J×$7Ñ$7Ð#8Ð8WðYóð ô
 (=Ü  v§|¡|Ô4Ø—‘“Ü" 6Ó*ó(Ð Ò$ð	ô !Ð5HÔIˆ�r'   )NF)
r@   rm   rn   ro   r   r   r   rs   rz   Ú__classcell__)r}   s   @r%   r   r   –   sK   ø„ ñð. (,Ø$ñ	Jà#ðJð ˜8Ñ$ðJð ð	Jð
 
÷Jñ Jr'   r   Údcp_checkpoint_dirÚtorch_save_pathc                 ót   — i }t        |t        | «      t        «       d¬«       t        j                  ||«       y)aq  
    Given a directory containing a DCP checkpoint, this function will convert it into a
    Torch save file.

    Args:
        dcp_checkpoint_dir: Directory containing the DCP checkpoint.
        torch_save_path: Filename to store the converted Torch save file.

    .. warning::
        To avoid OOM, it's recommended to only run this function on a single rank.
    T)Ústorage_readerr-   Úno_distN)r   r   r   r8   Úsave)r   r€   Úsds      r%   r   r   Ä   s6   € ð €BÜØ
Ü'Ð(:Ó;Ü*Ó,Øõ	ô 
‡J�Jˆr�?Õ#r'   c                 ó`   — t        j                  | d¬«      }t        |t        |«      d¬«       y)aB  
    Given the location of a torch save file, converts it into a DCP checkpoint.

    Args:
        torch_save_path: Filename of the Torch save file.
        dcp_checkpoint_dir: Directory to store the DCP checkpoint.

    .. warning::
        To avoid OOM, it's recommended to only run this function on a single rank.
    F)r1   T)Ústorage_writerrƒ   N)r8   r9   r   r	   )r€   r   rF   s      r%   r   r   Ý   s-   € ô —‘˜O¸%Ô@€Jô ØÔ#3Ð4FÓ#GÐQUör'   Ú__main__c                   ó   — e Zd ZdZdZy)Ú
FormatModeÚtorch_to_dcpÚdcp_to_torchN)r@   rm   rn   ÚTORCH_TO_DCPÚDCP_TO_TORCHr`   r'   r%   rŠ   rŠ   ö   s   „ Ø%ˆØ%‰r'   rŠ   ÚmodezConversion mode)r;   ÚhelpÚchoicesÚdefaultr5   zPath to the source model)r;   r�   ÚdstzPath to the destination modelzConverting checkpoint from z to z using method: 'ú'zNo checkpoint found at z. Skipping conversion.zUnknown conversion mode: )DÚargparserh   Úenumr   Útypingr   r   r   r8   Útorch.distributedÚdistributedrA   Útorch.distributed._shard._utilsr   Útorch.distributed.checkpointr   r	   Ú)torch.distributed.checkpoint._nested_dictr
   Ú,torch.distributed.checkpoint.default_plannerr   r   Ú%torch.distributed.checkpoint.metadatar   r   r   r   r   Ú$torch.distributed.checkpoint.plannerr   r   r   Ú,torch.distributed.checkpoint.planner_helpersr   Ú.torch.distributed.checkpoint.state_dict_loaderr   Ú-torch.distributed.checkpoint.state_dict_saverr   Ú$torch.distributed.checkpoint.storager   Útorch.futuresr   Ú__all__r   r   rp   rq   r   r   r@   rŠ   ÚArgumentParserÚparserÚadd_argumentÚvaluer�   Ú
parse_argsÚargsÚprintr5   r“   r�   Úcheckpoint_missing_warningri   rj   rŽ   ÚisdirÚ
ValueError)Úms   0r%   ú<module>r±      se  ðã Û 	Ý ß (Ñ (ã Ý  Ý Bß KÝ H÷÷õ ÷ UÑ TÝ KÝ KÝ JÝ >Ý  ò€ôl- -ô l-ô^+JÐ/ô +Jð\$Ø˜c 2§;¡;Ð.Ñ/ð$à˜3 §¡Ð+Ñ,ó$ð2Ø˜3 §¡Ð+Ñ,ðà˜c 2§;¡;Ð.Ñ/óð. ˆzÓô&�Tô &ð
 %ˆX×$Ñ$Ó&€FØ
×ÑØØØØ",Ö-˜Q�—“Ò-Ø×'Ñ'ð ô ð ×Ñ˜ CÐ.HÐÔIØ
×Ñ˜ CÐ.MÐÔNØ×ÑÓ€Dá	Ø
% d§h¡h Z¨t°D·H±H°:Ð=MÈdÏiÉiÈ[ÐXYÐZôð " $§(¡( Ð+AÐBð ð ‡y�y�J×+Ñ+×1Ñ1Ò1Ø�7‰7�>‰>˜$Ÿ(™(Ô#Ù˜dŸh™h¨¯©Õ1áÐ,Õ-Ø	�‰�j×-Ñ-×3Ñ3Ò	3Ø�7‰7�=‰=˜Ÿ™Ô"Ù˜dŸh™h¨¯©Õ1áÐ,Õ-áÐ4°T·Y±Y°KÐ@ÓAÐAðI ùò .s   ÄJ!