
      i4                         d dl mZmZm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 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  G d de      Zy)    )AnyOptionalUnionN)Tensor)DataParallelModule)override)Accelerator)CheckpointIO)	Precision)ParallelStrategy)_StrategyRegistry)
TBroadcastTReduce)apply_to_collection)ReduceOpc                   \    e Zd ZdZ	 	 	 	 d$dee   deeej                        dee	   dee
   f fdZeedej                  fd	              Zeed%d
              Zededefd       Zededdfd       Zed&dedeej                     defd       Ze	 d'd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d(dededefd       Zed)dededefd       Z edede!eeee"f   f   f fd       Z#e	 d)d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' xZ(S )*DataParallelStrategyzImplements data-parallel training in a single process, i.e., the model gets replicated to each device and each
    gets a split of the data.Nacceleratorparallel_devicescheckpoint_io	precisionc                 .    t         |   ||d ||       y )N)r   r   cluster_environmentr   r   )super__init__)selfr   r   r   r   	__class__s        s/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/lightning/fabric/strategies/dp.pyr   zDataParallelStrategy.__init__#   s&     	#- $' 	 	
    returnc                 <    | j                   J | j                   d   S )Nr   )r   r   s    r   root_devicez DataParallelStrategy.root_device2   s'     $$000$$Q''r    c                      y N r#   s    r   distributed_sampler_kwargsz/DataParallelStrategy.distributed_sampler_kwargs8   s     r    modulec                 0    t        || j                        S )zDWraps the given model into a :class:`~torch.nn.DataParallel` module.)r)   
device_ids)r   r   r   r)   s     r   setup_modulez!DataParallelStrategy.setup_module=   s     6d6K6KLLr    c                 :    |j                  | j                         y r&   )tor$   r,   s     r   module_to_devicez%DataParallelStrategy.module_to_deviceB   s    		$""#r    batchdevicec                     |S r&   r'   )r   r1   r2   s      r   batch_to_devicez$DataParallelStrategy.batch_to_deviceF   s	     r    
collectiongroup	reduce_opc                 D    dt         dt         fd}t        |t         |      S )Ntr!   c                 t    | j                   }| j                         j                         j                  |      S r&   )dtypefloatmeanr/   )r9   original_dtypes     r   r=   z-DataParallelStrategy.all_reduce.<locals>.meanO   s)    WWN779>>#&&~66r    )r   r   )r   r5   r6   r7   r=   s        r   
all_reducezDataParallelStrategy.all_reduceK   s&    	7F 	7v 	7 #:vt<<r    argskwargsc                      y r&   r'   )r   r@   rA   s      r   barrierzDataParallelStrategy.barrierU   s    r    objsrcc                     |S r&   r'   )r   rD   rE   s      r   	broadcastzDataParallelStrategy.broadcastY   s    
r    decisionallc                     |S r&   r'   )r   rH   rI   s      r   reduce_boolean_decisionz,DataParallelStrategy.reduce_boolean_decision]   s    r    c                 Z    t        |t              r|j                  }t        |   |      S r&   )
isinstancer   r)   r   get_module_state_dict)r   r)   r   s     r   rN   z*DataParallelStrategy.get_module_state_dicta   s&    fl+]]Fw,V44r    
state_dictstrictc                 b    t        |t              r|j                  }t        |   |||       y )N)r)   rO   rP   )rM   r   r)   r   load_module_state_dict)r   r)   rO   rP   r   s       r   rR   z+DataParallelStrategy.load_module_state_dictg   s-     fl+]]F&fTZ&[r    strategy_registryc                 @    |j                  d| | j                         y )Ndp)description)register__name__)clsrS   s     r   register_strategiesz(DataParallelStrategy.register_strategieso   s     	""4#,,"Gr    )NNNN)r!   Nr&   )Nr=   )r   )T))rX   
__module____qualname____doc__r   r
   listtorchr2   r   r   r   propertyr	   r$   r(   r   r   r-   r0   r   r4   r   r   r   strr?   rC   r   intrG   boolrK   dictr   rN   rR   classmethodr   rZ   __classcell__)r   s   @r   r   r      s   !
 .29=04)-
k*
 #4#56
  -	

 I&
 (U\\ (  (    M6 Ml M M $v $$ $ $ S (5<<2H TW   lr=!=*23-=KSTYZbdgZgThKi=	= = S C D   Z c *    4 4   5F 5tCsF{AS<S7T 5 5
 X\\\*.sE#v+4F/F*G\QU\	\ \ H4E H$ H  Hr    r   )typingr   r   r   r_   r   torch.nnr   r   typing_extensionsr	   lightning.fabric.acceleratorsr
   )lightning.fabric.plugins.io.checkpoint_ior   "lightning.fabric.plugins.precisionr   $lightning.fabric.strategies.parallelr   $lightning.fabric.strategies.registryr   $lightning.fabric.strategies.strategyr   r   %lightning.fabric.utilities.apply_funcr   &lightning.fabric.utilities.distributedr   r   r'   r    r   <module>rr      sB    ( '   ) & 5 B 8 A B D E ;SH+ SHr    