Ë
    g^(hb3  ã                   ór  — 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 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 d dlmZmZmZmZmZmZm Z  d dl!m"Z"m#Z# d dl$m%Z%m&Z& d dl'm(Z( d dl)m*Z* d dl+m,Z,m-Z-m.Z. d dl/m0Z0 d dl1m2Z2 d dl3m4Z4 d dl5m6Z6 e7e8e9eee:      ee:   f   f   Z;dgZ<d*de:de8de8fd„Z=	 d+dee
j|                     defd„Z?dej€                  deAfd„ZB	 d*dedee:   de8dej€                  fd „ZCd!ede9e;ee
j|                     f   fd"„ZD G d#„ d$e«      ZE	 d+d%ed&e8d'e*d(ee#   def
d)„ZFy),é    N)ÚSequence)ÚcastÚOptionalÚUnion)Ú_get_device_module)ÚShardedTensor)ÚTensorProperties)ÚShard)ÚChunkShardingSpec)Úunflatten_state_dict)ÚDefaultLoadPlanner)ÚBytesStorageMetadataÚChunkStorageMetadataÚMetadataÚMetadataIndexÚSTATE_DICT_TYPEr	   ÚTensorStorageMetadata)ÚLoadPlanÚLoadPlanner)Ú_create_read_itemsÚ create_read_items_for_chunk_list)Úload_state_dict)ÚStorageReader)Ú_element_wise_addÚ_element_wise_subÚ_normalize_device_info)Ú_get_default_group)Ú_create_chunk_sharded_tensor)Ú_remote_device)ÚDTensorÚ!load_sharded_optimizer_state_dictÚglobal_rankÚdevice_typeÚreturnc                 ó€   — |dk(  ryt        |«      }|j                  «       rt        || |j                  «       z  «      S y)NÚcpu)r   Úis_availabler   Údevice_count)r"   r#   Údevice_modules      úd/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/distributed/checkpoint/optimizer.pyÚ_gen_rank_devicer+   6   sH   € Ø�eÒØÜ& {Ó3€MØ×!Ñ!Ô#Ü%Ø˜ }×'AÑ'AÓ'CÑCó
ð 	
ð ó    Úpgc                 óÔ  — t         j                  j                  | «      j                  }| €;t	        t        j
                  «       «      D �cg c]  }d|› dt        ||«      › �‘Œ }}nJt	        | j                  «       «      D �cg c](  }d|› dt        t        j                  | |«      |«      › �‘Œ* }}t        dt        t        t        t        t        f      |«      ¬«      S c c}w c c}w )Núrank:ú/r   ©ÚdimÚ
placements)ÚdistÚdistributed_c10dÚ_get_pg_default_deviceÚtypeÚrangeÚget_world_sizer+   ÚsizeÚget_global_rankr   r   Úlistr   r   Ústr)r-   Úpg_device_typeÚidxr3   s       r*   Ú_create_colwise_specr@   A   sê   € ô ×*Ñ*×AÑAÀ"ÓE×JÑJ€NØ	€zô œT×0Ñ0Ó2Ó3ö
àð �C�5˜Ô*¨3°Ó?Ð@ÒAð
ˆ
ñ 
ô ˜RŸW™W›YÓ'ö
àð �C�5˜Ô*¬4×+?Ñ+?ÀÀCÓ+HÈ.ÓYÐZÒ[ð
ˆ
ð 
ô ØÜœœU¤>´3Ð#6Ñ7Ñ8¸*ÓEôð ùò
ùò

s   ÁC Â-C%Úvalc                 óÎ  — t        | «      t        u r‚t        | j                  «       «      dk(  ryt        | j                  «       d   j                  «      t        u ryt        | j                  «       d   j                  «      t
        u rt        d«      ‚yt        | «      t
        u rAt        | j                  «      t
        u st        | j                  «      t        u rt        d«      ‚y)Nr   FTz2Cannot handle DTensor nested insided ShardedTensorzCannot handle nested DTensor)r7   r   ÚlenÚlocal_shardsÚtensorr    Ú
ValueErrorÚ_local_tensor)rA   s    r*   Ú_is_nested_tensorrH   U   s¿   € ÜˆCƒy”MÑ!Üˆs×ÑÓ!Ó" aÒ'ØÜ�× Ñ Ó" 1Ñ%×,Ñ,Ó-´Ñ>ØÜ�× Ñ Ó" 1Ñ%×,Ñ,Ó-´Ñ8ÜÐQÓRÐRð
 ô	 
ˆc‹”gÑ	ÜˆS×ÑÓ¤7Ñ*¬d°3×3DÑ3DÓ.EÌÑ.VäÐ7Ó8Ð8Ør,   Úpropsr:   c                 óP  — |dk(  r2t        t        j                  t        |«      j	                  «       «      }n-t        j                  |t        |«      j	                  «       «      }t        j
                  || j                  | j                  | j                  | j                  |¬«      S )Nr&   )r:   ÚdtypeÚlayoutÚrequires_gradÚ
pin_memoryÚdevice)
r   ÚtorchrO   r   Úcurrent_deviceÚemptyrK   rL   rM   rN   )rI   r:   r#   rO   s       r*   Ú_alloc_tensorrS   d   s†   € ð �eÒÜ”e—l‘lÔ$6°{Ó$C×$RÑ$RÓ$TÓU‰ä—‘ØÔ+¨KÓ8×GÑGÓIó
ˆô �;‰;ØØ�k‰kØ�|‰|Ø×)Ñ)Ø×#Ñ#Øôð r,   Ú
state_dictc                 ó¨  — i }d}| j                  «       D ]¸  \  }}d|j                  «       f||<   t        |«      sŒ't        |j	                  «       «      dk(  sJ d«       ‚t        |t        «      sJ d«       ‚|j	                  «       d   }|j                  j                  |j                  j                  f||<   |j                  j                  }Œº ||fS )a+  
    Load the right TP slice of the optimizer state.

    This is not easy since the per-tensor slicing can't be inferred from checkpoint metadata.
    We take advantage of the model state_dict producing a sliced ST to figure out what we need to load.
    This is pretty fragile and it might be easier for FSDP to compute this info for us.
    Returns a dictionary where keys are the same of the state_dict and the value is a tuple of
    (offset, size) for the current rank TP slice.
    N.B. The state_dict *MUST* come from FSDP.sharded_state_dict.
    Né   z%Cannot handle ST with multiple shardsz$Can only handle nested ShardedTensorr   )Úitemsr:   rH   rC   rD   Ú
isinstancer   ÚmetadataÚshard_offsetsÚshard_sizesrE   Ú_process_group)rT   ÚspecsÚdp_pgÚkeyÚvalueÚshards         r*   Ú_get_state_dict_2d_layoutrb   x   så   € ð #%€EØ)-€EØ ×&Ñ&Ó(ò 0‰
ˆˆUØ˜EŸJ™J›LÐ)ˆˆc‰
Ü˜UÕ#Ü�u×)Ñ)Ó+Ó,°Ò1ð Ø7óÐ1ô ˜e¤]Ô3ð Ø6óÐ3ð ×&Ñ&Ó(¨Ñ+ˆEà—‘×,Ñ,Ø—‘×*Ñ*ðˆE�#‰Jð —L‘L×/Ñ/‰Eð0ð" 	Øðð r,   c                   ó–   ‡ — e Zd ZU eeef   ed<   eed<   eed<   deee	e
   f   ddfˆ fd„Zdefd„Zd	edej                  fˆ fd
„Zˆ xZS )Ú_ReaderWithOffsetÚtranslationrT   rY   Úfqn_to_offsetr$   Nc                 ól   •— t         ‰| �  «        || _        t        i «      | _        i | _        i | _        y ©N)ÚsuperÚ__init__rf   r   rY   rT   re   )Úselfrf   Ú	__class__s     €r*   rj   z_ReaderWithOffset.__init__¢   s0   ø€ Ü‰ÑÔØ*ˆÔÜ  ›ˆŒØˆŒØˆÕr,   c           	      óÈ  — g }i | _         | j                  j                  «       D �]±  \  }}| j                  j                  |   }t        |t        «      s|t        |||«      z  }ŒA|| j                  vr|t        |||«      z  }Œ`| j                  |   }t        |j                  «       «      dk(  sJ ‚|j                  «       d   }t        t        j                  t        |j                  j                  |«      «      t        j                  |j                  j                   «      ¬«      g}t#        |t%        t&        |«      |«      }|D ]‡  }	|	j(                  j*                  €J ‚t-        |	j(                  j*                  |«      }
t/        j0                  |	j(                  t        j                  |
«      ¬«      }|| j                   |	j(                  <   Œ‰ ||z  }�Œ´ t3        |«      S )NrV   r   )ÚoffsetsÚsizes)Úoffset)re   rT   rW   rY   Ústate_dict_metadatarX   r   r   rf   rC   rD   r   rP   ÚSizer   rZ   r[   r   r   r   Ú
dest_indexrp   r   ÚdataclassesÚreplacer   )rk   ÚrequestsÚfqnÚobjÚmdrp   Úoriginal_shardÚlocal_chunksÚreqsÚriÚoriginal_offsetÚoriginal_indexs               r*   Úcreate_local_planz#_ReaderWithOffset.create_local_plan©   sÂ  € ØˆØˆÔØŸ™×-Ñ-Ó/ó $	‰HˆC�Ø—‘×2Ñ2°3Ñ7ˆBÜ˜c¤=Ô1ØÔ.¨s°B¸Ó<Ñ<�Øà˜$×,Ñ,Ñ,ØÔ.¨s°B¸Ó<Ñ<�Øà×'Ñ'¨Ñ,ˆFä�s×'Ñ'Ó)Ó*¨aÒ/Ð/Ð/Ø ×-Ñ-Ó/°Ñ2ˆNä$Ü!ŸJ™JÜ)¨.×*AÑ*A×*OÑ*OÐQWÓXóô  Ÿ*™* ^×%<Ñ%<×%HÑ%HÓIô	ðˆLô 4Ø”TÔ/°Ó4°lóˆDð
 ò A�Ø—}‘}×+Ñ+Ð7Ð7Ð7Ü"3°B·M±M×4HÑ4HÈ&Ó"Q�Ü!,×!4Ñ!4Ø—M‘M¬%¯*©*°_Ó*Eô"�ð 3A�× Ñ  §¡Ò/ðAð ˜ÑŠHðI$	ôJ ˜Ó!Ð!r,   Úindexc                 óV   •— t         ‰| �  | j                  j                  ||«      «      S rh   )ri   Úlookup_tensorre   Úget)rk   r�   rl   s     €r*   rƒ   z_ReaderWithOffset.lookup_tensorÓ   s&   ø€ Ü‰wÑ$ T×%5Ñ%5×%9Ñ%9¸%ÀÓ%GÓHÐHr,   )Ú__name__Ú
__module__Ú__qualname__Údictr   Ú__annotations__r   r   r=   r   Úintrj   r   r€   rP   ÚTensorrƒ   Ú__classcell__)rl   s   @r*   rd   rd   �   sm   ø… Ø�m ]Ð2Ñ3Ó3ØÓØÓð d¨3°¸±Ð+=Ñ&>ð À4õ ð(" 8ó ("ðTI =ð I°U·\±\÷ Iñ Ir,   rd   Úmodel_state_dictÚoptimizer_keyÚstorage_readerÚplannerc                 ó.  — |j                  «       }t        | «      \  }}t        j                  j	                  |«      j
                  }t        |«      }|€fg }	t        t        j                  «       «      D ]6  }
t        ||
|j                  «       z  «      }|	j                  d|
› d|› �«       Œ8 t        d|	¬«      }nt        |«      }i }i }|j                  j                  «       D �]|  \  }}|j                   |   }|d   |k7  rŒt#        |t$        «      rd||<   Œ5|j&                  j)                  «       dk(  r%t+        |j,                  |j&                  |«      ||<   Œw|€mt/        t+        |j,                  |j&                  |«      t        j0                  «       t        j                  «       |j                  «       t3        «       ¬«      ||<   Œæ|d	   }|j5                  |d|j&                  f«      d   }t7        |j,                  j8                  |j,                  j:                  |j,                  j<                  |j,                  j>                  |j,                  j@                  ¬
«      }|jC                  tE        jF                  |«      |«      }g }t        j0                  |«      }|jH                  D ]i  }tK        tL        |jN                  «      jQ                  «       |k7  rŒ/|j                  tS        t+        |j,                  |jT                  |«      |¬«      «       Œk tW        jX                  |||¬«      }||v r(||   d   � tK        tZ        t\           ||   d   «      ||<   |||<   �Œ t_        |||�ta        |«      n|¬«       tc        ||j                   «      }|S )aç  
    Load a state_dict in conjunction with FSDP sharded optimizer state.

    This is the current recommended way to checkpoint FSDP.
    >>> # xdoctest: +SKIP
    >>> import torch.distributed.checkpoint as dist_cp
    >>> # Save
    >>> model: torch.nn.Model
    >>> optim_params = model.parameters()
    >>> optim = torch.optim.SGD(optim_params, lr=0.01)
    >>> # Save
    >>> with FSDP.state_dict_type(model, StateDictType.SHARDED_STATE_DICT):
    >>>     state_dict = {
    >>>         "optimizer": FSDP.optim_state_dict(model, optim),
    >>>         "model": model.state_dict()
    >>>     }
    >>>     dist_cp.save_state_dict(
    >>>         state_dict=optim_state,
    >>>         storage_writer=dist_cp.FileSystemWriter("checkpoint"),
    >>>         planner=dist_cp.DefaultSavePlanner(),
    >>>     )
    >>>
    >>> # Load
    >>> with FSDP.state_dict_type(model_tp, StateDictType.SHARDED_STATE_DICT):
    >>>     model_state_dict = model_tp.state_dict()
    >>>     checkpoint = {
    >>>         "model": model_state_dict
    >>>     }
    >>>     dist_cp.load_state_dict(
    >>>         state_dict=checkpoint,
    >>>         storage_reader=dist_cp.FileSystemReader(checkpoint_file),
    >>>         planner=dist_cp.DefaultLoadPlanner(),
    >>>     )
    >>>     model.load_state_dict(checkpoint["model_state"])
    >>>
    >>>     optim_state = dist_cp.load_sharded_optimizer_state_dict(
    >>>         model_state_dict,
    >>>         optimizer_key="optimizer",
    >>>         storage_reader=dist_cp.FileSystemReader("checkpoint"),
    >>>     )
    >>>
    >>>     flattened_osd = FSDP.optim_state_dict_to_load(
    >>>        model, optim, optim_state["optimizer"]
    >>>     )
    >>>
    >>>     optim.load_state_dict(flattened_osd)
    Nr/   r0   r   r1   z
<bytes_io>rV   )ÚrankÚ
world_sizeÚnum_devices_per_noder-   é   )rK   rL   rM   Úmemory_formatrN   )rE   rY   )Úprocess_group)rT   r�   r�   )2Úread_metadatarb   r4   r5   r6   r7   r   r8   r9   r   r(   Úappendr   r@   rq   rW   Úplanner_datarX   r   r:   ÚnumelrS   Ú
propertiesr   Úget_rankr   r„   ÚShardTensorPropertiesrK   rL   rM   r–   rN   Úbuild_metadatarP   rr   Úshards_metadatar   r   Ú	placementr’   r
   r[   r   Ú+_init_from_local_shards_and_global_metadatar   rŠ   r   rd   r   )r�   rŽ   r�   r�   rY   Úlayout_specsr^   Údp_pg_device_typer)   r3   ÚiÚdevice_infoÚsharding_specrT   rf   r_   r`   Úkey_pathÚspec_keyÚ
alloc_sizerœ   Úst_mdrD   Úcurrent_rankÚshard_mdÚsts                             r*   r!   r!   ×   sf  € ðj ×+Ñ+Ó-€Hä3Ð4DÓEÑ€L�%Ü×-Ñ-×DÑDÀUÓK×PÑPÐÜ&Ð'8Ó9€Mà€}Øˆ
Ü”t×*Ñ*Ó,Ó-ò 	9ˆAÜ0Ø! 1 }×'AÑ'AÓ'CÑ#CóˆKð ×Ñ  a S¨¨+¨Ð7Õ8ð		9ô
 *¨a¸JÔG‰ä,¨UÓ3ˆð #%€Jà.0€MØ×2Ñ2×8Ñ8Ó:ó 8!‰
ˆˆUØ×(Ñ(¨Ñ-ˆØ�A‰;˜-Ò'Øä�eÔ1Ô2Ø*ˆJ�s‰OØð �:‰:×ÑÓ Ò"Ü+Ø× Ñ  %§*¡*Ð.?óˆJ�sŠOð ˆ]Ü:Ü˜e×.Ñ.°·
±
Ð<MÓNÜ—]‘]“_Ü×.Ñ.Ó0Ø%2×%?Ñ%?Ó%AÜ%Ó'ôˆJ�sŠOð   ‘{ˆHØ%×)Ñ)¨(°T¸5¿:¹:Ð4FÓGÈÑJˆJä.Ø×&Ñ&×,Ñ,Ø×'Ñ'×.Ñ.Ø#×.Ñ.×<Ñ<Ø#×.Ñ.×<Ñ<Ø ×+Ñ+×6Ñ6ôˆJð "×0Ñ0´·±¸JÓ1GÈÓTˆEØˆLÜŸ=™=¨Ó/ˆLØ!×1Ñ1ò 
�Üœ¨×(:Ñ(:Ó;×@Ñ@ÓBÀlÒRØØ×#Ñ#ÜÜ,Ø!×,Ñ,¨h×.BÑ.BÐDUó ð "*ô	õð
ô ×JÑJØ˜e°5ôˆBð ˜<Ñ'¨L¸Ñ,BÀ1Ñ,EÐ,QÜ%)¬(´3©-¸ÀhÑ9OÐPQÑ9RÓ%S�˜cÑ"à ˆJ�s‹Oðq8!ôv ØØ%à49Ð4EÔ! -Ô0È7õ	ô & j°(×2GÑ2GÓH€JàÐr,   )Úcudarh   )Grt   Úcollections.abcr   Útypingr   r   r   rP   Útorch.distributedÚdistributedr4   Útorch._utilsr   Ú+torch.distributed._shard.sharded_tensor.apir   Ú0torch.distributed._shard.sharded_tensor.metadatar	   rž   Ú-torch.distributed._shard.sharded_tensor.shardr
   Ú:torch.distributed._shard.sharding_spec.chunk_sharding_specr   Ú)torch.distributed.checkpoint._nested_dictr   Ú,torch.distributed.checkpoint.default_plannerr   Ú%torch.distributed.checkpoint.metadatar   r   r   r   r   r   Ú$torch.distributed.checkpoint.plannerr   r   Ú,torch.distributed.checkpoint.planner_helpersr   r   Ú.torch.distributed.checkpoint.state_dict_loaderr   Ú$torch.distributed.checkpoint.storager   Ú"torch.distributed.checkpoint.utilsr   r   r   Ú"torch.distributed.distributed_c10dr   Ú#torch.distributed.fsdp._shard_utilsr   Útorch.distributed.remote_devicer   Útorch.distributed.tensorr    rˆ   r=   ÚtuplerŠ   ÚSTATE_DICT_2D_LAYOUTÚ__all__r+   ÚProcessGroupr@   r‹   ÚboolrH   rS   rb   rd   r!   © r,   r*   ú<module>rË      s­  ðó Ý $ß (Ñ (ã Ý  Ý +Ý Eõõ @Ý XÝ JÝ K÷÷ ñ ÷ G÷õ KÝ >÷ñ õ
 BÝ LÝ :Ý ,ð ˜C  x°¸±Ñ'>ÀÈÁÐ'MÑ!NÐNÑOÐ ð
 (ð€ñ
 #ð °Cð ÀSó ð '+ñØ�×"Ñ"Ñ#ðàóð(˜5Ÿ<™<ð ¨Dó ð  FLñØðØ#+¨C¡=ðØ?Bðà
‡\�\óð("Øð"à
Ð ¨$×*;Ñ*;Ñ!<Ð<Ñ=ó"ôJ7IÐ*ô 7Ið| &*ñ	NØ%ðNàðNð "ðNð �kÑ"ð	Nð
 ôNr,   