Ë
      çih  ã                   ój  — d dl Z d dlZd dl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mZ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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( d dl)m*Z* d dl+m,Z, d dl-m.Z.m/Z/m0Z0m1Z1 d dl2m3Z3m4Z4m5Z5m6Z6m7Z7 d dl2m8Z9 d dl:m;Z;m<Z< d dl=m>Z> d dl?m@Z@mAZAmBZB d dlCmZ d dlDmEZE d dlFmGZGmHZH erd dlImJZJ  ede¬«      ZK G d„ de,«      ZL G d„ de0«      ZM G d „ d!e«      ZN	 d;d"e
d#eOePeeeef   f   d$eQd%eRd&eeOePeePegeQf   f      d'dfd(„ZS	 	 	 d<d"e
d#eOePeeeef   f   d)eQd*eQd+eeQ   d'eOePef   fd,„ZTd-eRd.eRd/eRd0ejª                  d'd1f
d2„ZVd3eWd'ee   fd4„ZXd=d"e
d3ed/eRd)eQd'df
d5„ZY	 d>d6eOePef   d3ed/eRd)eQd'df
d7„ZZd3ed'efd8„Z[d9eOePef   d3ed'eOePef   fd:„Z\y)?é    N)Ú	Generator)ÚAbstractContextManagerÚ	ExitStack)Ú	timedelta)ÚPath)ÚTYPE_CHECKINGÚAnyÚCallableÚLiteralÚOptionalÚTypeVarÚUnion)Úrank_zero_only)ÚTensor)ÚModule)Ú	Optimizer)Ú	TypeGuardÚoverride)ÚCheckpointIO)Údefault_pg_timeout)Ú_distributed_checkpoint_loadÚ_distributed_checkpoint_saveÚ_get_full_state_dict_contextÚ_is_full_checkpointÚ_is_sharded_checkpoint)Ú_SubprocessScriptLauncher)ÚParallelStrategy)Ú
TBroadcastÚ_apply_filterÚ_BackwardSyncControlÚ!_validate_keys_for_strict_loading)ÚReduceOpÚ_distributed_is_initializedÚ-_get_default_process_group_backend_for_deviceÚ_init_dist_connectionÚ_sync_ddp_if_available©Úgroup)Ú_TORCH_GREATER_EQUAL_2_3Ú_TORCH_GREATER_EQUAL_2_4)Ú_materialize_distributed_module)Ú_METADATA_FILENAMEÚ
_lazy_loadÚ_move_state_into)Ú
reset_seed)Ú_PATHÚ	_Stateful)Ú
DeviceMeshÚTModel)Úboundc                   ó¤  ‡ — e Zd ZdZddddefdeedgef   deed   e	f   deed   e	f   d	e
d
ee   dee   ddfˆ fd„Zed5d„«       Zeedefd„«       «       Zej(                  ededdfd„«       «       Zeedej,                  fd„«       «       Zede	fd„«       Zej(                  de	ddfd„«       Zede	fd„«       Zeedeeef   fd„«       «       Zedee   fd„«       Zed6d„«       Zed6ˆ fd„«       Zede de fd„«       Z!ede ddfd„«       Z"ed7dee
   de#fd„«       Z$e	 d8d e%d!ee   d"eee&ef      de%fd#„«       Z'ed$ed%eddfd&„«       Z(ed9d'e)d(e	de)fd)„«       Z*e	 	 d:d*e+d+eeee e,ef   f   d,ee   d-eeeeeege
f   f      ddf
d.„«       Z-e	 	 	 d;d*e+d+eee e,eeee e,ef   f   f      d/e
d0ee
   deeef   f
d1„«       Z.d6d2„Z/defd3„Z0d6d4„Z1ˆ xZ2S )<ÚModelParallelStrategyaß  Enables user-defined parallelism applied to a model.

    .. warning::  This is an :ref:`experimental <versioning:Experimental API>` feature.

    Currently supports up to 2D parallelism. Specifically, it supports the combination of
    Fully Sharded Data-Parallel 2 (FSDP2) with Tensor Parallelism (DTensor). These PyTorch APIs are currently still
    experimental in PyTorch. Requires PyTorch 2.4 or newer.

    Arguments:
        parallelize_fn: A function that applies parallelisms to a module. The strategy will provide the
            model and device mesh as input.
        data_parallel_size: The number of devices within a data-parallel group. Defaults to ``"auto"``, which
            sets this size to the number of nodes in the cluster.
        tensor_parallel_size: The number of devices within a tensor-parallel group. Defaults to ``"auto"``, which
            sets this size to the number of GPUs in a single node.
        save_distributed_checkpoint: If ``True``, each rank saves its shard of weights and optimizer states to a file.
            The checkpoint is a folder with as many files as the world size.
            If ``False``, the full weights and optimizer states get assembled on rank 0 and saved to a single file.

    ÚautoTNÚparallelize_fnr2   Údata_parallel_sizeÚtensor_parallel_sizeÚsave_distributed_checkpointÚprocess_group_backendÚtimeoutÚreturnc                 óþ   •— t         ‰| �  «        t        s!t        t	        | «      j
                  › d�«      ‚|| _        || _        || _        d| _	        || _
        || _        || _        t        «       | _        d | _        y )Nz  requires PyTorch 2.4 or higher.é   )ÚsuperÚ__init__r*   ÚImportErrorÚtypeÚ__name__Ú_parallelize_fnÚ_data_parallel_sizeÚ_tensor_parallel_sizeÚ
_num_nodesÚ_save_distributed_checkpointÚ_process_group_backendÚ_timeoutÚ_ParallelBackwardSyncControlÚ_backward_sync_controlÚ_device_mesh)Úselfr8   r9   r:   r;   r<   r=   Ú	__class__s          €ú/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/lightning_fabric/strategies/model_parallel.pyrB   zModelParallelStrategy.__init__Y   sz   ø€ ô 	‰ÑÔÝ'Ü¤ d£×!4Ñ!4Ð 5Ð5UÐVÓWÐWØ-ˆÔØ#5ˆÔ Ø%9ˆÔ"ØˆŒØ,GˆÔ)Ø5JˆÔ#Ø-4ˆŒÜ&BÓ&DˆÔ#à26ˆÕó    c                 óH   — | j                   €t        d«      ‚| j                   S )NzKAccessing the device mesh before processes have initialized is not allowed.)rO   ÚRuntimeError©rP   s    rR   Údevice_meshz!ModelParallelStrategy.device_meshp   s&   € à×ÑÐ$ÜÐlÓmÐmØ× Ñ Ð rS   c                 óF   — t        dt        | «      j                  › d�«      ‚)NúThe `z3` does not use the `CheckpointIO` plugin interface.©ÚNotImplementedErrorrD   rE   rV   s    rR   Úcheckpoint_ioz#ModelParallelStrategy.checkpoint_iov   ó$   € ô " E¬$¨t«*×*=Ñ*=Ð)>Ð>qÐ"rÓsÐsrS   Úioc                 óF   — t        dt        | «      j                  › d�«      ‚)NrY   z3` does not support setting a `CheckpointIO` plugin.rZ   )rP   r^   s     rR   r\   z#ModelParallelStrategy.checkpoint_io{   r]   rS   c                 óP   — | j                   €J ‚| j                   | j                     S ©N)Úparallel_devicesÚ
local_rankrV   s    rR   Úroot_devicez!ModelParallelStrategy.root_device€   s+   € ð ×$Ñ$Ð0Ð0Ð0Ø×$Ñ$ T§_¡_Ñ5Ð5rS   c                 ó   — | j                   S ra   ©rI   rV   s    rR   Ú	num_nodeszModelParallelStrategy.num_nodes†   s   € à�‰ÐrS   rg   c                 ó   — || _         y ra   rf   )rP   rg   s     rR   rg   zModelParallelStrategy.num_nodesŠ   s	   € à#ˆ�rS   c                 óH   — | j                   �t        | j                   «      S dS )Nr   )rb   ÚlenrV   s    rR   Únum_processesz#ModelParallelStrategy.num_processesŽ   s$   € à-1×-BÑ-BÐ-NŒs�4×(Ñ(Ó)ÐUÐTUÐUrS   c                 ó~   — | j                   €J ‚| j                   d   }|j                  «       |j                  «       dœS )NÚdata_parallel)Únum_replicasÚrank)rW   ÚsizeÚget_local_rank)rP   Údata_parallel_meshs     rR   Údistributed_sampler_kwargsz0ModelParallelStrategy.distributed_sampler_kwargs’   sE   € ð ×ÑÐ+Ð+Ð+Ø!×-Ñ-¨oÑ>ÐØ 2× 7Ñ 7Ó 9ÐCU×CdÑCdÓCfÑgÐgrS   c                 ó   — | j                   S ra   )rK   rV   s    rR   r<   z+ModelParallelStrategy.process_group_backend™   s   € à×*Ñ*Ð*rS   c                 ó®   — | j                   €J ‚| j                   j                  s1t        | j                   | j                  | j                  «      | _        y y ra   )Úcluster_environmentÚcreates_processes_externallyr   rk   rg   Ú	_launcherrV   s    rR   Ú_configure_launcherz)ModelParallelStrategy._configure_launcher�   sM   € à×'Ñ'Ð3Ð3Ð3Ø×'Ñ'×DÒDÜ6°t×7OÑ7OÐQU×QcÑQcÐei×esÑesÓtˆD�Nð ErS   c                 ó8  •— t         ‰| �  «        | j                  «        | j                  dk(  r| j                  | _        | j
                  dk(  r| j                  | _        t        | j                  | j
                  | j                  | j                  «      | _
        y )Nr7   )rA   Úsetup_environmentÚ_setup_distributedrG   rg   rH   rk   Ú_setup_device_meshÚ
world_sizerd   rO   )rP   rQ   s    €rR   r{   z'ModelParallelStrategy.setup_environment£   s}   ø€ ä‰Ñ!Ô#Ø×ÑÔ!Ø×#Ñ# vÒ-Ø'+§~¡~ˆDÔ$Ø×%Ñ%¨Ò/Ø)-×);Ñ);ˆDÔ&Ü.Ø×$Ñ$ d×&@Ñ&@À$Ç/Á/ÐSW×ScÑScó
ˆÕrS   Úmodulec                 ód  ‡— ddl mŠ t        ˆfd„|j                  «       D «       «      r#t	        d| j
                  j                  › d�«      ‚| j                  || j                  «      }t        |t        «      s!t	        dt        |«      j                  › �«      ‚t        || j                  «       |S )Nr   ©ÚFullyShardedDataParallelc              3   ó6   •K  — | ]  }t        |‰«      –— Œ y ­wra   ©Ú
isinstance)Ú.0Úmodr‚   s     €rR   Ú	<genexpr>z5ModelParallelStrategy.setup_module.<locals>.<genexpr>³   s   øè ø€ ÐUÑDT¸SŒz˜#Ð7×8ÑDTùó   ƒz\Found modules that are wrapped with `torch.distributed.fsdp.FullyShardedDataParallel`. The `z5` only supports the new FSDP2 APIs in PyTorch >= 2.4.zBThe `parallelize_fn` must return a `nn.Module` instance, but got: )Útorch.distributed.fsdpr‚   ÚanyÚmodulesÚ	TypeErrorrQ   rE   rF   rW   r…   r   rD   r+   rd   )rP   r   r‚   s     @rR   Úsetup_modulez"ModelParallelStrategy.setup_module¯   s¦   ø€ åCäÓUÀFÇNÁNÔDTÓUÔUÜðØŸ™×0Ñ0Ð1Ð1fðhóð ð
 ×%Ñ% f¨d×.>Ñ.>Ó?ˆÜ˜&¤&Ô)ÜØTÔUYÐZ`ÓUa×UjÑUjÐTkÐlóð ô 	(¨°×0@Ñ0@ÔAØˆrS   c                  ó   — y ra   © )rP   r   s     rR   Úmodule_to_devicez&ModelParallelStrategy.module_to_deviceÁ   s   € àrS   Ú
empty_initc                 ó¼   — | j                   j                  «       }t        «       }|r$|j                  t	        j
                  d«      «       |j                  |«       |S )NÚmeta)Ú	precisionÚmodule_init_contextr   Úenter_contextÚtorchÚdevice)rP   r’   Úprecision_init_ctxÚstacks       rR   r–   z)ModelParallelStrategy.module_init_contextÅ   sL   € à!Ÿ^™^×?Ñ?ÓAÐÜ“ˆÙð ×Ñ¤§¡¨VÓ 4Ô5Ø×ÑÐ.Ô/ØˆrS   Útensorr(   Ú	reduce_opc                 óB   — t        |t        «      rt        |||¬«      S |S )N)r�   )r…   r   r&   )rP   rœ   r(   r�   s       rR   Ú
all_reducez ModelParallelStrategy.all_reduceÐ   s"   € ô �fœfÔ%Ü)¨&°%À9ÔMÐMØˆrS   ÚargsÚkwargsc                 ó  — t        «       sy t        j                  j                  «       dk(  r6t        j                  j	                  | j
                  j                  g¬«       y t        j                  j	                  «        y )NÚnccl)Ú
device_ids)r#   r˜   ÚdistributedÚget_backendÚbarrierrd   Úindex)rP   r    r¡   s      rR   r§   zModelParallelStrategy.barrierØ   sZ   € ä*Ô,ØÜ×Ñ×(Ñ(Ó*¨fÒ4Ü×Ñ×%Ñ%°$×2BÑ2B×2HÑ2HÐ1IÐ%ÕJä×Ñ×%Ñ%Õ'rS   ÚobjÚsrcc                 óŠ   — t        «       s|S |g}t        j                  j                  ||t        j
                  ¬«       |d   S )Nr'   r   )r#   r˜   r¥   Úbroadcast_object_listÚ_groupÚWORLD)rP   r©   rª   s      rR   Ú	broadcastzModelParallelStrategy.broadcastá   s<   € ä*Ô,ØˆJàˆeˆÜ×Ñ×/Ñ/°°SÄÇÁÐ/ÔMØ�1‰vˆrS   ÚpathÚstateÚstorage_optionsÚfilterc                 óT  — |�8t        dt        | «      j                  › dt        | «      j                  › d�«      ‚|�-| j                  r!t	        t        | «      j                  › d�«      ‚t        | j                  |«      «      }t        ||| j                   | j                  |¬«       y)aN  Save model, optimizer, and other state to a checkpoint on disk.

        If distributed checkpointing is enabled (default), the checkpoint gets saved as a directory containing one file
        per process, with model- and optimizer shards stored per file. Additionally, it creates a metadata file
        `meta.pt` with the rest of the user's state (only saved from rank 0).
        If distributed checkpointing is disabled (``save_distributed_checkpoint=False``), the checkpoint will be
        written to a single file containing the weights, optimizer state and other metadata.

        NÚ`zF.save_checkpoint(..., storage_options=...)` is not supported because `z"` does not use the `CheckpointIO`.zV doesn't support loading distributed filtered checkpoints, so saving them is disabled.)r°   r±   Úfull_state_dictro   r³   )	r�   rD   rE   rJ   r[   r   r¯   Ú_save_checkpointÚglobal_rank)rP   r°   r±   r²   r³   s        rR   Úsave_checkpointz%ModelParallelStrategy.save_checkpointê   s¹   € ð" Ð&ÜØ”D˜“J×'Ñ'Ð(ð )Ü˜$“Z×(Ñ(Ð)Ð)KðMóð ð Ð $×"CÒ"Cä%Ü˜“:×&Ñ&Ð'ð (/ð /óð ô
 �D—N‘N 4Ó(Ó)ˆÜØØØ!%×!BÑ!BÐBØ×!Ñ!Øö	
rS   ÚstrictÚweights_onlyc           
      óˆ  — |s;t        dt        | «      j                  › d|›dt        | «      j                  › d�«      ‚t        | j	                  |«      «      }t        |t        «      rt        ||| j                  |¬«       i S t        |t        «      r"t        dt        | «      j                  › d�«      ‚t        ||||¬«      S )	zOLoad the contents from a checkpoint and restore the state of the given objects.zGot z.load_checkpoint(..., state=zY) but a state with at least  a model instance to reload is required. Pass it in like so: z2.load_checkpoint(..., state={'model': model, ...}))r   r~   rº   zNLoading a single optimizer object from a checkpoint is not supported yet with Ú.)r°   r±   rº   r»   )Ú
ValueErrorrD   rE   r   r¯   r…   r   Ú _load_raw_module_state_from_pathr~   r   r[   Ú_load_checkpoint)rP   r°   r±   rº   r»   s        rR   Úload_checkpointz%ModelParallelStrategy.load_checkpoint  sÉ   € ñ ÜØ”t˜D“z×*Ñ*Ð+Ð+GÈÀyð Qä˜“J×'Ñ'Ð(Ð(\ð^óð ô �D—N‘N 4Ó(Ó)ˆä�eœVÔ$Ü,¨T¸%ÈDÏOÉOÐdjÕkØˆIä�eœYÔ'Ü%Ø`ÔaeÐfjÓak×atÑatÐ`uÐuvÐwóð ô   T°¸vÐT`ÔaÐarS   c                 ó<  — t        «        | j                  «        | j                  «       | _        | j                  €J ‚d| j
                  i}t        r*| j                  j                  dk7  r| j                  nd |d<   t        | j                  | j                  fi |¤Ž y )Nr=   ÚcpuÚ	device_id)
r/   Ú_set_world_ranksÚ_get_process_group_backendrK   rv   rL   r)   rd   rD   r%   )rP   r¡   s     rR   r|   z(ModelParallelStrategy._setup_distributed-  s‰   € ÜŒØ×ÑÔØ&*×&EÑ&EÓ&GˆÔ#Ø×'Ñ'Ð3Ð3Ð3Ø"+¨T¯]©]Ð!;ˆÝ#Ø6:×6FÑ6F×6KÑ6KÈuÒ6T $×"2Ò"2ÐZ^ˆF�;ÑÜ˜d×6Ñ6¸×8SÑ8SÑ^ÐW]Ó^rS   c                 óH   — | j                   xs t        | j                  «      S ra   )rK   r$   rd   rV   s    rR   rÆ   z0ModelParallelStrategy._get_process_group_backend7  s    € Ø×*Ñ*ÒmÔ.[Ð\`×\lÑ\lÓ.mÐmrS   c                 ó>  — | j                   �q| j                   j                  | j                  | j                  z  | j                  z   «       | j                   j                  | j                  | j                  z  «       | j                  xt        _	        t        _	        y ra   )rv   Úset_global_rankÚ	node_rankrk   rc   Úset_world_sizerg   r¸   r   ro   Úutils_rank_zero_onlyrV   s    rR   rÅ   z&ModelParallelStrategy._set_world_ranks:  sy   € Ø×#Ñ#Ð/Ø×$Ñ$×4Ñ4°T·^±^Àd×FXÑFXÑ5XÐ[_×[jÑ[jÑ5jÔkØ×$Ñ$×3Ñ3°D·N±NÀT×EWÑEWÑ4WÔXð ;?×:JÑ:JÐJŒÔÔ2Õ7rS   )r>   r2   ©r>   Nra   )NÚmean)r   )NN)NTN)3rE   Ú
__module__Ú__qualname__Ú__doc__r   r
   r3   r   r   ÚintÚboolr   Ústrr   rB   ÚpropertyrW   r   r   r\   Úsetterr˜   r™   rd   rg   rk   Údictr	   rs   r<   ry   r{   r   rŽ   r‘   r   r–   r   r"   rŸ   r§   r   r¯   r0   r   r¹   rÁ   r|   rÆ   rÅ   Ú__classcell__)rQ   s   @rR   r6   r6   C   s7  ø„ ñð0 ;AØ<BØ,0Ø/3Ø'9ñ7à  &¨,Ð!7¸Ð!?Ñ@ð7ð " '¨&¡/°3Ð"6Ñ7ð7ð $ G¨F¡O°SÐ$8Ñ9ð	7ð
 &*ð7ð  (¨™}ð7ð ˜)Ñ$ð7ð 
õ7ð. ò!ó ð!ð
 Øðt˜|ò tó ó ðtð ×ÑØðt ð t°ò tó ó ðtð Øð6˜UŸ\™\ò 6ó ó ð6ð ð˜3ò ó ðð ×Ñð$ 3ð $¨4ò $ó ð$ð ðV˜sò Vó ðVð Øðh¨D°°c°©Nò hó ó ðhð
 ð+ x°¡}ò +ó ð+ð òuó ðuð
 ô	
ó ð	
ð ð 6ð ¨fò ó ðð" ð vð °$ò ó ðð ñ¨h°t©nð ÐH^ò ó ðð àgmñØðØ%-¨c¡]ðØFNÈuÐU]Ð_bÐUbÑOcÑFdðà	òó ðð ð(˜Sð (¨Cð (°Dò (ó ð(ð ñ˜Zð ¨cð ¸*ò ó ðð ð
 *.ØBFñ#
àð#
ð �C˜˜v y°#Ð5Ñ6Ð6Ñ7ð#
ð " #™ð	#
ð
 ˜˜c 8¨S°#¨J¸Ð,<Ñ#=Ð=Ñ>Ñ?ð#
ð 
ò#
ó ð#
ðJ ð _cØØ'+ñbàðbð ˜˜f i°°c¸5ÀÈÐTWÐAWÑ;XÐ6XÑ1YÐYÑZÑ[ðbð ð	bð
 ˜t‘nðbð 
ˆc�3ˆh‰òbó ðbó8_ðn¨Có n÷KrS   r6   c                   ó*   — e Zd Zedededefd„«       Zy)rM   r   Úenabledr>   c                 ó   — t        ||¬«      S )z9Blocks gradient synchronization inside the FSDP2 modules.)r   rÚ   )Ú_FSDPNoSync©rP   r   rÚ   s      rR   Úno_backward_syncz-_ParallelBackwardSyncControl.no_backward_syncD  s   € ô  &°'Ô:Ð:rS   N)rE   rÏ   rÐ   r   r   rÓ   r   rÞ   r�   rS   rR   rM   rM   C  s*   „ Øð; vð ;¸ð ;ÐAWò ;ó ñ;rS   rM   c                   óP   — e Zd Zdededdfd„Zdeddfd„Zdd„Zd	ed
ededdfd„Z	y)rÜ   r   rÚ   r>   Nc                 ó    — || _         || _        y ra   )Ú_moduleÚ_enabledrÝ   s      rR   rB   z_FSDPNoSync.__init__K  s   € ØˆŒØˆ�rS   Úrequires_grad_syncc                 óŽ   — ddl m} | j                  j                  «       D ]"  }t	        ||«      sŒ|j                  |d¬«       Œ$ y )Nr   )Ú
FSDPModuleF©Úrecurse)Ú"torch.distributed._composable.fsdprå   rá   rŒ   r…   Úset_requires_gradient_sync)rP   rã   rå   r‡   s       rR   Ú_set_requires_grad_syncz#_FSDPNoSync._set_requires_grad_syncO  s:   € ÝAà—<‘<×'Ñ'Ö)ˆCÜ˜#˜zÕ*Ø×.Ñ.Ð/AÈ5Ð.ÕQñ *rS   c                 ó<   — | j                  | j                   «       y ra   ©rê   râ   rV   s    rR   Ú	__enter__z_FSDPNoSync.__enter__V  s   € Ø×$Ñ$¨¯©Ð%6Õ7rS   Úexc_typeÚ	exc_valueÚ	tracebackc                 ó:   — | j                  | j                  «       y ra   rì   )rP   rî   rï   rð   s       rR   Ú__exit__z_FSDPNoSync.__exit__Y  s   € Ø×$Ñ$ T§]¡]Õ3rS   rÍ   )
rE   rÏ   rÐ   r   rÓ   rB   rê   rí   r	   rò   r�   rS   rR   rÜ   rÜ   J  sX   „ ð ˜vð  °ð  ¸ó  ðR¸$ð RÀ4ó Ró8ð4 ð 4°ð 4Àð 4Èô 4rS   rÜ   r°   r±   r¶   ro   r³   r>   c                 óÊ  — | j                  «       r|rt        | «      st        d| › �«      ‚|j                  «       D �cg c]  }t	        |«      sŒ|‘Œ }}t        |«      dk(  rt        d«      ‚t        |«      dkD  rt        d«      ‚|d   }ddlm}m	}m
}	  ||d¬«      }
i }i }|j                  «       D ]v  \  }}t        |t        «      r |||
¬	«      }|}nBt        |t        «      r |	|||
¬	«      }|}n$t        |t        «      r|j!                  «       n|}|}t#        ||xs i ||«       Œx |rNt        | «      rt%        j&                  | «       |j)                  |«       |dk(  rt+        j,                  || «       y y | j/                  «       r| j1                  «        | j3                  dd¬
«       t5        || «       |dk(  rt+        j,                  || t6        z  «       y y c c}w )Nz/The checkpoint path exists and is a directory: r   a  Could not find a distributed model in the provided checkpoint state. Please provide the model as part of the state like so: `save_checkpoint(..., state={'model': model, ...})`. Make sure you set up the model (and optimizers if any) through the strategy before saving the checkpoint.r@   zêFound multiple distributed models in the given state. Saving distributed checkpoints is currently limited to a single model per checkpoint. To save multiple models, call the save method for each model separately with a different path.)ÚStateDictOptionsÚget_model_state_dictÚget_optimizer_state_dictT)r¶   Úcpu_offload©Úoptions)ÚparentsÚexist_ok)Úis_dirr   ÚIsADirectoryErrorÚvaluesÚ_has_dtensor_modulesrj   r¾   Ú'torch.distributed.checkpoint.state_dictrô   rõ   rö   Úitemsr…   r   r   r1   Ú
state_dictr   ÚshutilÚrmtreeÚupdater˜   ÚsaveÚis_fileÚunlinkÚmkdirr   r,   )r°   r±   r¶   ro   r³   r   rŒ   rô   rõ   rö   Ústate_dict_optionsÚconverted_stateÚmetadataÚkeyr©   Ú	convertedÚtarget_dicts                    rR   r·   r·   ]  sË  € ð ‡{�{„}™Ô1GÈÔ1MÜÐ"QÐRVÐQWÐ XÓYÐYà$)§L¡L¤NÓS¡N˜&Ô6JÈ6Õ6RŠv N€GÐSÜ
ˆ7ƒ|�qÒÜðoó
ð 	
ô
 ˆ7ƒ|�aÒÜðLó
ð 	
ð
 �Q‰Z€FçxÑxá)¸/ÐW[Ô\Ðð ')€OØ!€HØ—K‘K–M‰ˆˆSä�cœ6Ô"Ù,¨SÐ:LÔMˆIØ)‰KÜ˜œYÔ'Ù0°¸ÐFXÔYˆIØ)‰Kä,6°s¼IÔ,F˜Ÿ™Ô(ÈCˆIØ"ˆKÜ�c˜6š< R¨°KÕ@ð "ñ Ü! $Ô'Ü�M‰M˜$ÔØ×Ñ˜xÔ(Ø�1Š9Ü�J‰J�¨Õ-ð ð �<‰<Œ>Ø�K‰KŒMØ�
‰
˜4¨$ˆ
Ô/Ü$ _°dÔ;Ø�1Š9Ü�J‰J�x Ô(:Ñ!:Õ;ð ùò_ Ts   ¾G ÁG rº   Úoptimizer_states_from_listr»   c                 óè  — ddl m}m}m}m} |j                  «       D �	�
ci c]  \  }	}
t        |
«      sŒ|	|
“Œ }}	}
t        |«      dk(  rt        d«      ‚|j                  «       D �	�ci c]  \  }	}t        |t        «      sŒ|	|“Œ }}	}t        |«      dkD  rt        d«      ‚t        |j                  «       «      d   \  }}
t        | «      �r |d¬«      }| ||
«      i}t        || «       |
j                  ||   |¬«       |j                  «       D ]+  \  }}| ||
|«      i}t        || «        ||
|||   |¬	«       Œ- t        j                   | t"        z  |¬
«      }|j%                  «       |j%                  «       z
  |j%                  «       z
  }t'        ||j%                  «       |¬«       |D ]  }	|	|vrŒ|j)                  |	«      ||	<   Œ |S t+        | «      r÷t        j                   | dd|¬«      }t-        |j)                  |«      |
|¬«        |dd|¬«      }t/        |j                  «       «      D ]<  \  }\  }}|r	|d   |   }n|j)                  |«      }t1        ||
«      } ||
|||¬	«       Œ> |j%                  «       |j%                  «       z
  |j%                  «       z
  }t'        ||j%                  «       |¬«       t3        |||¬«       |S t        dt5        | «      ›d�«      ‚c c}
}	w c c}}	w )Nr   )rô   rõ   rö   Úset_optimizer_state_dicta  Could not find a distributed model in the provided checkpoint state. Please provide the model as part of the state like so: `load_checkpoint(..., state={'model': model, ...})`. Make sure you set up the model (and optimizers if any) through the strategy before loading the checkpoint.r@   zëFound multiple distributed models in the given state. Loading distributed checkpoints is currently limited to a single model per checkpoint. To load multiple models, call the load method for each model separately with a different path.T)r÷   ©rº   )Úoptim_state_dictrù   )r»   rÃ   )ÚmmapÚmap_locationr»   ©Úbroadcast_from_rank0r¶   rº   Úoptimizer_states)ÚsourceÚdestinationÚkeysz	The path z£ does not point to a valid checkpoint. Make sure the path points to either a directory with distributed checkpoint shards, or a single file with a full checkpoint.)r   rô   rõ   rö   r  r  rÿ   rj   r¾   r…   r   Úlistr   r   Úload_state_dictr˜   Úloadr,   r  r!   Úpopr   Ú_load_raw_module_stateÚ	enumerateÚ _rekey_optimizer_state_if_neededr.   rÔ   )r°   r±   rº   r  r»   rô   rõ   rö   r  r  r   rŒ   ÚoptimÚ
optimizersÚ
module_keyr
  Úmodule_stateÚ	optim_keyÚoptim_stater  Úrequested_metadata_keysÚ
checkpointÚoptimizer_idxÚoptimizer_nameÚ	optimizerÚoptimizer_states                             rR   rÀ   rÀ   š  s   € ÷ó ð /4¯k©k¬mÔ\©m™{˜s FÔ?SÐTZÕ?[ˆs�F‰{¨m€GÑ\Ü
ˆ7ƒ|�qÒÜðpó
ð 	
ð
 05¯{©{¬}Ô]©}¡  eÄ
È5ÔR[Õ@\�#�u‘*¨}€JÑ]Ü
ˆ7ƒ|�aÒÜðLó
ð 	
ô
 ˜gŸm™m›oÓ.¨qÑ1Ñ€J�ä˜dÕ#Ù-¸$Ô?Ðà"Ñ$8¸Ó$@ÐAˆÜ$ \°4Ô8Ø×Ñ˜|¨JÑ7ÀÐÔGð !+× 0Ñ 0Ö 2ÑˆI�uØ$Ñ&>¸vÀuÓ&MÐNˆKÜ(¨°dÔ;Ù$ V¨UÀ[ÐQZÑE[ÐewÖxð !3ô —:‘:˜dÔ%7Ñ7ÀlÔSˆØ"'§*¡*£,°·±³Ñ"?À*Ç/Á/ÓBSÑ"SÐÜ)Ð*AÀ8Ç=Á=Ã?Ð[aÕbÛ*ˆCØ˜(Ñ"ØØ!Ÿ™ cÓ*ˆE�#ŠJð +ð ˆä˜4Ô Ü—Z‘Z ¨4¸eÐR^Ô_ˆ
Ü˜zŸ~™~¨jÓ9¸6È&ÕQá-Ø!%Ø Øô
Ðô
 ;DÀJ×DTÑDTÓDVÖ:WÑ6ˆMÑ6˜N¨IÙ)ð #-Ð-?Ñ"@ÀÑ"O‘à",§.¡.°Ó"@�ä>¸ÐPVÓWˆOÙ$ØØØ!0Ø*ö	ð ;Xð  #(§*¡*£,°·±³Ñ"?À*Ç/Á/ÓBSÑ"SÐÜ)Ð*AÀ:Ç?Á?ÓCTÐ]cÕdô 	 
¸ÐD[Õ\ð Ðä
Ø
”C˜“I�=ð !bð 	bóð ùóW ]ùó ^s    K(´K(Á(K.ÂK.r9   r:   r~   r™   r2   c           	      óv   — ddl m} | |z  |k7  rt        d| › d|› d|› d�«      ‚ ||j                  | |fd¬«      S )	Nr   )Úinit_device_meshzThe sizes `data_parallel_size=z` and `tensor_parallel_size=z*` multiplied should equal the world size (z).)rm   Útensor_parallel)Údevice_typeÚ
mesh_shapeÚmesh_dim_names)Útorch.distributed.device_meshr1  rU   rD   )r9   r:   r~   r™   r1  s        rR   r}   r}   ù  sj   € õ ?àÐ0Ñ0°JÒ>ÜØ,Ð-?Ð,@ð A&Ø&:Ð%;ð <Ø�˜Bð ó
ð 	
ñ
 Ø—K‘KØ&Ð(<Ð=Ø;ôð rS   r   c                 óx   ‡— ddl mŠ t        | t        «      xr" t	        ˆfd„| j                  «       D «       «      S )Nr   )ÚDTensorc              3   ó6   •K  — | ]  }t        |‰«      –— Œ y ­wra   r„   )r†   Útr8  s     €rR   rˆ   z'_has_dtensor_modules.<locals>.<genexpr>  s   øè ø€ Ð-bÑNaÈ¬j¸¸G×.DÑNaùr‰   )Útorch.distributed._tensorr8  r…   r   r‹   Ú
parameters)r   r8  s    @rR   rÿ   rÿ     s,   ø€ Ý1ä�fœfÓ%Òb¬#Ó-bÈf×N_ÑN_ÔNaÓ-bÓ*bÐbrS   c                 ó¦   — t        | «      st        d| › �«      ‚t        rt        j                  | dd¬«      n
t        | «      }t        ||||¬«       y)z;Loads the state dict from a file path into the FSDP module.zxFailed to load checkpoint directly into the model. The given path must be a single file containing the full state dict: TrÃ   )r  r  )r  r   r~   rº   N)r   r¾   r)   r˜   r  r-   r!  )r°   r   r~   rº   r  s        rR   r¿   r¿     sV   € ä˜tÔ$Üð!Ø!% ð(ó
ð 	
õ
 E]”—‘˜D t¸%Õ@ÔblÐmqÓbr€JÜ j¸ÈJÐ_eÖfrS   r  c                 ó¸  — ddl m} t        |«      rsddlm}m}  |ddd¬«      }|j                  «       D ]L  \  }}	t        |	«      D ]9  \  }
}|› |rdnd› |
› �}|| vr|sŒt        d	|› d
�«      ‚|
| |   i} ||	||¬«       Œ; ŒN yt        ||«      r+t        ||d¬«      5  |j                  | |¬«       ddd«       y|j                  | |¬«       y# 1 sw Y   yxY w)zlLoads the state dict into the module by gathering all weights first and then and writing back to each shard.r   r�   )rô   Úset_model_state_dictTFr  r½   Ú zThe model contains a key 'z^' that does not exist in the loaded checkpoint. To disable strict loading, set `strict=False`.rø   )r~   Ú
rank0_onlyr  N)rŠ   r‚   rÿ   r   rô   r?  Únamed_modulesÚ%_named_parameters_and_buffers_to_loadÚKeyErrorr…   r   r  )r  r   r~   rº   ÚFSDPrô   r?  r
  Úsubmodule_nameÚ	submoduleÚ
param_nameÚ_Úfull_param_nameÚlocal_state_dicts                 rR   r!  r!     s  € õ Hä˜FÔ#ßbá-Ø!%Ø àô	
Ðð *0×)=Ñ)=Ö)?Ñ%ˆN˜IÜ!FÀyÖ!Q‘�
˜AØ%3Ð$4¹N±SÐPRÐ4SÐT^ÐS_Ð"`�Ø"¨*Ñ4Ù!Ø Ü"Ø4°_Ð4Eð FJð Jóð ð %/°
¸?Ñ0KÐ#LÐ Ù$ YÐ0@ÐJ\Ö]ñ "Rñ *@ô 
�F˜DÔ	!Ü)¨&¸ZÐTYÖZØ×"Ñ" :°fÐ"Ô=÷ [ÐZð 	×Ñ˜z°&ÐÕ9÷ [ÐZús   ÂCÃCc              #   ó²   K  — t        j                  | j                  d¬«      | j                  d¬«      «      D ]  \  }}|| j                  v rŒ||f–— Œ y­w)zEReturns parameters and buffers, with non-persistent buffers excluded.Fræ   N)Ú	itertoolsÚchainÚnamed_buffersÚnamed_parametersÚ_non_persistent_buffers_set)r   rH  Úparams      rR   rC  rC  D  s^   è ø€ ä&Ÿ_™_Ø×Ñ UÐÓ+Ø×Ñ¨ÐÓ.öÑˆ
�Eð ˜×;Ñ;Ñ;ØØ˜%ÐÓñùs   ‚AAÚoptimizer_state_dictc                 ó²   — ddl m} ddl m} t        t	        | d   j                  «       «      d   t        «      r|j                  | |j                  |«      } | S )zyHandles the case where the optimizer state is saved from a normal optimizer and converts the keys to parameter
    names.r   r�   )ÚOptimStateKeyTyper±   )	rŠ   r‚   rU  r…   r  r  rÒ   Úrekey_optim_state_dictÚ
PARAM_NAME)rS  r   rE  rU  s       rR   r#  r#  O  sR   € õ HÝ8ä”$Ð+¨GÑ4×9Ñ9Ó;Ó<¸QÑ?ÄÔEØ#×:Ñ:Ð;OÐQb×QmÑQmÐouÓvÐØÐrS   ra   )TFN)T)r@   T)]rM  r  Úcollections.abcr   Ú
contextlibr   r   Údatetimer   Úpathlibr   Útypingr   r	   r
   r   r   r   r   r˜   Ú"lightning_utilities.core.rank_zeror   rÌ   r   Útorch.nnr   Útorch.optimr   Útyping_extensionsr   r   Úlightning_fabric.pluginsr   Ú5lightning_fabric.plugins.collectives.torch_collectiver   Ú lightning_fabric.strategies.fsdpr   r   r   r   r   Ú7lightning_fabric.strategies.launchers.subprocess_scriptr   Ú$lightning_fabric.strategies.parallelr   Ú$lightning_fabric.strategies.strategyr   r   r    r!   Ú&lightning_fabric.utilities.distributedr"   r#   r$   r%   r&   r(   r­   Ú"lightning_fabric.utilities.importsr)   r*   Úlightning_fabric.utilities.initr+   Úlightning_fabric.utilities.loadr,   r-   r.   Ú$lightning_fabric.utilities.rank_zeroÚlightning_fabric.utilities.seedr/   Ú lightning_fabric.utilities.typesr0   r1   r6  r2   r3   r6   rM   rÜ   r×   rÔ   rÓ   rÒ   r·   rÀ   r™   r}   Úobjectrÿ   r¿   r!  rC  r#  r�   rS   rR   Ú<module>ro     s¥  ðó Û Ý %ß 8Ý Ý ß R× RÑ Rã Ý UÝ Ý Ý !ß 1å 1Ý T÷õ õ ^Ý A÷ó ÷õ õ Cß aÝ Kß \Ñ \Ý ?Ý 6ß =áÝ8á	� Ô	(€ô}KÐ,ô }Kô@;Ð#7ô ;ô4Ð(ô 4ð0 ?Cñ:<Ø
ð:<à��U˜6 9¨cÐ1Ñ2Ð2Ñ3ð:<ð ð:<ð ð	:<ð
 �T˜#˜x¨¨c¨
°DÐ(8Ñ9Ð9Ñ:Ñ;ð:<ð 
ó:<ð@ Ø',Ø#'ñ\Ø
ð\à��U˜6 9¨cÐ1Ñ2Ð2Ñ3ð\ð ð\ð !%ð	\ð
 ˜4‘.ð\ð 
ˆ#ˆsˆ(�^ó\ð~Øðàðð ðð �L‰Lð	ð
 óð*c ð c¨I°fÑ,=ó cñ	g¨4ð 	g¸ð 	gÈSð 	gÐZ^ð 	gÐjnó 	gð UYñ!:Ø�S˜#�X‘ð!:Ø(.ð!:Ø<?ð!:ØMQð!:à	ó!:ðH °&ð  ¸Yó  ð ¸4ÀÀSÀ¹>ð  ÐSYð  Ð^bÐcfÐhkÐckÑ^lô  rS   