
      i                         d dl mZmZ d dlmZ d dlmZ d dl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mZ d d	lmZmZ d d
lmZ d dlmZ d dlmZ  G d dee      Zy)    )ABCabstractmethod)	Generator)contextmanager)AnyOptionalN)Tensor)override)CheckpointIOClusterEnvironment)ReduceOp_all_gather_ddp_if_available)	LayerSync)	Precision)Strategyc                   L    e Zd ZdZ	 	 	 	 	 dded   deeej                        dee   dee	   dee
   f
 fd	Zeeed
ej                  fd                     Zed
efd       Zed
efd       Zed
efd       Zed
efd       Zeed
efd              Zed
eeej                        fd       Zej0                  deeej                        d
dfd       Zed
eeef   fd       Zeddedee   ded
efd       Zeddeded
efd       Ze d
e!fd       Z"ed  fd       Z# xZ$S )!ParallelStrategyz:Strategy for training with multiple processes in parallel.Nacceleratorzpl.accelerators.Acceleratorparallel_devicescluster_environmentcheckpoint_ioprecision_pluginc                 T    t         |   |||       || _        || _        d | _        y )N)r   r   r   )super__init__r   r   _layer_sync)selfr   r   r   r   r   	__class__s         z/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/lightning/pytorch/strategies/parallel.pyr   zParallelStrategy.__init__"   s2     	[`pq 0AT 04    returnc                      y)zReturn the root device.N r   s    r   root_devicezParallelStrategy.root_device/   s    r    c                 R    | j                   | j                   j                         S dS Nr   )r   global_rankr$   s    r   r(   zParallelStrategy.global_rank5   s(    9=9Q9Q9]t''335dcddr    c                 R    | j                   | j                   j                         S dS r'   )r   
local_rankr$   s    r   r*   zParallelStrategy.local_rank9   (    8<8P8P8\t''224cbccr    c                 R    | j                   | j                   j                         S dS r'   )r   	node_rankr$   s    r   r-   zParallelStrategy.node_rank=   s(    7;7O7O7[t''113babbr    c                 R    | j                   | j                   j                         S dS )N   )r   
world_sizer$   s    r   r0   zParallelStrategy.world_sizeA   r+   r    c                      | j                   dk(  S r'   )r(   r$   s    r   is_global_zerozParallelStrategy.is_global_zeroE   s     1$$r    c                     | j                   S N_parallel_devicesr$   s    r   r   z!ParallelStrategy.parallel_devicesJ   s    %%%r    c                     || _         y r4   r5   )r   r   s     r   r   z!ParallelStrategy.parallel_devicesN   s
    !1r    c                 b    | j                   t        | j                         nd| j                  dS )Nr   )num_replicasrank)r   lenr(   r$   s    r   distributed_sampler_kwargsz+ParallelStrategy.distributed_sampler_kwargsR   s3     ;?:O:O:[C 5 56ab$$
 	
r    tensorgroup
sync_gradsc                     t        |||      S )z&Perform a all_gather on all processes.)r>   r?   )r   )r   r=   r>   r?   s       r   
all_gatherzParallelStrategy.all_gatherY   s     ,F%JWWr    decisionallc                     t        j                  t        |      | j                        }| j	                  |t
        j                        }|rt        || j                  k(        }|S t        |      }|S )a  Reduces a boolean decision over distributed processes. By default is analogous to ``all`` from the standard
        library, returning ``True`` only if all input decisions evaluate to ``True``. If ``all`` is set to ``False``,
        it behaves like ``any`` instead.

        Args:
            decision: A single input decision.
            all: Whether to logically emulate ``all`` or ``any``. Defaults to True.

        Returns:
            bool: The reduced boolean decision.

        )device)	reduce_op)	torchr=   intr%   reducer   SUMboolr0   )r   rB   rC   s      r   reduce_boolean_decisionz(ParallelStrategy.reduce_boolean_decision^   si     <<Hd6F6FG;;ll  
 9<4DOO34 BFhr    c              #      K   t        | j                  t        j                  j                  j
                        r(| j                  j                         5  d ddd       yd y# 1 sw Y   yxY ww)zBlocks ddp sync gradients behaviour on backwards pass.

        This is useful for skipping sync when accumulating gradients, reducing communication overhead
        Returns: context manager with sync behaviour off

        N)
isinstancemodelpl	utilitiestypesDistributedDataParallelno_syncr$   s    r   block_backward_syncz$ParallelStrategy.block_backward_synct   sR      djj",,"4"4"L"LM##%
 &% J &%s   AA4A(A4(A1-A4c                 r    | j                   J | j                   j                          t        |           y r4   )r   teardownr   )r   r   s    r   rW   zParallelStrategy.teardown   s2    ''333  ))+r    )NNNNN)NF)T)r!   N)%__name__
__module____qualname____doc__r   listrG   rE   r   r   r   r   propertyr   r
   r%   rH   r(   r*   r-   r0   rK   r2   r   setterdictstrr   r<   r	   rA   rL   r   r   rU   rW   __classcell__)r   s   @r   r   r      sS   D @D9=<@04045;<5 #4#565 &&89	5
  -5 #9-5 &U\\ &   & eS e e dC d d c3 c c dC d d % %  % &(4+="> & & 2$u||:L1M 2RV 2 2 
DcN 
 
 X X XRV Xci X X  4 4  * Y    r    r   )abcr   r   collections.abcr   
contextlibr   typingr   r   rG   r	   typing_extensionsr
   lightning.pytorchpytorchrP   lightning.fabric.pluginsr   r   &lightning.fabric.utilities.distributedr   r   lightning.pytorch.pluginsr   #lightning.pytorch.plugins.precisionr   %lightning.pytorch.strategies.strategyr   r   r#   r    r   <module>rn      s>    $ % %     &  E Y / 9 :gx gr    