
      i                         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 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  G d dee      Zy)    )ABC)AnyOptionalN)Tensor)override)Accelerator)ClusterEnvironment)CheckpointIO)	Precision)Strategy_all_gather_ddp_if_available)ReduceOpc                       e Zd ZdZ	 	 	 	 	 ddee   deeej                        dee	   dee
   dee   f
 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j.                  deeej                        d	dfd       Zed	e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 fd       Z xZ S )ParallelStrategyz:Strategy for training with multiple processes in parallel.Nacceleratorparallel_devicescluster_environmentcheckpoint_io	precisionc                 F    t         |   |||       || _        || _        y )N)r   r   r   )super__init__r   r   )selfr   r   r   r   r   	__class__s         y/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/lightning/fabric/strategies/parallel.pyr   zParallelStrategy.__init__!   s*     	[Ybc 0AT     returnc                 R    | j                   | j                   j                         S dS Nr   )r   global_rankr   s    r   r!   zParallelStrategy.global_rank-   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_rank1   (    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_rank5   s(    7;7O7O7[t''113babbr   c                 R    | j                   | j                   j                         S dS )N   )r   
world_sizer"   s    r   r*   zParallelStrategy.world_size9   r%   r   c                      | j                   dk(  S r    )r!   r"   s    r   is_global_zerozParallelStrategy.is_global_zero=   s     1$$r   c                     | j                   S N_parallel_devicesr"   s    r   r   z!ParallelStrategy.parallel_devicesB   s    %%%r   c                     || _         y r.   r/   )r   r   s     r   r   z!ParallelStrategy.parallel_devicesF   s
    !1r   c                 4    | j                   | j                  dS )zArguments for the ``DistributedSampler``.

        If this method is not defined, or it returns ``None``, then the ``DistributedSampler`` will not be used.

        )num_replicasrank)r*   r!   r"   s    r   distributed_sampler_kwargsz+ParallelStrategy.distributed_sampler_kwargsJ   s     !%9I9IJJr   tensorgroup
sync_gradsc                     t        |||      S )z&Perform a all_gather on all processes.)r7   r8   r   )r   r6   r7   r8   s       r   
all_gatherzParallelStrategy.all_gatherS   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)	torchr6   introot_device
all_reducer   SUMboolr*   )r   r;   r<   s      r   reduce_boolean_decisionz(ParallelStrategy.reduce_boolean_decisionX   si     <<Hd6F6FG??ll # 
 9<4DOO34 BFhr   c                 p    | j                   J | j                   j                          t        |          S r.   )r   teardownr   )r   r   s    r   rH   zParallelStrategy.teardownn   s5    ''333  ))+w!!r   )NNNNN)NF)T)r   N)!__name__
__module____qualname____doc__r   r   listr@   r>   r	   r
   r   r   propertyrA   r!   r$   r'   r*   r   rE   r,   r   setterdictstrr   r5   r   r:   rF   rH   __classcell__)r   s   @r   r   r      s   D .29=<@04)-
Uk*
U #4#56
U &&89	
U
  -
U I&
U eS e e dC d d c3 c c dC d d % %  % &(4+="> & & 2$u||:L1M 2RV 2 2 KHT#s(^,D K K X X XRV Xci X X  4 4  * " "r   r   )abcr   typingr   r   r@   r   typing_extensionsr   )lightning.fabric.accelerators.acceleratorr   9lightning.fabric.plugins.environments.cluster_environmentr	   )lightning.fabric.plugins.io.checkpoint_ior
   "lightning.fabric.plugins.precisionr   $lightning.fabric.strategies.strategyr   &lightning.fabric.utilities.distributedr    lightning.fabric.utilities.typesr   r    r   r   <module>r^      s;         & A X B 8 9 O 5T"x T"r   