Ë
      çi+  ã                   óp  — d dl mZmZ d dlmZ d dlmZmZmZm	Z	 d dl
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 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* d dl+m,Z,m-Z-m.Z.m/Z/m0Z0 d dl+m1Z2 d dl3m4Z4 d dl5mZ dZ6 G d„ de%«      Z7 G d„ de*«      Z8y)é    )ÚAbstractContextManagerÚnullcontext)Ú	timedelta)ÚAnyÚLiteralÚOptionalÚUnionN)Úrank_zero_only)ÚTensor)ÚModule)ÚDistributedDataParallel)Úoverride)ÚAccelerator)Údefault_pg_timeout)ÚClusterEnvironment)ÚCheckpointIO)Ú	Precision)Ú_MultiProcessingLauncher)Ú_SubprocessScriptLauncher)ÚParallelStrategy)Ú_StrategyRegistry)Ú
TBroadcastÚ_BackwardSyncControl)ÚReduceOpÚ_distributed_is_initializedÚ-_get_default_process_group_backend_for_deviceÚ_init_dist_connectionÚ_sync_ddp_if_available©Úgroup)Ú_TORCH_GREATER_EQUAL_2_3)Úddp_forkÚddp_notebookc                   ó  ‡ — e Zd ZdZddddddedfdee   deeej                        dee
   dee   dee   d	ee   d
ee   ded   deddfˆ 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d/d„«       Zed/ˆ fd„«       Zededefd„«       Z ededdfd„«       Z!e	 d0d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d1d"e'd#ede'fd$„«       Z(ededeee#ee"f   f   fˆ fd%„«       Z)e	 d2ded&eee#ee"f   f   d'e*ddfˆ fd(„«       Z+e,ed)e-ddfd*„«       «       Z.d/d+„Z/defd,„Z0d/d-„Z1deee      fd.„Z2ˆ xZ3S )3ÚDDPStrategyzKStrategy for multi-process single-device training on one or multiple nodes.NÚpopenÚacceleratorÚparallel_devicesÚcluster_environmentÚcheckpoint_ioÚ	precisionÚprocess_group_backendÚtimeoutÚstart_method)r&   ÚspawnÚforkÚ
forkserverÚkwargsÚreturnc	                 ó’   •— t         ‰
| �  |||||¬«       d| _        || _        || _        || _        t        «       | _        |	| _        y )N)r'   r(   r)   r*   r+   é   )	ÚsuperÚ__init__Ú
_num_nodesÚ_process_group_backendÚ_timeoutÚ_start_methodÚ_DDPBackwardSyncControlÚ_backward_sync_controlÚ_ddp_kwargs)Úselfr'   r(   r)   r*   r+   r,   r-   r.   r2   Ú	__class__s             €út/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/lightning_fabric/strategies/ddp.pyr7   zDDPStrategy.__init__8   sY   ø€ ô 	‰ÑØ#Ø-Ø 3Ø'Øð 	ô 	
ð ˆŒØ5JˆÔ#Ø-4ˆŒØ)ˆÔÜ&=Ó&?ˆÔ#Ø!ˆÕó    c                 óP   — | j                   €J ‚| j                   | j                     S ©N)r(   Ú
local_rank©r?   s    rA   Úroot_devicezDDPStrategy.root_deviceR   s+   € ð ×$Ñ$Ð0Ð0Ð0Ø×$Ñ$ T§_¡_Ñ5Ð5rB   c                 ó   — | j                   S rD   ©r8   rF   s    rA   Ú	num_nodeszDDPStrategy.num_nodesX   s   € à�‰ÐrB   rJ   c                 ó   — || _         y rD   rI   )r?   rJ   s     rA   rJ   zDDPStrategy.num_nodes\   s   € ð $ˆ�rB   c                 óH   — | j                   �t        | j                   «      S dS )Nr   )r(   ÚlenrF   s    rA   Únum_processeszDDPStrategy.num_processesa   s$   € à-1×-BÑ-BÐ-NŒs�4×(Ñ(Ó)ÐUÐTUÐUrB   c                 óN   — | j                   | j                  z  | j                  dœS )N)Únum_replicasÚrank)rJ   rN   Úglobal_rankrF   s    rA   Údistributed_sampler_kwargsz&DDPStrategy.distributed_sampler_kwargse   s$   € ð "&§¡°$×2DÑ2DÑ!DÈt×O_ÑO_Ñ`Ð`rB   c                 ó   — | j                   S rD   )r9   rF   s    rA   r,   z!DDPStrategy.process_group_backendj   s   € à×*Ñ*Ð*rB   c                 óØ   — | j                   €J ‚| j                  dk(  r1t        | j                   | j                  | j                  «      | _        y t        | | j                  ¬«      | _        y )Nr&   )r.   )r)   r;   r   rN   rJ   Ú	_launcherr   rF   s    rA   Ú_configure_launcherzDDPStrategy._configure_launchern   sZ   € à×'Ñ'Ð3Ð3Ð3Ø×Ñ Ò(Ü6°t×7OÑ7OÐQU×QcÑQcÐei×esÑesÓtˆD�Nä5°dÈ×I[ÑI[Ô\ˆD�NrB   c                 óB   •— t         ‰| �  «        | j                  «        y rD   )r6   Úsetup_environmentÚ_setup_distributed)r?   r@   s    €rA   rY   zDDPStrategy.setup_environmentv   s   ø€ ä‰Ñ!Ô#Ø×ÑÕ!rB   Úmodulec                 ó  — | j                  «       }|�;t        j                  j                  t        j                  j	                  «       «      n	t        «       }|5  t        d||dœ| j                  ¤Žcddd«       S # 1 sw Y   yxY w)z^Wraps the model into a :class:`~torch.nn.parallel.distributed.DistributedDataParallel` module.N)r[   Ú
device_ids© )Ú_determine_ddp_device_idsÚtorchÚcudaÚstreamÚStreamr   r   r>   )r?   r[   r]   Úctxs       rA   Úsetup_modulezDDPStrategy.setup_module{   sd   € ð ×3Ñ3Ó5ˆ
à8BÐ8NŒe�j‰j×Ñ¤§
¡
× 1Ñ 1Ó 3Ô4ÔT_ÓTaˆÚÜ*Ðd°&ÀZÑdÐSW×ScÑScÑd÷ �SŠSús   ÁA<Á<Bc                 ó:   — |j                  | j                  «       y rD   )ÚtorG   )r?   r[   s     rA   Úmodule_to_devicezDDPStrategy.module_to_device„   s   € à�	‰	�$×"Ñ"Õ#rB   Ú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

        )rj   )Ú
isinstancer   r   )r?   ri   r    rj   s       rA   Ú
all_reducezDDPStrategy.all_reduceˆ   s"   € ô  �fœfÔ%Ü)¨&°%À9ÔMÐMØˆrB   Úargsc                 óÂ  — t        «       sy t        j                  j                  «       dk(  r/t        j                  j	                  | j                  «       ¬«       y 	 t        j                  j	                  «        y # t        $ rY}dt        |«      v rAt        j                  d| j                  ¬«      }t        j                  j                  |«       n‚ Y d }~y d }~ww xY w)NÚnccl)r]   ÚPrivateUse1HooksInterfaceg        )Údevice)r   r`   ÚdistributedÚget_backendÚbarrierr_   ÚRuntimeErrorÚstrri   rG   rm   )r?   rn   r2   ÚeÚdummy_tensors        rA   ru   zDDPStrategy.barrierœ   s©   € ä*Ô,ØÜ×Ñ×(Ñ(Ó*¨fÒ4Ü×Ñ×%Ñ%°×1OÑ1OÓ1QÐ%ÕRð	Ü×!Ñ!×)Ñ)Õ+øÜò Ø.´#°a³&Ñ8ô $)§<¡<°¸D×<LÑ<LÔ#M�LÜ×%Ñ%×0Ñ0°Õ>àô ?ûðús   ÁA< Á<	CÂACÃCÚobjÚsrcc                 óŠ   — t        «       s|S |g}t        j                  j                  ||t        j
                  ¬«       |d   S )Nr   r   )r   r`   rs   Úbroadcast_object_listÚ_groupÚWORLD)r?   rz   r{   s      rA   Ú	broadcastzDDPStrategy.broadcast¯   s<   € ä*Ô,ØˆJàˆeˆÜ×Ñ×/Ñ/°°SÄÇÁÐ/ÔMØ�1‰vˆrB   c                 óZ   •— t        |t        «      r|j                  }t        ‰| �  |«      S rD   )rl   r   r[   r6   Úget_module_state_dict)r?   r[   r@   s     €rA   r‚   z!DDPStrategy.get_module_state_dict¸   s'   ø€ ä�fÔ5Ô6Ø—]‘]ˆFÜ‰wÑ,¨VÓ4Ð4rB   Ú
state_dictÚstrictc                 ób   •— t        |t        «      r|j                  }t        ‰| �  |||¬«       y )N)r[   rƒ   r„   )rl   r   r[   r6   Úload_module_state_dict)r?   r[   rƒ   r„   r@   s       €rA   r†   z"DDPStrategy.load_module_state_dict¾   s.   ø€ ô �fÔ5Ô6Ø—]‘]ˆFÜ‰Ñ&¨fÀÐTZÐ&Õ[rB   Ústrategy_registryc                 óz   — d}|D ]  \  }}|j                  || d|›d�|¬«       Œ  |j                  d| ddd¬	«       y )
N))Úddpr&   )Ú	ddp_spawnr/   )r"   r0   )r#   r0   z DDP strategy with `start_method=Ú`)Údescriptionr.   Úddp_find_unused_parameters_truezBAlias for `find_unused_parameters_true` and `start_method='popen'`Tr&   )rŒ   Úfind_unused_parametersr.   )Úregister)Úclsr‡   ÚentriesÚnamer.   s        rA   Úregister_strategieszDDPStrategy.register_strategiesÆ   sg   € ð
ˆó #*ÑˆD�,Ø×&Ñ&ØØØ>¸|Ð>NÈaÐPØ)ð	 'õ ð #*ð 	×"Ñ"Ø-ØØ\Ø#'Ø ð 	#õ 	
rB   c                 ó(  — | 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)	Ú_set_world_ranksÚ_get_process_group_backendr9   r)   r:   r!   rG   Útyper   )r?   r2   s     rA   rZ   zDDPStrategy._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]Ó^rB   c                 óH   — | j                   xs t        | j                  «      S rD   )r9   r   rG   rF   s    rA   r˜   z&DDPStrategy._get_process_group_backendç   s    € Ø×*Ñ*ÒmÔ.[Ð\`×\lÑ\lÓ.mÐmrB   c                 ó>  — | j                   �q| j                   j                  | j                  | j                  z  | j                  z   «       | j                   j                  | j                  | j                  z  «       | j                  xt        _	        t        _	        y rD   )r)   Úset_global_rankÚ	node_rankrN   rE   Úset_world_sizerJ   rR   r
   rQ   Úutils_rank_zero_onlyrF   s    rA   r—   zDDPStrategy._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Õ7rB   c                 óf   — | j                   j                  dk(  rd S | j                   j                  gS )Nr•   )rG   r™   ÚindexrF   s    rA   r_   z%DDPStrategy._determine_ddp_device_idsò   s/   € Ø×'Ñ'×,Ñ,°Ò5ˆtÐS¸D×<LÑ<L×<RÑ<RÐ;SÐSrB   )r3   N)NÚmean)r   )T)4Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r   Úlistr`   rr   r   r   r   rw   r   r   r   r7   Úpropertyr   rG   ÚintrJ   ÚsetterrN   ÚdictrS   r,   rW   rY   r   r   re   rh   r   r	   r   rm   ru   r   r€   r‚   Úboolr†   Úclassmethodr   r“   rZ   r˜   r—   r_   Ú__classcell__)r@   s   @rA   r%   r%   5   s€  ø„ ÙUð .2Ø9=Ø<@Ø04Ø)-Ø/3Ø'9ØHOñ"à˜kÑ*ð"ð # 4¨¯©Ñ#5Ñ6ð"ð &Ð&8Ñ9ð	"ð
   Ñ-ð"ð ˜IÑ&ð"ð  (¨™}ð"ð ˜)Ñ$ð"ð ÐDÑEð"ð ð"ð 
õ"ð4 Øð6˜UŸ\™\ò 6ó ó ð6ð ð˜3ò ó ðð ×Ñð$ 3ð $¨4ò $ó ð$ð ðV˜sò Vó ðVð Øða¨D°°c°©Nò aó ó ðað ð+ x°¡}ò +ó ð+ð ò]ó ð]ð ô"ó ð"ð ðe 6ð eÐ.Eò eó ðeð ð$ vð $°$ò $ó ð$ð àgmñØðØ%-¨c¡]ðØFNÈuÐU]Ð_bÐUbÑOcÑFdðà	òó ðð& ð˜Sð ¨Cð °Dò ó ðð$ ñ˜Zð ¨cð ¸*ò ó ðð ð5¨Fð 5°t¸CÀÀsÈFÀ{ÑASÐ<SÑ7Tô 5ó ð5ð
 àX\ñ\Øð\Ø*.¨s°E¸#¸v¸+Ñ4FÐ/FÑ*Gð\ØQUð\à	ô\ó ð\ð Øð
Ð4Eð 
È$ò 
ó ó ð
ó,_ðn¨Có nóKðT¨8°D¸±IÑ+>÷ TrB   r%   c                   ó*   — e Zd Zedededefd„«       Zy)r<   r[   Úenabledr3   c                 óÎ   — |s
t        «       S t        |t        «      s:t        d| j                  j
                  › d|j                  j
                  › d�«      ‚|j                  «       S )z{Blocks gradient synchronization inside the :class:`~torch.nn.parallel.distributed.DistributedDataParallel`
        wrapper.zABlocking backward sync is only possible if the module passed to `zA.no_backward_sync` is wrapped in `DistributedDataParallel`. Got: Ú.)r   rl   r   Ú	TypeErrorr@   r£   Úno_sync)r?   r[   r°   s      rA   Úno_backward_syncz(_DDPBackwardSyncControl.no_backward_sync÷   si   € ñ Ü“=Ð ä˜&Ô"9Ô:ÜðØ—^‘^×,Ñ,Ð-ð .Ø×)Ñ)×2Ñ2Ð3°1ð6óð ð
 �~‰~ÓÐrB   N)r£   r¤   r¥   r   r   r¬   r   rµ   r^   rB   rA   r<   r<   ö   s*   „ Øð  vð  ¸ð  ÐAWò  ó ñ rB   r<   )9Ú
contextlibr   r   Údatetimer   Útypingr   r   r   r	   r`   Útorch.distributedÚ"lightning_utilities.core.rank_zeror
   rŸ   r   Útorch.nnr   Útorch.nn.parallel.distributedr   Útyping_extensionsr   Ú)lightning_fabric.accelerators.acceleratorr   Ú5lightning_fabric.plugins.collectives.torch_collectiver   Ú9lightning_fabric.plugins.environments.cluster_environmentr   Ú)lightning_fabric.plugins.io.checkpoint_ior   Ú"lightning_fabric.plugins.precisionr   Ú5lightning_fabric.strategies.launchers.multiprocessingr   Ú7lightning_fabric.strategies.launchers.subprocess_scriptr   Ú$lightning_fabric.strategies.parallelr   Ú$lightning_fabric.strategies.registryr   Ú$lightning_fabric.strategies.strategyr   r   Ú&lightning_fabric.utilities.distributedr   r   r   r   r   r    r~   Ú"lightning_fabric.utilities.importsr!   Ú$lightning_fabric.utilities.rank_zeroÚ_DDP_FORK_ALIASESr%   r<   r^   rB   rA   Ú<module>rÌ      sƒ   ð÷ ;Ý ß 0Ó 0ã Û Ý UÝ Ý Ý AÝ &å AÝ TÝ XÝ BÝ 8Ý ZÝ ]Ý AÝ Bß Q÷õ õ CÝ GÝ ?ðÐ ô~TÐ"ô ~TôB Ð2õ  rB   