a
    di8                     @   s   d dl mZ d dlZd dlm  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mZmZ d
d Zd!ddZd"ddZdd Zdd Zd#ddZd$ddZdd Zdd Z dd Z!e"d kre   dS )%    )deepcopyN)AdamW)LambdaLR)
DataLoader)Accelerator)GradientState)RegressionDatasetRegressionModel)DistributedTypeis_torch_versionset_seedc              	   C   s   t |  | D ]~\}}|js"q|s\t|j|jdu sJ d| d|j d|j dqt|j|jdu sJ d| d|j d|j dqd S )	NF7Gradients in sync when they should not be at iteration z:
model_a grad (z) == model_b grad ()T7Gradients not in sync when they should be at iteration z) != model_b grad ()zip
parametersrequires_gradtorchallclosegrad)Zmodel_aZmodel_bdid_step	iterationparamZ
grad_param r   p/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/accelerate/test_utils/scripts/test_sync.pycheck_model_parameters   s    r   Tc                 C   sJ   |    | |}t|||j}|s<||j }|  n
|| d S N)trainFZmse_losstodevicegradient_accumulation_stepsZbackward)modelinputtargetacceleratorZdo_backwardoutputZlossr   r   r   
step_model-   s    

r'   Fc           	      C   s   t d t }t|}tdd}t|dd}|| j |r|t| dd}t| dd}t	|dd	 d
}t	|dd	 d
}|r| 
||||\}}}}n| 
||\}}|r|||||||fS |||fS )z3Returns everything needed to perform basic training*   P   length   Z
batch_sizegMbP?)paramslrc                 S   s   | d S Ng?r   epochr   r   r   <lambda>C       z$get_training_setup.<locals>.<lambda>)Z	lr_lambdac                 S   s   | d S r0   r   r1   r   r   r   r3   D   r4   )r   r	   r   r   r   r   r    r   r   r   prepare)	r%   schedr"   	ddp_modelZdset
dataloaderoptddp_opt	ddp_schedr   r   r   get_training_setup8   s"    
r<   c              	   C   s@  t | \}}}tt| \}}tdD ]}| ||f\}}|| j|| j }}t||||  |d dkr| 	| t||||  W d    q1 s0    Y  nt||||  t
||d| t| | D ]:\}	}
|	jsqt|	j|
jsJ d|	j d|
j dqtd|  |tt| }q*d S )	N      r   T7Gradients not in sync when they should be:
Model grad () != DDP grad (r   9  )r<   nextitervaluesrangegatherr   r    r'   no_syncr   r   r   r   r   r   r   manual_seedrandpermlenr%   r"   r7   r8   	ddp_input
ddp_targetr   r#   r$   r   	ddp_paramr   r   r   test_noop_syncO   s*    .rO   c              	   C   sv  t | \}}}tt| \}}tdD ]D}| ||f\}}|| j|| j }}t||||  |d dkr| 	| t||||  W d    q1 s0    Y  nt||||  t
| | D ]~\}	}
|	jsq|d dkr t|	j|
jdu sNJ d|	j d|
j dqt|	j|
jdu sJ d	|	j d
|
j dqtd|  |tt| }q*d S )Nr=   r>   r   Fz7Gradients in sync when they should not be:
Model grad () == DDP grad (r   Tr?   r@   rA   )r<   rB   rC   rD   rE   rF   r   r    r'   rG   r   r   r   r   r   r   rH   rI   rJ   rK   r   r   r   test_distributed_syncq   s0    .rQ   c              
   C   s  t | |dd}t|\}}}t|D ]Z\}}| \}}	|||	f\}
}|
|j||j }
}t||
||d || t|||	| W d    n1 s0    Y  t	|
 |
 D ]\}}|jsq|d d dks|t|d kr*t|j|jdu s^J d| d|j d	|j d
qt|j|jdu sJ d| d|j d|j d
qtd|  |tt| }q$t  d S )Nr>   split_batchesdispatch_batchesr!   F   r   Tr   z:
Model grad (r@   r   r   rP   rA   )r   r<   	enumeraterD   rF   r   r    r'   
accumulater   r   r   rJ   r   r   r   rH   rI   r   _reset_state)rS   rT   r%   r"   r7   r8   r   batchrL   rM   r#   r$   r   rN   r   r   r   test_gradient_accumulation   s4    ,"rZ   c              	   C   s  t | |dd}t|d\}}}}}}}	t|D ]\}
}| \}}|||f\}}||j||j }}|  |  t||||d |	  |
d d dks|
d t
|kr| r|	  nt|jD ]}|	  q|  ||6 t|||| |	  |		  |  W d    n1 s.0    Y  |jd d |jd d ksJ d|jd d  d	|jd d  d
|
d d dkp|
d t
|k}|jdkrt||||
 td|
  q.t  d S )Nr>   rR   TFrU   r   r/   z:Learning rates found in each optimizer did not align
opt: z

DDP opt: 
rA   )r   r<   rV   rD   rF   r   r    r   r'   steprJ   rE   Znum_processesZ	zero_gradrW   Zparam_groupsr   r   rH   r   rX   )rS   rT   r%   r"   r9   r6   r8   r7   r:   r;   r   rY   rL   rM   r#   r$   _r   r   r   r   1test_gradient_accumulation_with_opt_and_scheduler   s@     

($"r^   c                  C   s  t  } tdd}t|dd}tdd}t|dd}| ||\}}| jjd u sRJ t|D ]\}}t| jjt|kszJ |t|d k r| jj	rJ |dkrt|D ]J\}}t| jjt|ksJ |t|d k r| jj	rJ q| jj	sJ qqZ| jj	sZJ qZ| jjd u sJ d S )Nr)   r*   r,   r-   `   rU   )
r   r   r   r5   Zgradient_stateZactive_dataloaderrV   idrJ   Zend_of_dataloader)r%   Z
first_dsetZfirst_dataloaderZsecond_dsetZsecond_dataloaderr   r]   Z	batch_numr   r   r   test_dataloader_break   s&    

ra   c               	   C   s\  t  } | j}|jdkrtd t  |jtjkrJ|jdkrBtd t|  |jtj	tj
fv rv|jdkrntd t|  |jtj	krdD ]:}dD ]0}|jdkrtdd| d| d	 t|| qqtd
ds|jtjkrX|jdkrtdd t  |jtj	krXdD ]P}dD ]D}|s"|s"q|jdkrFtdd| d| d	 t|| qqd S )Nr   zA**Test `accumulate` gradient accumulation with dataloader break**z'**Test NOOP `no_sync` context manager**z.**Test Distributed `no_sync` context manager**)TFz+**Test `accumulate` gradient accumulation, z`split_batches=z` and `dispatch_batches=z`**<z2.0zH**Test `accumulate` gradient accumulation with optimizer and scheduler, z1`split_batches=False`, `dispatch_batches=False`**)r   stateZlocal_process_indexprintra   Zdistributed_typer
   NOrO   Z	MULTI_GPUZ	MULTI_CPUrQ   rZ   r   r^   )r%   rc   Zsplit_batchrT   r   r   r   main  sP    




rf   c                 C   s
   t   d S r   )rf   )indexr   r   r   _mp_fn0  s    rh   __main__)T)F)FF)FF)#copyr   r   Ztorch.nn.functionalnnZ
functionalr   Ztorch.optimr   Ztorch.optim.lr_schedulerr   Ztorch.utils.datar   Zaccelerate.acceleratorr   Zaccelerate.stater   Zaccelerate.test_utilsr   r	   Zaccelerate.utilsr
   r   r   r   r'   r<   rO   rQ   rZ   r^   ra   rf   rh   __name__r   r   r   r   <module>   s*   

"(
&
+-