a
    d                     @   s*  d dl Z d dl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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" d dl#m$Z$ d dl%m&Z& d dl'm(Z( d dl)m*Z* erd dl+m
  m,Z- d dl.Z/e 0e1Z2G dd de$Z3dS )    N)AnyCallableDictListOptionalUnion)Module)	Optimizer)CheckpointIOClusterEnvironmentgroup)_HPU_AVAILABLE)LightningDistributedModule)broadcast_object_list)HPUCheckpointIO)_WrappingCheckpointIO)PrecisionPlugin)DDPStrategy)MisconfigurationException)_TORCH_LESSER_EQUAL_1_10_2)STEP_OUTPUTc                       sv  e Zd ZdZdZd*ed eeej  ee	 ee
 ee ee ee ee ee ee edd fddZee
d	d
dZejee
 ddddZdd	 fddZdd	ddZdd	ddZdd	 fddZd+eeedddZdd	ddZd,eeeg ef eedef  eed fddZeed d!d"Z eed d#d$Z!e"e#dd%d&d'Z$dd	 fd(d)Z%  Z&S )-HPUParallelStrategyz:Strategy for distributed training on multiple HPU devices.Zhpu_parallelNhcclzpl.accelerators.Accelerator)acceleratorparallel_devicescluster_environmentcheckpoint_ioprecision_pluginddp_comm_stateddp_comm_hookddp_comm_wrappermodel_averaging_periodprocess_group_backendkwargsreturnc                    s8   t stdt jf |||||||||	|
d
| d S )Nz1`HPUParallelStrategy` requires HPU devices to run)
r   r   r   r   r   r   r    r!   r"   r#   )r   r   super__init__)selfr   r   r   r   r   r   r    r!   r"   r#   r$   	__class__ r/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/strategies/hpu_parallel.pyr'   0   s     zHPUParallelStrategy.__init__)r%   c                 C   s0   | j d u rt | _ nt| j tr*t | j _| j S N)_checkpoint_ior   
isinstancer   r   r(   r+   r+   r,   r   P   s
    


z!HPUParallelStrategy.checkpoint_io)ior%   c                 C   s
   || _ d S r-   )r.   )r(   r1   r+   r+   r,   r   Y   s    c                    s6   t | jtjd< | jdkr(t dtjd< t   d S )NIDr      HCCL_DISTRIBUTED_BACKEND)strZ
local_rankosenvironZ_process_group_backendr&   setup_environmentr0   r)   r+   r,   r8   ]   s    
z%HPUParallelStrategy.setup_environmentc                 C   s   d S r-   r+   r0   r+   r+   r,   determine_ddp_device_idse   s    z,HPUParallelStrategy.determine_ddp_device_idsc                 C   sN   | j dd| j d< d| _| j d}|r:d| j d< d| _|d urJ| j d= d S )NZfind_unused_parametersTFstatic_graph)Z_ddp_kwargsget_static_graph)r(   r:   r+   r+   r,   _pre_configure_ddph   s    
z&HPUParallelStrategy._pre_configure_ddpc                    sf   t rXt| jj d |   | t| j| _| j	j
dkrN| jrN| j  |   n
t   d S )Nz%: configuring DistributedDataParallelZhpu)r   logZdetailr*   __name__r=   Z_setup_modelr   modelZroot_devicetyper<   _modelZ_set_static_graphZ_register_ddp_hooksr&   configure_ddpr0   r)   r+   r,   rC   z   s    

z!HPUParallelStrategy.configure_ddpr   )objsrcr%   c                 C   s.   |g}| j |krd g}t||tjd |d S )Nr   r   )Zglobal_rankr   _groupZWORLD)r(   rD   rE   r+   r+   r,   	broadcast   s
    
zHPUParallelStrategy.broadcastc                 C   s   t   d S r-   htcore	mark_stepr0   r+   r+   r,   on_after_backward   s    z%HPUParallelStrategy.on_after_backwardzpl.LightningModule)	optimizeropt_idxclosurer@   r$   r%   c                    s&   t  j||||fi |}t  |S r-   )r&   optimizer_steprI   rJ   )r(   rL   rM   rN   r@   r$   Zoptimizer_outputr)   r+   r,   rO      s    z"HPUParallelStrategy.optimizer_step)step_outputr%   c                 C   s   t   |S r-   rH   r(   rP   r+   r+   r,   validation_step_end   s    z'HPUParallelStrategy.validation_step_endc                 C   s   t   |S r-   rH   rQ   r+   r+   r,   test_step_end   s    z!HPUParallelStrategy.test_step_end)strategy_registryr%   c                 C   s   |j | j| | jj d d S )N)description)registerstrategy_namer*   r?   )clsrT   r+   r+   r,   register_strategies   s
    z'HPUParallelStrategy.register_strategiesc                    s*   t    tjdd  tjdd  d S )Nr2   r4   )r&   teardownr6   r7   popr0   r)   r+   r,   rZ      s    
zHPUParallelStrategy.teardown)
NNNNNNNNNr   )r   )N)'r?   
__module____qualname____doc__rW   r   r   torchZdevicer   r
   r   objectr   intr5   r   r'   propertyr   setterr8   r9   r=   rC   rG   rK   r	   r   r   rO   r   rR   rS   classmethodr   rY   rZ   __classcell__r+   r+   r)   r,   r   +   sd              	 
r   )4loggingr6   typingr   r   r   r   r   r   Ztorch.distributedr_   Ztorch.nnr   Ztorch.optim.optimizerr	   Zpytorch_lightningplZlightning_fabric.pluginsr
   r   Z&lightning_fabric.utilities.distributedr   rF   Z"pytorch_lightning.accelerators.hpur   Zpytorch_lightning.overridesr   Z-pytorch_lightning.overrides.torch_distributedr   Z'pytorch_lightning.plugins.io.hpu_pluginr   Z$pytorch_lightning.plugins.io.wrapperr   Z#pytorch_lightning.plugins.precisionr   Z pytorch_lightning.strategies.ddpr   Z&pytorch_lightning.utilities.exceptionsr   Z#pytorch_lightning.utilities.importsr   Z!pytorch_lightning.utilities.typesr   Zhabana_frameworks.torch.corecorerI   Z(habana_frameworks.torch.distributed.hcclZhabana_frameworks	getLoggerr?   r>   r   r+   r+   r+   r,   <module>   s.    
