a
    d9E                     @   sP  d dl Z d dlZd dlmZmZ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mZmZ d dlZd dlmZm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'm(Z( d dl)m*Z* d dl+m,Z,m-Z- d dl.m/Z/ d dl0m1Z1 d dl2m3Z3 er8d dl4Z4ndZ4G dd de!Z5dS )    N)AnyCallableDictIterableListOptionalTupleUnion)apply_to_collection)Tensor)
DataLoaderSampler)CheckpointIOClusterEnvironment)get_filesystem)_IPU_AVAILABLE_POPTORCH_AVAILABLE)_LightningModuleWrapperBase)PrecisionPlugin)ParallelStrategy)
TBroadcast)_fp_to_half)RunningStage	TrainerFn)rank_zero_warn)$_get_dataloader_init_args_and_kwargs_reinstantiate_wrapped_cls)MisconfigurationException)is_overridden)STEP_OUTPUTc                       s&  e Zd ZdZdZdjed eeee ee	e
j  ee ee ee ed ed dd fd	d
Zddd fddZddd fddZeedddZeddddZeddddZeddddZdkeeeef ee ddddZdddd Zeedd!d"Zeed#d$d%Z dleee
j eed'd(d)Z!ddd*d+Z"eeee#d,d-d.Z$eee#d/d0d1Z%eeee# d/d2d3Z&eeee# d/d4d5Z'eee#d/d6d7Z(dd fd8d9Z)eed:d;d<Z*ddd=d>Z+edd?d@dAZ,dddBdCZ-dddDdEZ.dddFdGZ/dddHdIZ0dddJdKZ1dddLdMZ2dddNdOZ3dddPdQZ4eeddRdSdTZ5ee
jddUdVZ6dddWdXZ7eeddYdZZ8ee9ef eeee9ef d[d\d]Z:dmee dd^d_d`Z;dne9ee ee9dadbdcZ<doe=ee=dddedfZ>e?e@ddgdhdiZA  ZBS )pIPUStrategyz#Plugin for training on IPU devices.Zipu_strategyN   Fzpl.accelerators.Acceleratorzpoptorch.Options)acceleratordevice_iterations
autoreportautoreport_dirparallel_devicescluster_environmentcheckpoint_ioprecision_plugintraining_optsinference_optsreturnc                    s   t  j|||||d ts"td|| _|| _|| _i | _|	| _|
| _	| jrd| ji}| jrt
t| j| _| jj| jdd | j|d< t|tjd< d| _d| _dS )	a4  
        Arguments:

            device_iterations: Number of iterations to run on device at once before returning to host.
                This can be used as an optimization to speed up training.
                https://docs.graphcore.ai/projects/poptorch-user-guide/en/latest/batching.html
            autoreport: Enable auto-reporting for IPUs using PopVision
                https://docs.graphcore.ai/projects/graphcore-popvision-user-guide/en/latest/graph/graph.html
            autoreport_dir: Optional directory to store autoReport output.
            training_opts: Optional ``poptorch.Options`` to override the default created options for training.
            inference_opts: Optional ``poptorch.Options`` to override the default
                created options for validation/testing and predicting.
        )r"   r&   r'   r(   r)   z{The IPU Accelerator requires IPU devices to run. Learn more or get started with IPUs at https://www.graphcore.ai/getstartedzautoReport.allT)exist_okzautoReport.directoryZPOPLAR_ENGINE_OPTIONSN)super__init__r   r   r#   r$   r%   poptorch_models_training_opts_inference_optsr   strZ_fsmakedirsjsondumpsosenviron_update_dataloader_original_optimizer_zero_grad_original)selfr"   r#   r$   r%   r&   r'   r(   r)   r*   r+   options	__class__ i/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/strategies/ipu.pyr/   2   s4    

zIPUStrategy.__init__z
pl.Trainer)trainerr,   c                    sx  |    tjjjj| _| jtjjj_t 	| | j
d us>J | j
j| _|   t| j
| _i | _| j
jjj}|tjkr| j}| j}| j
jjd }tj| j||d}|| jtj< | j
jjrtj| j|d}|| jtj< | j
jjdkr|| jtj < n|tjkrtj| j| jd}|| jtj< nV|tj!krJtj| j| jd}|| jtj!< n*|tj"krttj| j| jd}|| jtj"< d S )Nr   )modelr<   	optimizer)rB   r<   )##_handle_gradient_accumulation_stepsplrA   
connectorsdata_connector_update_dataloaderr9   _convert_to_poptorch_loaderr.   setuplightning_moduleoptimizer_zero_gradr:   _disable_zero_gradr   rB   r0   statefnr   ZFITTINGr*   r+   
optimizerspoptorchZtrainingModelr   TRAININGZenable_validationZinferenceModel
VALIDATINGZnum_sanity_val_stepsZSANITY_CHECKINGTESTING
PREDICTING)r;   rA   Z
trainer_fnr*   r+   rC   rB   r=   r?   r@   rJ   k   s<    


zIPUStrategy.setupc                    s&   t  | t| jdkr"tdd S )Nr!   z*IPUs currently only support one optimizer.)r.   setup_optimizerslenrP   r   )r;   rA   r=   r?   r@   rV      s    zIPUStrategy.setup_optimizers)r,   c                 C   sh   | j r| js<| jr| jjS | jr(| jjS | js2J t| jS | j jjj	}|d usTJ | j| j
 d S )Nreplication_factor)rK   r0   r1   rX   r2   r&   rW   rA   rN   stage_optionsZtoDict)r;   rY   r?   r?   r@   rX      s    

zIPUStrategy.replication_factor)trainingr,   c                 C   sp   | j d usJ t }|| j || j |r<| j jjnd}|j	
| tjdrl|ttjd  |S )Nr!   ZPL_GLOBAL_SEED)rK   rQ   OptionsZdeviceIterationsr#   ZreplicationFactorrX   rA   accumulate_grad_batchesTrainingZgradientAccumulationr7   r8   getZ
randomSeedint)r;   r[   optsgradient_accumulationr?   r?   r@   _create_opts   s    zIPUStrategy._create_optsc                 C   s   | j d u r| jdd| _ | j S )NTr[   )r1   rc   r;   r?   r?   r@   r*      s    
zIPUStrategy.training_optsc                 C   s   | j d u r| jdd| _ | j S )NFrd   )r2   rc   re   r?   r?   r@   r+      s    
zIPUStrategy.inference_optszpoptorch.DataLoader)
dataloadersamplermoder,   c                 C   s`   t |tjr|S t|||| jdk\}}|tjkr8| jn| j}t	||g|R dtji|}|S )Nr!   Zexplicit_cls)

isinstancerQ   r   r   rX   r   rR   r*   r+   r   )r;   rf   rg   rh   Zdl_argsZ	dl_kwargsra   r?   r?   r@   rI      s     z'IPUStrategy._convert_to_poptorch_loaderc                 C   s@   | j dusJ | j jj}|jdgkr,td|jddi dS )zOverride the trainer.accumulation_scheduler to act as ``accumulate_grad_batches=1`` if gradient
        accumulation has been set.

        ``optimizer_step`` will be called on every batch, and the IPU will handle grad accumulation internally.
        Nr   zXIPUs currently does not support different `accumulate_grad_batches` at different epochs.r!   )rK   rA   accumulation_schedulerZepochsr   Z
schedulingupdate)r;   rj   r?   r?   r@   rD      s    
z/IPUStrategy._handle_gradient_accumulation_stepsc                 C   sB   | j d usJ | j jr| jn| j}|jj}|j}|j}|| | S N)rK   r[   r*   r+   r^   rb   r#   rX   )r;   ra   r]   r#   rX   r?   r?   r@   _n_replicate   s    zIPUStrategy._n_replicate)argsr,   c                    sH   t tddd}t td fdd}t|t|d}t|ttf|d}|S )N)xr,   c                 S   s   t | S rl   )tuplero   r?   r?   r@   to_tuple   s    z,IPUStrategy._prepare_input.<locals>.to_tuplec                    s   t | d jS Nr   )torchtensorZ	unsqueezerepeatrm   rq   re   r?   r@   	to_tensor   s    z-IPUStrategy._prepare_input.<locals>.to_tensor)Zdtypefunction)r   r   r   r
   listr`   float)r;   rn   rr   rw   r?   re   r@   _prepare_input   s
    zIPUStrategy._prepare_inputr   )batchdevicedataloader_idxr,   c                 C   s   t |tt| jjd}|S )N)rx   	precision)r
   r   r   r)   r   )r;   r|   r}   r~   r?   r?   r@   batch_to_device  s    zIPUStrategy.batch_to_devicec                 C   s:   | j }|d usJ td|r0|d us(J td d |_d S )NrL   zYou have overridden the `LightningModule.optimizer_zero_grad` hook but it will be ignored since IPUs handle the zeroing of gradients internally.)rK   r   r   rL   )r;   rK   r?   r?   r@   rM   
  s    
zIPUStrategy._disable_zero_grad)rY   rn   kwargsr,   c                 O   sR   |  |}| j| }tjj  ||i |W  d    S 1 sD0    Y  d S rl   )r{   r0   rE   coremoduleZ_jit_is_scripting)r;   rY   rn   r   Zpoptorch_modelr?   r?   r@   _step  s    

zIPUStrategy._step)rn   r   r,   c                 O   sH   | j  * | jtjg|R i |W  d    S 1 s:0    Y  d S rl   )r)   Ztrain_step_contextr   r   rR   r;   rn   r   r?   r?   r@   training_step  s    zIPUStrategy.training_stepc                 O   sH   | j  * | jtjg|R i |W  d    S 1 s:0    Y  d S rl   )r)   Zval_step_contextr   r   rS   r   r?   r?   r@   validation_step  s    zIPUStrategy.validation_stepc                 O   sH   | j  * | jtjg|R i |W  d    S 1 s:0    Y  d S rl   )r)   Ztest_step_contextr   r   rT   r   r?   r?   r@   	test_step#  s    zIPUStrategy.test_stepc                 O   sH   | j  * | jtjg|R i |W  d    S 1 s:0    Y  d S rl   )r)   Zpredict_step_contextr   r   rU   r   r?   r?   r@   predict_step'  s    zIPUStrategy.predict_stepc                    s`   | j d ur| j tjjj_| jd us&J | jd ur:| j| j_| j	
 D ]}|  qDt   d S rl   )r9   rE   rA   rF   rG   rH   rK   r:   rL   r0   valuesdestroyr.   teardownr;   rB   r=   r?   r@   r   +  s    



zIPUStrategy.teardown)rB   r,   c                 C   s
   |j d uS rl   )_executabler   r?   r?   r@   	_compiled:  s    zIPUStrategy._compiledc                 C   s2   | j  D ]"\}}| |r
| r
|  q
dS )z4Detaches all stage specific models from IPU devices.N)r0   itemsr   isAttachedToDeviceZdetachFromDevice)r;   krB   r?   r?   r@   _detach_models>  s    zIPUStrategy._detach_models)rY   r,   c                 C   s0   |    | j| }| |r,| s,|  dS )zLoads the stage specific accelerator model onto device if compiled and not attached to IPU devices.

        Args:
            stage: The stage to load
        N)r   r0   r   r   ZattachToDevice)r;   rY   rB   r?   r?   r@   _load_modelD  s    
zIPUStrategy._load_modelc                 C   s   |  tj d S rl   )r   r   rR   re   r?   r?   r@   on_train_startO  s    zIPUStrategy.on_train_startc                 C   s   |  tj d S rl   )r   r   rS   re   r?   r?   r@   on_validation_startR  s    zIPUStrategy.on_validation_startc                 C   s   |  tj d S rl   )r   r   rT   re   r?   r?   r@   on_test_startU  s    zIPUStrategy.on_test_startc                 C   s   |  tj d S rl   )r   r   rU   re   r?   r?   r@   on_predict_startX  s    zIPUStrategy.on_predict_startc                 C   s   |    d S rl   r   re   r?   r?   r@   on_train_end[  s    zIPUStrategy.on_train_endc                 C   s   |    d S rl   r   re   r?   r?   r@   on_validation_end^  s    zIPUStrategy.on_validation_endc                 C   s   |    d S rl   r   re   r?   r?   r@   on_test_enda  s    zIPUStrategy.on_test_endc                 C   s   |    d S rl   r   re   r?   r?   r@   on_predict_endd  s    zIPUStrategy.on_predict_end)r|   	batch_idxr,   c                 C   s    | j d }| jtj | d S rs   )rP   r0   r   rR   ZsetOptimizer)r;   r|   r   rC   r?   r?   r@   on_train_batch_startg  s    
z IPUStrategy.on_train_batch_startc                 C   s   d S rl   r?   re   r?   r?   r@   root_devicel  s    zIPUStrategy.root_devicec                 C   s   d S rl   r?   re   r?   r?   r@   model_to_devicep  s    zIPUStrategy.model_to_devicec                 C   s   dS )NTr?   re   r?   r?   r@   is_global_zeros  s    zIPUStrategy.is_global_zero)ru   rn   r   r,   c                 O   s   |S rl   r?   )r;   ru   rn   r   r?   r?   r@   reducew  s    zIPUStrategy.reduce)namer,   c                 C   s   d S rl   r?   )r;   r   r?   r?   r@   barrierz  s    zIPUStrategy.barrier)ru   group
sync_gradsr,   c                 C   s   |S rl   r?   )r;   ru   r   r   r?   r?   r@   
all_gather}  s    zIPUStrategy.all_gather)objsrcr,   c                 C   s   |S rl   r?   )r;   r   r   r?   r?   r@   	broadcast  s    zIPUStrategy.broadcast)strategy_registryr,   c                 C   s   |j | j| | jj d d S )N)description)registerstrategy_namer>   __name__)clsr   r?   r?   r@   register_strategies  s
    zIPUStrategy.register_strategies)
Nr!   FNNNNNNN)N)Nr   )N)NF)r   )Cr   
__module____qualname____doc__r   r   r`   boolr3   r   rt   r}   r   r   r   r/   rJ   rV   propertyrX   rc   r*   r+   r   r	   r   r   r   rI   rD   rm   r   r{   r   rM   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   classmethodr   r   __classcell__r?   r?   r=   r@   r    -   s             92 	$r    )6r5   r7   typingr   r   r   r   r   r   r   r	   rt   Z#lightning_utilities.core.apply_funcr
   r   Ztorch.utils.datar   r   Zpytorch_lightningrE   Zlightning_fabric.pluginsr   r   Z#lightning_fabric.utilities.cloud_ior   Z"pytorch_lightning.accelerators.ipur   r   Z pytorch_lightning.overrides.baser   Z#pytorch_lightning.plugins.precisionr   Z%pytorch_lightning.strategies.parallelr   Z%pytorch_lightning.strategies.strategyr   Z"pytorch_lightning.strategies.utilsr   Z pytorch_lightning.trainer.statesr   r   Zpytorch_lightning.utilitiesr   Z pytorch_lightning.utilities.datar   r   Z&pytorch_lightning.utilities.exceptionsr   Z)pytorch_lightning.utilities.model_helpersr   Z!pytorch_lightning.utilities.typesr   rQ   r    r?   r?   r?   r@   <module>   s2   (
