Ë
      çitx  ã                   óÂ  — d dl Z d dl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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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.m/Z/m0Z0m1Z1m2Z2m3Z3m4Z4m5Z5 d dl6m7Z7 d dl8m9Z9m:Z:m;Z;m<Z< d dl8m=Z> d dl?m@Z@mAZA d dlBmCZC d dlDmEZEmFZF d dlGmHZH d dlImJZJ d dlKmLZLmMZM d dlNmOZO d dlPmQZQ d dlRmSZS d dlTmUZU d dlVmWZW d dlXmYZY d dlZm[Z[ d d l\m]Z] d d!l^m_Z_mZm`Z` er6d d"lambZb d d#lcmdZdmeZemfZf d d$lgmhZh eeieje      eeekelgekf   ehf   Zmeefed%   f   Zn e jÞ                  ep«      Zq G d&„ d'eW«      Zry)(é    N)Ú	GeneratorÚMapping)ÚcontextmanagerÚnullcontext)Ú	timedelta)ÚPath)ÚTYPE_CHECKINGÚAnyÚCallableÚLiteralÚOptionalÚUnion)Úrank_zero_only)ÚTensor)ÚModule)Ú	Optimizer)Úoverride)ÚCheckpointIOÚClusterEnvironment)Údefault_pg_timeout)Ú_StrategyRegistry)Ú_METADATA_FILENAMEÚ _activation_checkpointing_kwargsÚ_auto_wrap_policy_kwargsÚ_distributed_checkpoint_loadÚ_distributed_checkpoint_saveÚ_get_full_state_dict_contextÚ_get_sharded_state_dict_contextÚ_init_cpu_offloadÚ_init_sharding_strategyÚ_is_full_checkpointÚ_is_sharded_checkpointÚ_move_torchmetrics_to_deviceÚ_optimizer_has_flat_paramsÚ_setup_activation_checkpointing)Ú_load_raw_module_state)Ú_distributed_is_initializedÚ-_get_default_process_group_backend_for_deviceÚ_init_dist_connectionÚ_sync_ddp_if_available©Úgroup)Ú_TORCH_GREATER_EQUAL_2_2Ú_TORCH_GREATER_EQUAL_2_3)Ú&_has_meta_device_parameters_or_buffers)Ú
_lazy_loadÚ_materialize_tensors)Ú_optimizers_to_device)Ú
reset_seed)Ú_PATHÚReduceOp)ÚLightningOptimizer)Ú	Precision)ÚFSDPPrecision)Ú_SubprocessScriptLauncher)ÚParallelStrategy)Ú
TBroadcast)Ú	TrainerFn)Úis_overridden)Úrank_zero_infor   Úrank_zero_warn)Ú
DeviceMesh)Ú
CPUOffloadÚMixedPrecisionÚShardingStrategy)ÚModuleWrapPolicy)Ú
FULL_SHARDÚSHARD_GRAD_OPÚNO_SHARDÚHYBRID_SHARDc            #       óT  ‡ — e Zd ZU dZdZg Zee   ed<   dddddde	ddddddddfde
d   d	e
eej                        d
e
e   de
e   de
e   de
e   de
e   deeddf   de
d   de
d   de
eee   eee      f      de
d   ddded   de
eee   df      deddf"ˆ fd„Zeedej                  fd „«       «       Zedefd!„«       Zede
e   fd"„«       Zede
d   fd#„«       Zeede fd$„«       «       Z!e!jD                  ede
e   ddfd%„«       «       Z!eede#fd&„«       «       Z$eedefd'„«       «       Z%eedefd(„«       «       Z&edQˆ fd)„«       Z'defd*„Z(dQd+„Z)edQd,„«       Z*ed-edefd.„«       Z+edRd/„«       Z,edRˆ fd0„«       Z-edQd1„«       Z.e/edSd2e
e   de0d3   fd4„«       «       Z1e/ede0d3   fd5„«       «       Z2edSd6e
e   ddfd7„«       Z3edTd8e4d9ede4fd:„«       Z5e	 	 dUd;ee6ef   d<e
e   d=e
ee7ef      de6fd>„«       Z8dee   fd?„Z9edQd@„«       Z:e;dee   fdA„«       Z<e;edBe=ddfdC„«       «       Z>ede#eef   fdD„«       Z?edVdEe@eef   dFeddfdG„«       ZAedHeBde#ee6f   fdI„«       ZCedEe@eef   ddfdJ„«       ZDe	 dSdEe#eef   dKeEdLe
e   ddfˆ fdM„«       ZFedSdNeEdOe
e   de#eef   fdP„«       ZGˆ xZHS )WÚFSDPStrategyae  Strategy for Fully Sharded Data Parallel provided by torch.distributed.

    Fully Sharded Training shards the entire model across all available GPUs, allowing you to scale model
    size, whilst using efficient communication to reduce overhead. In practice, this means we can remain
    at parity with PyTorch DDP, whilst scaling our model sizes dramatically. The technique is similar
    to ZeRO-Stage 3.

    For more information check out
    `this blogpost <https://pytorch.org/blog/introducing-pytorch-fully-sharded-data-parallel-api>`__.

    Defaults have been set and options have been exposed, but may require configuration
    based on your level of memory/speed efficiency. We suggest having a look at
    `this tutorial <https://pytorch.org/tutorials/intermediate/FSDP_tutorial.html>`__ for more information.

    Arguments:
        cpu_offload: See ``cpu_offload`` parameter in :class:`torch.distributed.fsdp.FullyShardedDataParallel`.
        mixed_precision: See ``mixed_precision`` parameter in :class:`torch.distributed.fsdp.FullyShardedDataParallel`.
        auto_wrap_policy: Same as ``auto_wrap_policy`` parameter in
            :class:`torch.distributed.fsdp.FullyShardedDataParallel`. For convenience, this also accepts a set of the
            layer classes to wrap.
        activation_checkpointing: Deprecated. Use ``activation_checkpointing_policy``.
        activation_checkpointing_policy: Same as ``auto_wrap_policy`` parameter in
            :class:`torch.distributed.fsdp.FullyShardedDataParallel` but used when selecting the modules for which you
            want to enable activation checkpointing. Enabling this can free up a significant amount of memory at the
            cost of speed since activations in these layers need to be recomputed during backpropagation. For
            convenience, this also accepts a set of the layer classes to wrap.
        sharding_strategy: Select whether to shard model parameters, gradients, optimizer states, or a combination of
            them. Available values are:

            - ``"FULL_SHARD"``: Shards model parameters, gradients, and optimizer states (default).
            - ``"SHARD_GRAD_OP"``: Shards gradients and optimizer states only. Model parameters get replicated.
            - ``"NO_SHARD"``: No sharding (identical to regular DDP).
            - ``"HYBRID_SHARD"``: Shards model parameters, gradients, and optimizer states within a single machine, but
              replicates across machines. See also the `device_mesh` parameter below.

            Also accepts a :class:`torch.distributed.fsdp.ShardingStrategy` enum value.

        device_mesh: A tuple `(replication size, sharding size)` that defines over how many devices to shard and
            replicate the model. The product of the two numbers must equal the world size. Only valid in combination
            with the `HYBRID_SHARD` sharding strategy.

        state_dict_type: The format in which the state of the model and optimizers gets saved into the checkpoint.

            - ``"full"``: The full weights and optimizer states get assembled on rank 0 and saved to a single file.
            - ``"sharded"``: 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.

        \**kwargs: See available parameters in :class:`torch.distributed.fsdp.FullyShardedDataParallel`.

    ÚfsdpÚ_registered_strategiesNrE   ÚfullÚacceleratorzpl.accelerators.AcceleratorÚparallel_devicesÚcluster_environmentÚcheckpoint_ioÚprecision_pluginÚprocess_group_backendÚtimeoutÚcpu_offloadrA   Úmixed_precisionrB   Úauto_wrap_policyÚ_POLICYÚactivation_checkpointingÚactivation_checkpointing_policyÚsharding_strategyÚ_SHARDING_STRATEGYÚstate_dict_type)rM   ÚshardedÚdevice_meshr@   ÚkwargsÚreturnc                 óŠ  •— t         ‰| �  |||||¬«       d| _        || _        || _        t        |«      | _        |	| _        t        |
|«      | _	        |� t        st        d«      ‚|| j                  d<   t        || j                  «      | _        | j                  j                  dd«       t        ||«      | _        || _        y )N)rN   rO   rP   rQ   rR   é   z=The `device_mesh` argument is only supported in torch >= 2.2.r_   Úuse_orig_paramsT)ÚsuperÚ__init__Ú	num_nodesÚ_process_group_backendÚ_timeoutr   rU   rV   r   r`   r-   Ú
ValueErrorr    r[   Ú
setdefaultr   Ú_state_dict_type)ÚselfrN   rO   rP   rQ   rR   rS   rT   rU   rV   rW   rY   rZ   r[   r]   r_   r`   Ú	__class__s                    €úv/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/pytorch_lightning/strategies/fsdp.pyrf   zFSDPStrategy.__init__“   sÍ   ø€ ô& 	‰ÑØ#Ø-Ø 3Ø'Ø-ð 	ô 	
ð ˆŒØ&;ˆÔ#Ø-4ˆŒÜ,¨[Ó9ˆÔØ.ˆÔÜ.Ð/?ÀÓHˆŒàÐ"Ý+Ü Ð!`ÓaÐaØ)4ˆD�K‰K˜Ñ&ä!8Ð9JÈDÏKÉKÓ!XˆÔð 	�‰×ÑÐ0°$Ô7ä0PØ$Ð&Eó1
ˆÔ-ð !0ˆÕó    c                 óP   — | j                   €J ‚| j                   | j                     S ©N)rO   Ú
local_rank©rm   s    ro   Úroot_devicezFSDPStrategy.root_deviceÄ   s+   € ð ×$Ñ$Ð0Ð0Ð0Ø×$Ñ$ T§_¡_Ñ5Ð5rp   c                 óH   — | j                   �t        | j                   «      S dS )Nr   )rO   Úlenrt   s    ro   Únum_processeszFSDPStrategy.num_processesÊ   s$   € à-1×-BÑ-BÐ-NŒs�4×(Ñ(Ó)ÐUÐTUÐUrp   c                 ó   — | j                   S rr   )rh   rt   s    ro   rS   z"FSDPStrategy.process_group_backendÎ   s   € à×*Ñ*Ð*rp   c                 ó„   — | j                   r| j                   S | j                  }t        |t        «      r|j                  S y rr   )rV   rR   Ú
isinstancer8   Úmixed_precision_config©rm   Úplugins     ro   r|   z#FSDPStrategy.mixed_precision_configÒ   s;   € à×ÒØ×'Ñ'Ð'Ø×&Ñ&ˆÜ�fœmÔ,Ø×0Ñ0Ð0Ørp   c                 ó\   — | j                   }|�t        |t        «      sJ ‚|S t        d«      S )Nz32-true)Ú_precision_pluginr{   r8   r}   s     ro   rR   zFSDPStrategy.precision_pluginÛ   s5   € ð ×'Ñ'ˆØÐÜ˜f¤mÔ4Ð4Ð4ØˆMÜ˜YÓ'Ð'rp   c                 óR   — |�t        |t        «      st        d|› �«      ‚|| _        y )NzGThe FSDP strategy can only work with the `FSDPPrecision` plugin, found )r{   r8   Ú	TypeErrorr€   )rm   rR   s     ro   rR   zFSDPStrategy.precision_pluginä   s7   € ð Ð'´
Ð;KÌ]Ô0[ÜØYÐZjÐYkÐlóð ð "2ˆÕrp   c                 óN   — | j                   | j                  z  | j                  dœS )N)Únum_replicasÚrank)rg   rx   Úglobal_rankrt   s    ro   Údistributed_sampler_kwargsz'FSDPStrategy.distributed_sampler_kwargsí   s$   € ð "&§¡°$×2DÑ2DÑ!DÈt×O_ÑO_Ñ`Ð`rp   c                  ó   — y)NT© rt   s    ro   Úrestore_checkpoint_after_setupz+FSDPStrategy.restore_checkpoint_after_setupò   s   € ð rp   c                  ó   — y)NFr‰   rt   s    ro   Úlightning_restore_optimizerz(FSDPStrategy.lightning_restore_optimizer÷   s   € ð rp   c                 óX  •— t         ‰| �  «        t        j                  | j                  j
                  › d�«       t        «        | j                  «        | j                  «       | _	        | j                  €J ‚d| j                  i}t        r*| j                  j                  dk7  r| j                  nd |d<   t        | j                  | j                  fi |¤Ž t!        | j"                  j%                  d«      t&        «      r*ddlm}  |d| j"                  d   «      | j"                  d<   y y )	Nz: setting up distributed...rT   ÚcpuÚ	device_idr_   r   )Úinit_device_meshÚcuda)re   Úsetup_environmentÚlogÚdebugrn   Ú__name__r3   Úset_world_ranksÚ_get_process_group_backendrh   rP   ri   r.   ru   Útyper)   r{   r`   ÚgetÚtupleÚtorch.distributed.device_meshr�   )rm   r`   r�   rn   s      €ro   r’   zFSDPStrategy.setup_environmentü   sõ   ø€ ä‰Ñ!Ô#Ü�	‰	�T—^‘^×,Ñ,Ð-Ð-HÐIÔJÜŒð 	×ÑÔà&*×&EÑ&EÓ&GˆÔ#Ø×'Ñ'Ð3Ð3Ð3Ø"+¨T¯]©]Ð!;ˆÝ#Ø6:×6FÑ6F×6KÑ6KÈuÒ6T $×"2Ò"2ÐZ^ˆF�;ÑÜ˜d×6Ñ6¸×8SÑ8SÑ^ÐW]Ò^ô �d—k‘k—o‘o mÓ4´eÔ<ÝFá)9¸&À$Ç+Á+ÈmÑB\Ó)]ˆD�K‰K˜Ò&ð =rp   c                 óH   — | j                   xs t        | j                  «      S rr   )rh   r(   ru   rt   s    ro   r—   z'FSDPStrategy._get_process_group_backend  s    € Ø×*Ñ*ÒmÔ.[Ð\`×\lÑ\lÓ.mÐmrp   c                 ó>  — | j                   �q| j                   j                  | j                  | j                  z  | j                  z   «       | j                   j                  | j                  | j                  z  «       | j                  xt        _	        t        _	        y rr   )rP   Úset_global_rankÚ	node_rankrx   rs   Úset_world_sizerg   r†   r   r…   Úutils_rank_zero_onlyrt   s    ro   r–   zFSDPStrategy.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Õ7rp   c                 ó®   — | j                   €J ‚| j                   j                  s1t        | j                   | j                  | j                  «      | _        y y rr   )rP   Úcreates_processes_externallyr9   rx   rg   Ú	_launcherrt   s    ro   Ú_configure_launcherz FSDPStrategy._configure_launcher  sM   € à×'Ñ'Ð3Ð3Ð3Ø×'Ñ'×DÒDÜ6°t×7OÑ7OÐQU×QcÑQcÐei×esÑesÓtˆD�Nð Erp   Úmodelc           	      ó0  ‡— ddl mŠ t        ˆfd„|j                  «       D «       «      r=t	        |«      rt        d«       d| j                  v rœt        d«       | j                  d= nƒt        j                  d| j                  j                  › d| j                  › �«        ‰d
|| j                  | j                  | j                  | j                  j                  d	œ| j                  ¤Ž}t        || j                  «       t        || j                   «       |S )z|Wraps the model into a :class:`~torch.distributed.fsdp.fully_sharded_data_parallel.FullyShardedDataParallel`
        module.r   ©ÚFullyShardedDataParallelc              3   ó6   •K  — | ]  }t        |‰«      –— Œ y ­wrr   )r{   )Ú.0Úmodr©   s     €ro   Ú	<genexpr>z,FSDPStrategy._setup_model.<locals>.<genexpr>)  s   øè ø€ ÐTÁO¸SŒz˜#Ð7×8ÁOùs   ƒzYThe model is already wrapped in `FSDP` but there are still parameters on the meta device.rW   z_A FSDP `auto_wrap_policy` is set, but the model is already wrapped. The policy will be ignored.z&setting up FSDP model with device id: z
, kwargs: )ÚmodulerU   rV   r[   r�   r‰   )Útorch.distributed.fsdpr©   ÚanyÚmodulesr/   r?   r`   r“   r”   ru   ÚindexrU   r|   r[   r#   r%   r   )rm   r¦   r©   s     @ro   Ú_setup_modelzFSDPStrategy._setup_model#  sû   ø€ õ 	DäÓTÀEÇMÁMÄOÓTÔTÜ5°eÔ<ÜØoôð " T§[¡[Ñ0äØuôð —K‘KÐ 2Ñ3ä�I‰IÐ>¸t×?OÑ?O×?UÑ?UÐ>VÐV`Ðae×alÑalÐ`mÐnÔoÙ,ð ØØ ×,Ñ,Ø $× ;Ñ ;Ø"&×"8Ñ"8Ø×*Ñ*×0Ñ0ñð —+‘+ñˆEô 	% U¨D×,<Ñ,<Ô=ô 	(¨¨t×/TÑ/TÔUàˆrp   c                 óD  — | j                   €J ‚| j                   j                  |«       | j                  €J ‚|j                  j                  t
        j                  k(  r6| j                  r*| j                  j                  | j                  «      | _        | j                  j                  | j                  «      | _        t        d| j                  «      rt        d«       n | j                  | j                  «      | _        | j                  «        |j                  j                  t
        j                  k(  r| j!                  |«       | j#                  «        |j                  j                  t
        j                  k(  r!t%        | j&                  | j(                  «       y y )NÚconfigure_sharded_modelzÉYou have overridden `LightningModule.configure_sharded_model` hook. It will assume that all the layers are already wrapped for sharding and won't wrap the entire model using `FullyShardedDataParallel`.)rN   Úsetupr¦   ÚstateÚfnr<   ÚFITTINGÚ_layer_syncÚapplyrR   Úconvert_moduler=   Úlightning_moduler>   r³   ÚbarrierÚsetup_optimizersÚsetup_precision_pluginr2   Ú
optimizersru   )rm   Útrainers     ro   r¶   zFSDPStrategy.setupF  s3  € à×ÑÐ+Ð+Ð+Ø×Ñ×Ñ˜wÔ'à�z‰zÐ%Ð%Ð%Ø�=‰=×Ñœy×0Ñ0Ò0°T×5EÒ5EØ×)Ñ)×/Ñ/°·
±
Ó;ˆDŒJà×*Ñ*×9Ñ9¸$¿*¹*ÓEˆŒ
äÐ2°D×4IÑ4IÔJäðvõð
 ×*Ñ*¨4¯:©:Ó6ˆDŒJØ�‰Œà�=‰=×Ñœy×0Ñ0Ò0Ø×!Ñ! 'Ô*Ø×#Ñ#Ô%Ø�=‰=×Ñœy×0Ñ0Ò0Ü! $§/¡/°4×3CÑ3CÕDð 1rp   c                 ó<  •— | j                  «        | j                  j                  d«      rt        ‰| �  |«      S d}	 t        ‰| �  |«       |st        d„ | j                  D «       «      rt        d«      ‚y # t
        $ r}dt        |«      vr‚ d}Y d }~ŒHd }~ww xY w)Nrd   Fz%optimizer got an empty parameter listTc              3   ó4   K  — | ]  }t        |«       –— Œ y ­wrr   )r$   )r«   Ú	optimizers     ro   r­   z0FSDPStrategy.setup_optimizers.<locals>.<genexpr>u  s   è ø€ Ð&rÑbqÐU^Ô+EÀiÓ+PÔ'PÑbqùs   ‚z×The optimizer does not seem to reference any FSDP parameters. HINT: Make sure to create the optimizer after setting up the model by referencing `self.trainer.model.parameters()` in the `configure_optimizers()` hook.)	Ú _reset_optimizers_and_schedulersr`   r™   re   r¿   rj   Ústrr°   rÁ   )rm   rÂ   Úinvalid_params_errorÚexrn   s       €ro   r¿   zFSDPStrategy.setup_optimizersa  s£   ø€ ð
 	×-Ñ-Ô/à�;‰;�?‰?Ð,Ô-Ü‘7Ñ+¨GÓ4Ð4à$Ðð	(ô ‰GÑ$ WÔ-ñ  ¤3Ñ&rÐbf×bqÒbqÓ&rÔ#räð2óð ð
 øô ò 	(Ø6¼cÀ"»gÑEØØ#'Õ ûð	(ús   ¿A8 Á8	BÂBÂBc                  ó   — y rr   r‰   rt   s    ro   Úmodel_to_devicezFSDPStrategy.model_to_device~  ó   € ð 	rp   Ú
empty_init)NNNc              #   óâ   K  — |rt        j                  d«      n	t        «       }|5  | j                  j	                  «       5  d –— d d d «       d d d «       y # 1 sw Y   ŒxY w# 1 sw Y   y xY w­w)NÚmeta)ÚtorchÚdevicer   rR   Útensor_init_context)rm   rÍ   Úempty_init_contexts      ro   rÒ   z FSDPStrategy.tensor_init_contextƒ  sN   è ø€ ñ 6@œUŸ\™\¨&Ô1Ä[Ã]ÐÚ ×!6Ñ!6×!JÑ!JÕ!LÛ÷ "M×Ð×!LÐ!Lú×Ðüs4   ‚$A/¦A#ÁAÁA#Á	A/ÁA 	ÁA#Á#A,Á(A/c           	   #   óB  K  — t         j                  | j                  j                  › d�«       ddlm} ddlm}  |d|| j                  | j                  | j                  | j                  j                  dœ| j                  ¤Ž5  d –— d d d «       y # 1 sw Y   y xY w­w)Nz : entered model_sharded_context.r   r¨   )Úenable_wrap)Úwrapper_clsrU   rV   r[   r�   r‰   )r“   r”   rn   r•   Ú2torch.distributed.fsdp.fully_sharded_data_parallelr©   Útorch.distributed.fsdp.wraprÕ   rU   r|   r[   ru   r²   r`   )rm   r©   rÕ   s      ro   Úmodel_sharded_contextz"FSDPStrategy.model_sharded_context�  s‰   è ø€ ô 	�	‰	�T—^‘^×,Ñ,Ð-Ð-MÐNÔOÝ_Ý;áð 
Ø0Ø×(Ñ(Ø ×7Ñ7Ø"×4Ñ4Ø×&Ñ&×,Ñ,ñ
ð �k‰kó
ó ÷
÷ 
ñ 
üs   ‚BBÂBÂ
	BÂBÂBÚnamec                 óö   — t        «       sy t        j                  j                  «       dk(  r/t        j                  j	                  | j                  «       ¬«       y t        j                  j	                  «        y )NÚnccl)Ú
device_ids)r'   rÐ   ÚdistributedÚget_backendr¾   Ú_determine_device_ids)rm   rÚ   s     ro   r¾   zFSDPStrategy.barrierž  sT   € ä*Ô,ØÜ×Ñ×(Ñ(Ó*¨fÒ4Ü×Ñ×%Ñ%°×1KÑ1KÓ1MÐ%ÕNä×Ñ×%Ñ%Õ'rp   ÚobjÚsrcc                 óŠ   — t        «       s|S |g}t        j                  j                  ||t        j
                  ¬«       |d   S )Nr+   r   )r'   rÐ   rÞ   Úbroadcast_object_listÚ_groupÚWORLD)rm   rá   râ   s      ro   Ú	broadcastzFSDPStrategy.broadcast§  s<   € ä*Ô,ØˆJàˆeˆÜ×Ñ×/Ñ/°°SÄÇÁÐ/ÔMØ�1‰vˆrp   Útensorr,   Ú	reduce_opc                 óB   — t        |t        «      rt        |||¬«      S |S )a  Reduces a tensor from several distributed processes to one aggregated tensor.

        Args:
            tensor: the tensor to sync and reduce
            group: the process group to gather results from. Defaults to all processes (world)
            reduce_op: the reduction operation. Defaults to 'mean'/'avg'.
                Can also be a string 'sum' to calculate the sum during reduction.

        Return:
            reduced value, except when the input was not a tensor the output remains is unchanged

        )ré   )r{   r   r*   )rm   rè   r,   ré   s       ro   ÚreducezFSDPStrategy.reduce°  s"   € ô& �fœfÔ%Ü)¨&°%À9ÔMÐMØˆrp   c                 ó0   — | j                   j                  gS rr   )ru   r²   rt   s    ro   rà   z"FSDPStrategy._determine_device_idsÇ  s   € Ø× Ñ ×&Ñ&Ð'Ð'rp   c                 óN  — t         j                  | j                  j                  › d�«       | j                  }|��|j
                  �u|j
                  j                  j                  t        j                  k(  rD| j                  r8| j                  €J ‚| j                  j                  | j                  «      | _        | j                  €J ‚| j                  €J ‚| j                  j                  «        | j                   j                  «        | j                  j                  «        y )Nz: tearing down strategy...)r“   r”   rn   r•   r½   Ú_trainerr·   r¸   r<   r¹   rº   r¦   ÚrevertrP   rN   ÚteardownrR   )rm   Ú	pl_modules     ro   rð   zFSDPStrategy.teardownÊ  sï   € ä�	‰	�T—^‘^×,Ñ,Ð-Ð-GÐHÔIà×)Ñ)ˆ	àÐ!ð ×"Ñ"Ð.Ø×"Ñ"×(Ñ(×+Ñ+¬y×/@Ñ/@Ò@Ø× Ò à—:‘:Ð)Ð)Ð)Ø×)Ñ)×0Ñ0°·±Ó<ˆDŒJà×'Ñ'Ð3Ð3Ð3Ø×ÑÐ+Ð+Ð+Ø× Ñ ×)Ñ)Ô+Ø×Ñ×&Ñ&Ô(Ø×Ñ×!Ñ!Õ#rp   c                 ó   — | j                   S rr   )rL   )Úclss    ro   Úget_registered_strategiesz&FSDPStrategy.get_registered_strategiesà  s   € à×)Ñ)Ð)rp   Ústrategy_registryc                 ó   — t         j                  j                  «       sy |j                  d| d¬«       | j                  j                  d«       |j                  d| dd¬«       | j                  j                  d«       y )NrK   z+Fully Sharded Data Parallel (FSDP) training)ÚdescriptionÚfsdp_cpu_offloadzQFully Sharded Data Parallel (FSDP) training with Full Sharding and CPU OffloadingT)r÷   rU   )rÐ   rÞ   Úis_availableÚregisterrL   Úappend)ró   rõ   s     ro   Úregister_strategiesz FSDPStrategy.register_strategiesä  s�   € ô × Ñ ×-Ñ-Ô/ØØ×"Ñ"ØØØEð 	#ô 	
ð
 	×"Ñ"×)Ñ)¨&Ô1à×"Ñ"ØØØkØð	 	#ô 	
ð 	×"Ñ"×)Ñ)Ð*<Õ=rp   c                 ó^  — | j                   €J ‚| j                  dk(  rt        | j                   «      }nI| j                  dk(  r"t        | j                   | j                  ¬«      }nt        d| j                  › �«      ‚|5  | j                   j                  «       cd d d «       S # 1 sw Y   y xY w)Nr^   rM   ©Ú
world_sizeúUnknown state_dict_type: )r¦   rl   r   r   rÿ   rj   Ú
state_dict)rm   Ústate_dict_ctxs     ro   Úlightning_module_state_dictz(FSDPStrategy.lightning_module_state_dictø  sŠ   € à�z‰zÐ%Ð%Ð%Ø× Ñ  IÒ-Ü<¸T¿Z¹ZÓH‰NØ×"Ñ" fÒ,Ü9¸$¿*¹*ÐQU×Q`ÑQ`Ôa‰NäÐ8¸×9NÑ9NÐ8OÐPÓQÐQÚØ—:‘:×(Ñ(Ó*÷ �^Š^ús   Á?B#Â#B,Ú
checkpointÚstrictc                  ó   — y rr   r‰   )rm   r  r  s      ro   Úload_model_state_dictz"FSDPStrategy.load_model_state_dict  rÌ   rp   rÅ   c                 ó~  — ddl m} ddl m} t        |t        «      r|j
                  }| j                  €J ‚| j                  dk(  r;t        | j                  «      5  |j                  | j                  |«      cd d d «       S | j                  dk(  rt        | j                  | j                  ¬«      5  |j                  | j                  |«      }| j                  dk(  r'|j                  ||j                  | j                  «      }|cd d d «       S t        d| j                  › �«      ‚# 1 sw Y   Œ!xY w# 1 sw Y   Œ-xY w)Nr   r¨   ©ÚOptimStateKeyTyper^   rM   rþ   r   )r¯   r©   r
  r{   r6   Ú
_optimizerr¦   rl   r   Úoptim_state_dictr   rÿ   r†   Úrekey_optim_state_dictÚPARAM_IDrj   )rm   rÅ   ÚFSDPr
  r  s        ro   Úoptimizer_statezFSDPStrategy.optimizer_state	  s  € åKÝ<ä�iÔ!3Ô4Ø!×,Ñ,ˆIà�z‰zÐ%Ð%Ð%Ø× Ñ  IÒ-Ü0°·±Õ<Ø×,Ñ,¨T¯Z©Z¸ÓC÷ =Ñ<ð ×"Ñ" fÒ,Ü-¨d¯j©jÀTÇ_Á_ÖUØ!×2Ñ2°4·:±:¸yÓI�
Ø×#Ñ# qÒ(à!%×!<Ñ!<¸ZÐIZ×IcÑIcÐei×eoÑeoÓ!p�JØ!÷ VÑUô Ð4°T×5JÑ5JÐ4KÐLÓMÐM÷ =Ð<ú÷ VÐUús   ÁD'Â1AD3Ä'D0Ä3D<c                  ó   — y rr   r‰   )rm   r  s     ro   Úload_optimizer_state_dictz&FSDPStrategy.load_optimizer_state_dict   rÌ   rp   ÚfilepathÚstorage_optionsc                 ó  •— |�t        d«      ‚t        | j                  |«      «      }|j                  «       r(| j                  dk(  rt        |«      st        d|› �«      ‚| j                  dk(  rÁ|j                  «       r|j                  «        |j                  dd¬«       d|j                  d«      i}|j                  t        |j                  d	g «      «      D ��ci c]  \  }}d
|› �|“Œ c}}«       t        ||«       | j                  dk(  rt        j                   ||t"        z  «       y y | j                  dk(  r1t        |«      rt%        j&                  |«       t(        ‰| �U  ||¬«      S t-        d| j                  › �«      ‚c c}}w )Nz�`FSDPStrategy.save_checkpoint(..., storage_options=...)` is not supported because `FSDPStrategy` does not use the `CheckpointIO`.rM   z/The checkpoint path exists and is a directory: r^   T)ÚparentsÚexist_okr¦   r  Úoptimizer_statesÚ
optimizer_r   )r  r  r   )r‚   r   rç   Úis_dirrl   r"   ÚIsADirectoryErrorÚis_fileÚunlinkÚmkdirÚpopÚupdateÚ	enumerater   r†   rÐ   Úsaver   ÚshutilÚrmtreere   Úsave_checkpointrj   )	rm   r  r  r  ÚpathÚconverted_stateÚidxÚoptim_statern   s	           €ro   r%  zFSDPStrategy.save_checkpoint%  s  ø€ ð Ð&ÜðCóð ô
 �D—N‘N 8Ó,Ó-ˆØ�;‰;Œ=˜T×2Ñ2°fÒ<ÔE[Ð\`ÔEaÜ#Ð&UÐVZÐU[Ð$\Ó]Ð]à× Ñ  IÒ-Ø�|‰|Œ~Ø—‘”Ø�J‰J˜t¨dˆJÔ3à&¨
¯©°|Ó(DÐEˆOØ×"Ñ"ä(1°*·.±.ÐASÐUWÓ2XÔ(Yô$á(YÑ$�C˜ð ˜S˜EÐ" KÑ/Ø(Yò$ô ô
 )¨¸$Ô?à×Ñ 1Ò$Ü—
‘
˜: tÔ.@Ñ'@ÕAð %à×"Ñ" fÒ,Ü% dÔ+Ü—‘˜dÔ#Ü‘7Ñ*°jÈ4Ð*ÓPÐPäÐ8¸×9NÑ9NÐ8OÐPÓQÐQùó$s   ÃF	
Úcheckpoint_pathÚweights_onlyc                 ón  — t        | j                  |«      «      }ddlm} | j                  €J ‚| j
                  €J ‚t        |«      �rZddlm} t        | j                  «      }|5  d| j                  j                  «       i}t        ||«       | j                  j                  |d   | j
                  j                  ¬«       | j
                  j                  j                  j                   t"        j$                  k(  r}| j&                  rqddlm}  ||¬«      }	t-        | j&                  «      D ]J  \  }
}d|
› �} ||d   ||	¬	«      }|j/                  ||   | j                  |¬
«      }|j                  |«       ŒL d d d «       t1        j2                  |t4        z  |¬«      }|S t7        |«      �rÖt9        |«      }t;        |j=                  d«      | j                  | j>                  | j
                  j                  ¬«       tA        |«      }ddlm} ddlm!} |jE                  d«      }|�;| j
                  j                  j                  j                   t"        j$                  k7  r|S tG        | j&                  «      tG        |«      k7  r.tI        dtG        | j&                  «      › dtG        |«      › d�«      ‚tK        | j                  | j>                  d¬«      5  tM        | j&                  |«      D ]ˆ  \  }}tO        tQ        |d   jS                  «       «      d   tT        «      r'|jW                  ||jX                  | j                  «      }|j/                  || j                  |¬
«      }|j                  |«       ŒŠ 	 d d d «       |S t[        dt]        |«      ›d�«      ‚# 1 sw Y   �Œ$xY w# 1 sw Y   |S xY w)Nr   r¨   )Ú!load_sharded_optimizer_state_dictr¦   )r  )ÚFileSystemReader)r&  r  )Úmodel_state_dictÚoptimizer_keyÚstorage_reader)r  r¦   Úoptim)r+  r  )r®   rÿ   r  r	  r  zYou have configured z( optimizers but the checkpoint contains z€ optimizers to load. Please resume training with the same number of optimizers or edit the checkpoint manually to remove states.F)rÿ   Ú
rank0_onlyr·   z	The path zœ does not point to a valid checkpoint. Make sure the path points to either a directory with FSDP checkpoint shards, or a single file with a full checkpoint.)/r   rç   r¯   r©   r¦   r½   r"   Ú&torch.distributed.checkpoint.optimizerr-  r   r  r   Úload_state_dictÚstrict_loadingrÂ   r·   r¸   r<   r¹   rÁ   Útorch.distributed.checkpointr.  r!  Úoptim_state_dict_to_loadrÐ   Úloadr   r!   r0   r&   r  rÿ   r1   r
  r™   rw   ÚRuntimeErrorr   Úzipr{   ÚlistÚkeysÚintr  Ú
PARAM_NAMErj   rÇ   )rm   r*  r+  r&  r  r-  r  Úmodule_stater.  Úreaderr(  r2  Ú	optim_keyr)  Úflattened_osdÚmetadatar  r
  r  rÅ   Ú	opt_states                        ro   Úload_checkpointzFSDPStrategy.load_checkpointI  sU  € ô �D—N‘N ?Ó3Ó4ˆåKà�z‰zÐ%Ð%Ð%Ø×$Ñ$Ð0Ð0Ð0ä! $Õ'Ý`ä<¸T¿Z¹ZÓHˆNâØ '¨¯©×)>Ñ)>Ó)@ÐA�Ü,¨\¸4Ô@Ø—
‘
×*Ñ*¨<¸Ñ+@È×I^ÑI^×ImÑImÐ*Ônà×(Ñ(×0Ñ0×6Ñ6×9Ñ9¼Y×=NÑ=NÒNÐSW×SbÒSbÝMñ .°4Ô8�Fä&/°·±Ö&@™
˜˜UØ&0°°Ð$6˜	Ù&GØ-9¸'Ñ-BØ*3Ø+1ô'˜ð
 )-×(EÑ(EØ-8¸Ñ-CØ"&§*¡*Ø"'ð )Fó )˜ð
 ×-Ñ-¨mÕ<ð 'A÷  ô6 —z‘z $Ô);Ñ";È,ÔWˆHØˆOä˜tÕ$Ü# DÓ)ˆJÜ"Ø—‘˜|Ó,Ø—z‘zØŸ?™?Ø×,Ñ,×;Ñ;õ	ô .¨jÓ9ˆJåOÝ@à)Ÿ~™~Ð.@ÓAÐØÐ'¨4×+@Ñ+@×+HÑ+H×+NÑ+N×+QÑ+QÔU^×UfÑUfÒ+fà!Ð!Ü�4—?‘?Ó#¤sÐ+;Ó'<Ò<Ü"Ø*¬3¨t¯©Ó+?Ð*@ð AÜÐ,Ó-Ð.ð /WðWóð ô .¨d¯j©jÀTÇ_Á_ÐafÖgÜ,/°·±ÐAQÖ,RÑ(�I˜yÜ!¤$ y°Ñ'9×'>Ñ'>Ó'@Ó"AÀ!Ñ"DÄcÔJà$(×$?Ñ$?À	ÐK\×KgÑKgÐim×isÑisÓ$t˜	à $× =Ñ =Ø)2Ø"Ÿj™jØ'ð !>ó !�Ið
 ×-Ñ-¨iÕ8ñ -S÷ hð ÐäØœ˜D›	�}ð %_ð _ó
ð 	
÷S  ‘ú÷t hð Ðús   Á&DNËB"N*ÎN'Î*N4)ra   N)rÂ   z
pl.Trainerra   Nrr   )r   )NÚmean)T)Ir•   Ú
__module__Ú__qualname__Ú__doc__Ústrategy_namerL   r<  rÇ   Ú__annotations__r   r   rÐ   rÑ   r   r   r7   r   r   Úboolr˜   r   r   rš   r>  r
   rf   Úpropertyr   ru   rx   rS   r|   r8   rR   ÚsetterÚdictr‡   rŠ   rŒ   r’   r—   r–   r¥   r³   r¶   r¿   rË   r   r   rÒ   rÙ   r¾   r;   rç   r   r5   rë   rà   rð   Úclassmethodrô   r   rü   r  r   r  r   r  r  r4   r%  rF  Ú__classcell__)rn   s   @ro   rJ   rJ   \   s  ø… ñ1ðf €MØ(*Ð˜D ™IÓ*ð @DØ9=Ø<@Ø04Ø04Ø/3Ø'9Ø7;Ø6:Ø04ØVZØ?CØ2>Ø6<ØAEñ!/0àÐ;Ñ<ð/0ð # 4¨¯©Ñ#5Ñ6ð/0ð &Ð&8Ñ9ð	/0ð
   Ñ-ð/0ð # 9Ñ-ð/0ð  (¨™}ð/0ð ˜)Ñ$ð/0ð ˜4 ¨tÐ3Ñ4ð/0ð "Ð"2Ñ3ð/0ð # 9Ñ-ð/0ð #+¨5°°f±¸tÀDÈÁLÑ?QÐ1QÑ+RÑ"Sð/0ð *2°)Ñ)<ð/0ð 0ð/0ð !Ð!2Ñ3ð/0ð  ˜e E¨#¡J°Ð$<Ñ=Ñ>ð!/0ð" ð#/0ð$ 
õ%/0ðb Øð6˜UŸ\™\ò 6ó ó ð6ð ðV˜sò Vó ðVð ð+ x°¡}ò +ó ð+ð ð¨Ð1AÑ(Bò ó ðð Øð( -ò (ó ó ð(ð ×ÑØð2°¸)Ñ1Dð 2Èò 2ó ó ð2ð Øða¨Dò aó ó ðað Øð°ò ó ó ðð Øð¨Tò ó ó ðð ô^ó ð^ð*n¨Có nóKð òuó ðuð
 ð  &ð  ¨Vò  ó ð ðD òEó ðEð4 ôó ðð8 òó ðð Øñ¨h°t©nð È	ÐRbÑHcò ó ó ðð Øð yÐ1AÑ'Bò ó ó ðð ñ(˜H S™Mð (°Tò (ó ð(ð ñ˜Zð ¨cð ¸*ò ó ðð ð  $Ø4:ñ	à�f˜c�kÑ"ðð ˜‰}ðð ˜E (¨C -Ñ0Ñ1ð	ð
 
òó ðð,( t¨C¡yó (ð ò$ó ð$ð* ð*¨$¨s©)ò *ó ð*ð Øð>Ð4Eð >È$ò >ó ó ð>ð$ ð	+¨T°#°s°(©^ò 	+ó ð	+ð ñ°¸¸S¸Ñ0Að È4ð Ð[_ò ó ðð ðN¨ð N°t¸CÀ¸KÑ7Hò Nó ðNð, ð°G¸CÀ¸HÑ4Eð È$ò ó ðð à\`ñ!RØ˜s C˜x™.ð!RØ49ð!RØLTÐUXÉMð!Rà	ô!Ró ð!RðF ñZ
¨uð Z
ÀHÈTÁNð Z
Ð^bÐcfÐhkÐckÑ^lò Z
ó ôZ
rp   rJ   )sÚloggingr#  Úcollections.abcr   r   Ú
contextlibr   r   Údatetimer   Úpathlibr   Útypingr	   r
   r   r   r   r   rÐ   Ú"lightning_utilities.core.rank_zeror   r¡   r   Útorch.nnr   Útorch.optimr   Útyping_extensionsr   Úpytorch_lightningÚplÚlightning_fabric.pluginsr   r   Ú5lightning_fabric.plugins.collectives.torch_collectiver   Úlightning_fabric.strategiesr   Ú lightning_fabric.strategies.fsdpr   r   r   r   r   r   r   r   r    r!   r"   r#   r$   r%   Ú*lightning_fabric.strategies.model_parallelr&   Ú&lightning_fabric.utilities.distributedr'   r(   r)   r*   r,   rå   Ú"lightning_fabric.utilities.importsr-   r.   Úlightning_fabric.utilities.initr/   Úlightning_fabric.utilities.loadr0   r1   Ú$lightning_fabric.utilities.optimizerr2   Úlightning_fabric.utilities.seedr3   Ú lightning_fabric.utilities.typesr4   r5   Ú pytorch_lightning.core.optimizerr6   Ú#pytorch_lightning.plugins.precisionr7   Ú(pytorch_lightning.plugins.precision.fsdpr8   Ú8pytorch_lightning.strategies.launchers.subprocess_scriptr9   Ú%pytorch_lightning.strategies.parallelr:   Ú%pytorch_lightning.strategies.strategyr;   Ú pytorch_lightning.trainer.statesr<   Ú)pytorch_lightning.utilities.model_helpersr=   Ú%pytorch_lightning.utilities.rank_zeror>   r?   r›   r@   r×   rA   rB   rC   rØ   rD   Úsetr˜   rM  r>  rX   r\   Ú	getLoggerr•   r“   rJ   r‰   rp   ro   Ú<module>rv     s  ðó Û ß .ß 2Ý Ý ÷÷ ó Ý UÝ Ý Ý !Ý &ã ß EÝ TÝ 9÷÷ ÷ ÷ õ  N÷ó õ Cß aÝ Rß LÝ FÝ 6ß <Ý ?Ý 9Ý BÝ ^Ý BÝ <Ý 6Ý Cß `Ñ `áÝ8ßoÑoÝ<à�C˜˜V™Ñ% x°¸¸sÐ0CÀTÐ0IÑ'JÐL\Ð\Ñ]€GØÐ/°Ð9rÑ1sÐsÑtÐð €g×Ñ˜Ó!€ôH	
Ð#õ H	
rp   