Ë
    *êñi4  ã                   ój  — d dl Z d dlmZ d dl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-m.Z. d dl/m0Z0 d dl1m2Z2 d dl3m4Z4 e5e6e7ee8   dz  ee8   f   f   Z9dgZ:d*de8de6de6fd„Z;	 d+dejx                  dz  defd„Z=dej|                  de?fd„Z@	 d*dedee8   de6dej|                  fd „ZAd!ede7e9ejx                  dz  f   fd"„ZB G d#„ d$e«      ZC	 d+d%ed&e6d'e(d(e!dz  def
d)„ZDy),é    N)ÚSequence)Úcast)Ú_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      úh/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/torch/distributed/checkpoint/optimizer.pyÚ_gen_rank_devicer)   8   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        z     |«      ¬«      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   Ústr)r+   Úpg_device_typeÚidxr1   s       r(   Ú_create_colwise_specr>   C   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Ò[ð
ˆ
ð 
ô ØÜœœ^¬cÑ1Ñ2°JÓ?ôð ùò
ùò

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   FTz1Cannot handle DTensor nested inside ShardedTensorzCannot handle nested DTensor)r5   r   ÚlenÚlocal_shardsÚtensorr   Ú
ValueErrorÚ_local_tensor)r?   s    r(   Ú_is_nested_tensorrF   W   s¿   € ÜˆCƒy”MÑ!Üˆs×ÑÓ!Ó" aÒ'ØÜ�× Ñ Ó" 1Ñ%×,Ñ,Ó-´Ñ>ØÜ�× Ñ Ó" 1Ñ%×,Ñ,Ó-´Ñ8ÜÐPÓQÐQð
 ô	 
ˆc‹”gÑ	ÜˆS×ÑÓ¤7Ñ*¬d°3×3DÑ3DÓ.EÌÑ.VäÐ7Ó8Ð8Ør*   Úpropsr8   c                 óP  — |dk(  r2t        t        j                  t        |«      j	                  «       «      }n-t        j                  |t        |«      j	                  «       «      }t        j
                  || j                  | j                  | j                  | j                  |¬«      S )Nr$   )r8   ÚdtypeÚlayoutÚrequires_gradÚ
pin_memoryÚdevice)
r   ÚtorchrM   r   Úcurrent_deviceÚemptyrI   rJ   rK   rL   )rG   r8   r!   rM   s       r(   Ú_alloc_tensorrQ   f   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t        d«      ‚t        |t        «      st        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   )Úitemsr8   rF   rA   rB   ÚAssertionErrorÚ
isinstancer   ÚmetadataÚshard_offsetsÚshard_sizesrC   Ú_process_group)rR   ÚspecsÚdp_pgÚkeyÚvalueÚshards         r(   Ú_get_state_dict_2d_layoutra   z   sÙ   € ð #%€EØ&*€EØ ×&Ñ&Ó(ò 0‰
ˆˆUØ˜EŸJ™J›LÐ)ˆˆc‰
Ü˜UÕ#Ü�u×)Ñ)Ó+Ó,°Ò1Ü$Ð%LÓMÐMÜ˜e¤]Ô3Ü$Ð%KÓLÐLØ×&Ñ&Ó(¨Ñ+ˆ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ÚtranslationrR   rX   Úfqn_to_offsetr"   Nc                 ól   •— t         ‰| �  «        || _        t        i «      | _        i | _        i | _        y ©N)ÚsuperÚ__init__re   r   rX   rR   rd   )Úselfre   Ú	__class__s     €r(   ri   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t        d«      ‚|j                  «       d   }t        t        j                  t        |j                  j                   |«      «      t        j                  |j                  j"                  «      ¬«      g}t%        |t'        t(        |«      |«      }|D ]�  }	|	j*                  j,                  €t        d«      ‚t/        |	j*                  j,                  |«      }
t1        j2                  |	j*                  t        j                  |
«      ¬«      }|| j                   |	j*                  <   Œ’ ||z  }�ŒÆ t5        |«      S )NrT   z Expected exactly one local shardr   )ÚoffsetsÚsizesz"dest_index.offset must not be None)Úoffset)rd   rR   rU   rX   Ústate_dict_metadatarW   r   r   re   rA   rB   rV   r   rN   ÚSizer   rY   rZ   r   r   r   Ú
dest_indexro   r   ÚdataclassesÚreplacer   )rj   ÚrequestsÚfqnÚobjÚmdro   Ú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Ò/Ü$Ð%GÓHÐHØ ×-Ñ-Ó/°Ñ2ˆNä$Ü!ŸJ™JÜ)¨.×*AÑ*A×*OÑ*OÐQWÓXóô  Ÿ*™* ^×%<Ñ%<×%HÑ%HÓIô	ðˆLô 4Ø”TÔ/°Ó4°lóˆDð
 ò A�Ø—=‘=×'Ñ'Ð/Ü(Ð)MÓNÐNÜ"3°B·M±M×4HÑ4HÈ&Ó"Q�Ü!,×!4Ñ!4Ø—M‘M¬%¯*©*°_Ó*Eô"�ð 3A�× Ñ  §¡Ò/ðAð ˜ÑŠHðM&	ôN ˜Ó!Ð!r*   Úindexc                 óV   •— t         ‰| �  | j                  j                  ||«      «      S rg   )rh   Úlookup_tensorrd   Úget)rj   r€   rk   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   Úintri   r   r   rN   ÚTensorr‚   Ú__classcell__)rk   s   @r(   rc   rc   �   sm   ø… Ø�m ]Ð2Ñ3Ó3ØÓàÓð d¨3°¸±Ð+=Ñ&>ð À4õ ð*" 8ó *"ðXI =ð I°U·\±\÷ Iñ Ir*   rc   Ú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-   r.   r   r/   z
<bytes_io>rT   )ÚrankÚ
world_sizeÚnum_devices_per_noder+   é   )rI   rJ   rK   Úmemory_formatrL   )rC   rX   )Úprocess_group)rR   rŽ   r�   )2Úread_metadatara   r2   r3   r4   r5   r   r6   r7   r   r&   Úappendr	   r>   rp   rU   Úplanner_datarW   r   r8   ÚnumelrQ   Ú
propertiesr   Úget_rankr   rƒ   ÚShardTensorPropertiesrI   rJ   rK   r•   rL   Úbuild_metadatarN   rq   Úshards_metadatar   r   Ú	placementr‘   r   rZ   r   Ú+_init_from_local_shards_and_global_metadatar   r‰   r   rc   r
   )rŒ   r�   rŽ   r�   rX   Úlayout_specsr]   Údp_pg_device_typer'   r1   ÚiÚdevice_infoÚsharding_specrR   re   r^   r_   Úkey_pathÚspec_keyÚ
alloc_sizer›   Úst_mdrB   Ú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*   )Úcudarg   )Ers   Úcollections.abcr   Útypingr   rN   Útorch.distributedÚdistributedr2   Ú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Š   ÚboolrF   rQ   ra   rc   r   © r*   r(   ú<module>rÊ      sª  ðó Ý $Ý ã Ý  Ý +Ý Eõõ @Ý XÝ JÝ K÷÷ ñ ÷ G÷õ KÝ >÷ñ õ
 BÝ LÝ :Ý ,ð ˜C  x°¡}°tÑ';¸XÀc¹]Ð'JÑ!KÐKÑLÐ ð
 (ð€ñ
 #ð °Cð ÀSó ð $(ñØ×Ñ˜DÑ ðàóð(˜5Ÿ<™<ð ¨Dó ð  FLñØðØ#+¨C¡=ðØ?Bðà
‡\�\óð( Øð à
Ð ×!2Ñ!2°TÑ!9Ð9Ñ:ó ôF:IÐ*ô :IðB #'ñ	NØ%ðNàðNð "ðNð ˜4Ñð	Nð
 ôNr*   