+
    &j4                        ^ RI t ^ RIHt ^ RIHt ^ RIt^ RIHt ^ RI	H
t
 ^ RIHt ^ RIHt ^ RIHt ^ RIHt ^ R	IHt ^ R
IHt ^ RIHtHtHtHtHtHtHt ^ RIH t H!t! ^ RI"H#t#H$t$ ^ RI%H&t& ^ RI'H(t( ^ RI)H*t*H+t+H,t, ^ RI-H.t. ^ RI/H0t0 ^ RI1H2t2 ^ RI3H4t4 ]5]6]7]]8,          R,          ]]8,          3,          3,          t9R.t:R$R R llt;R%R R llt<R R lt=R$R R llt>R R lt? ! R  R!]4      t@R%R" R# lltAR# )&    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_dictc                <    V ^8  d   QhR\         R\        R\        /# )   global_rankdevice_typereturn)intstr)formats   "~/Users/jameslopez/projects/CWCArchive/cwc-podcast/.venv/lib/python3.14/site-packages/torch/distributed/checkpoint/optimizer.py__annotate__r)   8   s!      # C S     c                     VR 8X  d   R # \        V4      pVP                  4       '       d!   \        WVP                  4       ,          4      # R # )cpu)r   is_availabler   device_count)r"   r#   device_modules   && r(   _gen_rank_devicer0   8   sI    e&{3M!!##%}'A'A'CC
 	
 r*   c                R    V ^8  d   QhR\         P                  R,          R\        /# )r!   pgNr$   )distProcessGroupr	   )r'   s   "r(   r)   r)   C   s'      D r*   c                    \         P                  P                  V 4      P                  pV f>   \	        \         P
                  ! 4       4       Uu. uF  pRV R\        W!4       2NK  	  ppML\	        V P                  4       4       Uu. uF)  pRV R\        \         P                  ! W4      V4       2NK+  	  pp\        ^ \        \        \        \        ,          ,          V4      R7      # u upi u upi )Nrank:/dim
placements)r3   distributed_c10d_get_pg_default_devicetyperangeget_world_sizer0   sizeget_global_rankr	   r   listr   r&   )r2   pg_device_typeidxr:   s   &   r(   _create_colwise_specrE   C   s     **AA"EJJN	z T0023
3 C5*3?@A3 	 

 RWWY'
' C5*4+?+?+H.YZ[' 	 
 ^c12J? 


s   C(/C-c                D    V ^8  d   QhR\         P                  R\        /# )r!   valr$   )torchTensorbool)r'   s   "r(   r)   r)   W   s      5<< D r*   c                    \        V 4      \        J d   \        V P                  4       4      ^ 8X  d   R# \        V P                  4       ^ ,          P                  4      \        J d   R# \        V P                  4       ^ ,          P                  4      \
        J d   \        R4      h R# \        V 4      \
        J dF   \        V P                  4      \
        J g   \        V P                  4      \        J d   \        R4      hR# )r   FTz1Cannot handle DTensor nested inside ShardedTensorzCannot handle nested DTensor)r=   r   lenlocal_shardstensorr   
ValueError_local_tensor)rG   s   &r(   _is_nested_tensorrQ   W   s    CyM!s!"a'  "1%,,->  "1%,,-8PQQ 9 	 
cg	S7*d33D3D.E.V788r*   c                r    V ^8  d   QhR\         R\        \        ,          R\        R\        P
                  /# )r!   propsr@   r#   r$   )r   r   r%   r&   rH   rI   )r'   s   "r(   r)   r)   f   s4      #+C=?B
\\r*   c           	      X   VR 8X  d3   \        \        P                  \        V4      P	                  4       4      pM.\        P                  ! V\        V4      P	                  4       4      p\        P
                  ! VV P                  V P                  V P                  V P                  VR7      # )r,   )r@   dtypelayoutrequires_grad
pin_memorydevice)
r   rH   rY   r   current_deviceemptyrU   rV   rW   rX   )rS   r@   r#   rY   s   &&& r(   _alloc_tensorr\   f   s     eell$6{$C$R$R$TU+K8GGI
 ;;kk||))## r*   c                t    V ^8  d   QhR\         R\        \        \        P                  R,          3,          /# )r!   
state_dictr$   N)r   tupleSTATE_DICT_2D_LAYOUTr3   r4   )r'   s   "r(   r)   r)   z   s2         
!2!2T!99: r*   c                   / pRpV P                  4        F  w  r4RVP                  4       3W&   \        V4      '       g   K,  \        VP	                  4       4      ^8X  g   \        R4      h\        V\        4      '       g   \        R4      hVP	                  4       ^ ,          pVP                  P                  VP                  P                  3W&   VP                  P                  pK  	  VV3# )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.
Nz%Cannot handle ST with multiple shardsz$Can only handle nested ShardedTensor)itemsr@   rQ   rL   rM   AssertionError
isinstancer   metadatashard_offsetsshard_sizesrN   _process_group)r^   specsdp_pgkeyvalueshards   &     r(   _get_state_dict_2d_layoutrn   z   s     #%E&*E &&(
EJJL)
U##u))+,1$%LMMe]33$%KLL&&(+E,,**EJ LL//E ) 	 r*   c                   t   a a ] tR t^t oV3R lV 3R lltV3R lR ltV3R lV 3R lltV3R ltRtVt	V ;t
# )	_ReaderWithOffsetc                J   < V ^8  d   QhRS[ S[S[S[,          3,          RR/# )r!   fqn_to_offsetr$   N)dictr&   r   r%   )r'   __classdict__s   "r(   r)   _ReaderWithOffset.__annotate__   s)      d3+=&> 4 r*   c                l   < \         SV `  4        Wn        \        / 4      V n        / V n        / V n        R # N)super__init__rr   r   re   r^   translation)selfrr   	__class__s   &&r(   ry   _ReaderWithOffset.__init__   s.    * r*   c                    < V ^8  d   QhRS[ /# )r!   r$   )r   )r'   rt   s   "r(   r)   ru      s     *" *"8 *"r*   c           	     >   . p/ V n         V P                  P                  4        EF  w  r#V P                  P                  V,          p\        V\        4      '       g   V\        W$V4      ,          pKN  W P                  9  d   V\        W$V4      ,          pKs  V P                  V,          p\        VP                  4       4      ^8X  g   \        R4      hVP                  4       ^ ,          p\        \        P                  ! \        VP                  P                   V4      4      \        P                  ! VP                  P"                  4      R7      .p\%        V\'        \(        V4      V4      pV F  p	V	P*                  P,                  f   \        R4      h\/        V	P*                  P,                  V4      p
\0        P2                  ! V	P*                  \        P                  ! V
4      R7      pWP                   V	P*                  &   K  	  W,          pEK  	  \5        V4      # )   z Expected exactly one local shard)offsetssizesz"dest_index.offset must not be None)offset)rz   r^   rb   re   state_dict_metadatard   r   r   rr   rL   rM   rc   r   rH   Sizer   rf   rg   r   r   r   
dest_indexr   r   dataclassesreplacer   )r{   requestsfqnobjmdr   original_shardlocal_chunksreqsrioriginal_offsetoriginal_indexs   &           r(   create_local_plan#_ReaderWithOffset.create_local_plan   s   --/HC2237Bc=11.s<<,,,.s<<'',Fs'')*a/$%GHH --/2N$!JJ).*A*A*O*OQWX  **^%<%<%H%HI	L 4T/4lD
 ==''/()MNN"3BMM4H4H&"Q!,!4!4MM%**_*E" 3A  /  HM 0N !!r*   c                :   < V ^8  d   QhRS[ RS[P                  /# )r!   indexr$   )r   rH   rI   )r'   rt   s   "r(   r)   ru      s#     I I= IU\\ Ir*   c                T   < \         SV `  V P                  P                  W4      4      # rw   )rx   lookup_tensorrz   get)r{   r   r|   s   &&r(   r   _ReaderWithOffset.lookup_tensor   s$    w$T%5%5%9%9%%GHHr*   c                T   < V ^8  d   Qh/ S[ S[S[3,          ;R&   S[;R&   S[;R&   # )r!   rz   r^   re   )rs   r   r   r   )r'   rt   s   "r(   r)   ru      s2     m]233   	 r*   )rr   re   r^   rz   )__name__
__module____qualname____firstlineno__ry   r   r   __annotate_func____static_attributes____classdictcell____classcell__)r|   rt   s   @@r(   rp   rp      s.      *" *"XI Is  r*   rp   c          
      b    V ^8  d   QhR\         R\        R\        R\        R,          R\         /# )r!   model_state_dictoptimizer_keystorage_readerplannerNr$   )r   r&   r   r   )r'   s   "r(   r)   r)      sF     N N%NN "N 4	N
 Nr*   c                   VP                  4       p\        V 4      w  rV\        P                  P	                  V4      P
                  p\        V4      pVfm   . p	\        \        P                  ! 4       4       F:  p
\        WzVP                  4       ,          4      pV	P                  RV
 RV 24       K<  	  \        ^ V	R7      pM\        V4      p/ p/ pVP                  P                  4        EF  w  ppVP                   V,          pV^ ,          V8w  d   K*  \#        V\$        4      '       d   RW&   KF  VP&                  P)                  4       ^8X  d&   \+        VP,                  VP&                  V4      W&   K  Vfp   \/        \+        VP,                  VP&                  V4      \        P0                  ! 4       \        P                  ! 4       VP                  4       \3        4       R7      W&   K  V^,          pVP5                  VRVP&                  34      ^,          p\7        VP,                  P8                  VP,                  P:                  VP,                  P<                  VP,                  P>                  VP,                  P@                  R7      pVPC                  \D        PF                  ! V4      V4      p. p\        P0                  ! V4      pVPH                   Fm  p\K        \L        VPN                  4      PQ                  4       V8w  d   K2  VP                  \S        \+        VP,                  VPT                  V4      VR7      4       Ko  	  \V        PX                  ! VVVR	7      pVV9   d>   VV,          ^ ,          e,   \K        \Z        \\        ,          VV,          ^ ,          4      W&   VW&   EK  	  \_        TTVe   \a        V4      MTR
7       \c        WP                   4      pV# )a3  
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)
Nr6   r7   r8   z
<bytes_io>)rank
world_sizenum_devices_per_noder2   )rU   rV   rW   memory_formatrX   )rN   re   )process_group)r^   r   r   )2read_metadatarn   r3   r;   r<   r=   r   r>   r?   r   r.   appendr	   rE   r   rb   planner_datard   r   r@   numelr\   
propertiesr   get_rankr   r   ShardTensorPropertiesrU   rV   rW   r   rX   build_metadatarH   r   shards_metadatar   r   	placementr   r   rg   r   +_init_from_local_shards_and_global_metadatar   r%   r   rp   r
   )r   r   r   r   re   layout_specsrj   dp_pg_device_typer/   r:   idevice_infosharding_specr^   rr   rk   rl   key_pathspec_key
alloc_sizer   st_mdrM   current_rankshard_mdsts   &&&&                      r(   r   r      sJ   j ++-H34DEL--DDUKPP&'89M}
t**,-A0!}'A'A'C#CK aS+78	 .
 *aJG,U3 #%J.0M2288:
U((-A;-'e122*JO ::"+  %**.?JO ]:e..

<MN]]_..0%2%?%?%A%'JO  {H%))(T5::4FGJJ.&&,,''..#..<<#..<< ++66J "00J1GTEL==/L!11(:(:;@@BlR##,!,,h.B.BDU  "*	 2 JJe5B <'L,B1,E,Q%)(3-h9OPQ9R%S" JOq ;v %494E!-07	 &j2G2GHJr*   )cudarw   )Br   collections.abcr   typingr   rH   torch.distributeddistributedr3   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   rs   r&   r_   r%   r`   __all__r0   rE   rQ   r\   rn   rp   r    r*   r(   <module>r      s     $     + E @ X J K   G K > 
 B L : , Cx}t';Xc]'J!KKL 
 (
(( F:I* :IzN Nr*   