a
    d[                     @   sT  d dl Z d dlZd dlZd dlZd dlmZ d dl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mZ d d	lmZmZmZmZmZmZmZmZ e rd d
lmZ nd dlmZ dd Zdd Z dd Z!dd Z"dd Z#dd Z$dd Z%dd Z&dd Z'dd Z(d d! Z)d"d# Z*d$d% Z+d&d' Z,d(d) Z-e.d*krPe-  dS )+    N)deepcopy)Path)
DataLoader)Accelerator)prepare_data_loader)AcceleratorState)RegressionDatasetare_the_same_tensors)DistributedTypegatheris_bf16_availableis_ipex_availableis_npu_availableis_xpu_availableset_seedsynchronize_rng_states)RegressionModel4XPU)RegressionModelc                 C   s   t d| j  d S )NzPrinting from the main process printprocess_indexstate r   r/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/accelerate/test_utils/scripts/test_script.py
print_main2   s    r   c                 C   s   t d| j  d S )Nz%Printing from the local main process )r   local_process_indexr   r   r   r   print_local_main6   s    r   c                 C   s   t d| j  d S )NPrinting from the last process r   r   r   r   r   
print_last:   s    r   c                 C   s   t d| d| j  d S )NPrinting from process : r   )r   process_idxr   r   r   print_on>   s    r#   c               	   C   s  t  } | j}td}|   | jrdtd t|d}|d W d    q1 sX0    Y  n4t|d}|d W d    n1 s0    Y  W d    n1 s0    Y  | 	  | jrt|d}d
| }W d    n1 s0    Y  zl|dsJ d|d	kr2|ds2J d|d| jd	 kshJ d
|dd	  d| j W n ty   |   Y n0 | jr| r|  | 	  t }t|  | t| j W d    n1 s0    Y  |  }| jr|dks>J | dn |  dks>J | d|d |d t|  | t| j W d    n1 s0    Y  | jr|  dksJ n|  dksJ |d |d t|  |  t!| j W d    n1 s
0    Y  | j"rB|  d| jjd	  ksXJ n|  dksXJ |d |d t#|D ]}t|& | j$t%|d| j| W d    n1 s0    Y  | j&|kr|  d| d| j& ksJ n|  dksJ |d |d qtd S )Nzcheck_main_process_first.txt皙?za+zCurrently in the main process
zNow on another process
r zMain process was not first   zOnly wrote to file zNow on another processz times, not z Printing from the main process 0z$ != Printing from the main process 0z != ""r   z&Printing from the local main process 0r   )r   r    r!   )'r   num_processesr   Zmain_process_firstZis_main_processtimesleepopenwritewait_for_everyonejoin	readlines
startswithendswithcountAssertionErrorunlinkexistsioStringIO
contextlibredirect_stdoutZon_main_processr   r   getvaluerstriptruncateseekZon_local_main_processr   Zis_local_main_processZon_last_processr   is_last_processrangeZ
on_processr#   r   )acceleratorr(   pathftextresultr"   r   r   r   process_execution_checkB   sv    

*F,
0 

0

0&

6&
rE   c                  C   s$   t  } | jdkrtd t|  d S )Nr   zTesting, testing. 1, 2, 3.)r   r   r   r   r   r   r   init_state_check   s    
rF   c                  C   s   t  } tdg tt s$J d| jtjkrRtdg ttj s~J dn,| jtj	kr~tdg ttj
 s~J dt }tdg|d t| sJ d	| jd
krtd d S )Ntorchz*RNG states improperly synchronized on CPU.cudaz*RNG states improperly synchronized on GPU.xpuz*RNG states improperly synchronized on XPU.	generator)rJ   z0RNG states improperly synchronized in generator.r   zAll rng are properly synched.)r   r   r	   rG   Zget_rng_statedistributed_typer
   	MULTI_GPUrH   Z	MULTI_XPUrI   	GeneratorZ	get_stater   r   )r   rJ   r   r   r   rng_sync_check   s    



rN   c                  C   s(  t  } d| j }tt|dd}t|| j| j| jdd}g }|D ]}|t| q@t	
|}t| j|t| t	| t	d| sJ dtt|dd}t|| j| j| jddd}g }|D ]}|t| qt	
|}t	| t	d| s
J d| jdkrtd	 tt|ddd
}t|| j| j| jdd}g }|D ]}|t| qPt	
| }|  |tt|ksJ dtt|ddd
}t|| j| j| jddd}g }|D ]}|t| qt	
| }|  |tt|ksJ d| jdkr$td d S )N       
batch_sizeT)put_on_devicer   %Wrong non-shuffled dataloader result.)rS   split_batchesz Non-shuffled dataloader passing.rR   shuffle!Wrong shuffled dataloader result.zShuffled dataloader passing.)r   r(   r   r?   r   devicer   appendr   rG   catr   typeequalcpuarangelongtolistsortlistr   r   lengthdlrD   batchr   r   r   dl_preparation_check   sf    

$
&rh   c               	   C   s  t  } d| j }tt|dd}t|| j| j| jddd}g }|D ]}|t| qBt	
|}t	| t	d| sJ dtt|dd}t|| j| j| jdddd}g }|D ]}|t| qt	
|}t	| t	d| sJ d| jdkrtd	 tt|ddd
}t|| j| j| jddd}g }|D ]}|t| qBt	
| }|  |tt|ksJ dtt|ddd
}t|| j| j| jdddd}g }|D ]}|t| qt	
| }|  |tt|ksJ d| jdkrtd d S )NrO   rP   rQ   T)rS   dispatch_batchesr   rT   )rS   rU   ri   z(Non-shuffled central dataloader passing.rV   rX   z$Shuffled central dataloader passing.)r   r(   r   r?   r   rY   r   rZ   r   rG   r[   r]   r^   r_   r`   r   ra   rb   rc   r   rd   r   r   r   central_dl_preparation_check   sp    

$	
$	rj   c                 C   s   t d |d t| dd}t||d|d}t }tjj| dd}t	dD ]F}|D ]<}|
  ||d }	tjj|	|d	 }
|
  |  qXqP||fS )
N*   )re   seedTrR   rW   rJ   r$   lr   xy)r   manual_seedr   r   r   rG   optimSGD
parametersr?   	zero_gradnn
functionalmse_lossbackwardstep)re   rR   rJ   	train_settrain_dlmodel	optimizerepochrg   outputlossr   r   r   mock_training#  s    
r   c                  C   s  t  } t }d}|d | j }t||| j |\}}t|jsHJ dt|jsZJ dt }t	||d|d}t
 }tjj| dd}	||||	\}}}	td |d td	D ]H}
|D ]>}|  ||d
 }tjj||d }|| |	  qq|| }t|j|js*J dt|j|jsDJ d|d tdd}t	||| j d|d}t
 }tjj| dd}	||||	\}}}	td |d td	D ]L}|D ]@}|  ||d
 }tjj||d }|| |	  qq|| }t|j|js,J dt|j|jsFJ d|d tj sdt rftd t   tdd}t	||d|d}t
 }tjj| dd}	||||	\}}}	td |d td	D ]L}|D ]@}|  ||d
 }tjj||d }|| |	  qސq|| }t|j|jsLJ dt|j|jsfJ dtj rtd t   tdd}tj dd}||}|j|dd}t!ddgj"tj#|j$d}||}t% rtd t   tdd}t	||d|d}t
 }tjj| dd}	||||	\}}}	td |d td	D ]L}|D ]@}|  ||d
 }tjj||d }|| |	  qXqP|| }t|j|jsJ dt|j|jsJ dt& rtd t   tddd}t	||d|d}t
 }tjj| dd}	||||	\}}}	td |d td	D ]L}|D ]@}|  ||d
 }tjj||d }|| |	  qdq\|| }t|j|jsJ dt|j|jsJ dt' rtd t   tddd}t	||d|d}t
 }tjj| dd}	||||	\}}}	td |d td	D ]L}|D ]@}|  ||d
 }tjj||d }|| |	  qpqh|| }t|j|jsJ dt|j|jsJ dd S )NrP      z0Did not obtain the same model on both processes.Trm   r$   rn   rk   rp   rq   rr   z=Did not obtain the same model on CPU or distributed training.zVTraining yielded the same results on one CPU or distributed setup with no batch split.)rU   zSTraining yielded the same results on one CPU or distributes setup with batch split.zFP16 training check.Zfp16)mixed_precisionzKeep fp32 wrapper check.   )Zkeep_fp32_wrapperr'   )ZdtyperY   zBF16 training check.Zbf16zipex BF16 training check.)r   r^   zxpu BF16 training check.Fz=Did not obtain the same model on XPU or distributed training.)(r   rG   rM   r(   r   r	   abr   r   r   rt   ru   rv   preparer   rs   r?   rw   rx   ry   rz   r{   r|   Zunwrap_modelr^   allcloser   rH   Zis_availabler   Z_reset_stateZLinearZTensortofloat16rY   r   r   r   )r   rJ   rR   re   r}   Z	old_modelr@   r~   r   r   r   rg   r   r   _Zmodel_with_fp32_wrapperZinput_tensorr   r   r   training_check4  s   


















r   c                  C   s   t  } ttdd| j }| |6}t|dksLJ d| j dt| W d    n1 s`0    Y  ttdd| j d }| j|ddP}| jrt	t|| j }t||ksJ d| j dt| W d    n1 s0    Y  | 
  d S )	Nr   r   z4Each process did not have two items. Process index: z
; Length: rp   T)Zapply_paddingz;Last process did not get the extra item(s). Process index: )r   rc   r?   r(   split_between_processeslenr   r>   mathceilr-   )r   dataresultsZnum_samples_per_devicer   r   r   !test_split_between_processes_list  s     
"
"r   c                  C   s  t  } | jdv rxg dg dtg dd}t|}| |"}| jdkrt|d |d d d| j  ksJ nj| jd	kr|d |d d	d  ksJ nB| jd
kr|d |d dd  ksJ d|d d  d|d  | jdkr|d |d d d| j  ksfJ nV| jd	kr<|d |d d	d  ksfJ n*| jd
krf|d |d dd  ksfJ | jdkrt|d |d d d| j  sXJ d|d d d| j   d|d  n| jd	krt|d |d d	d  sXJ d|d d	d   d|d  nF| jd
krXt|d |d d
 sXJ d|d d
  d|d  W d    n1 sn0    Y  |   d S )N)r'   r   r   )r'   r   rp   r   )wrq   rr   zr   r'   r   rp   )r   r   cr   r   r   r   rp   z
Expected: z
, Actual: r   r   z7Did not obtain expected values on process 0, expected `z`, received: z7Did not obtain expected values on process 2, expected `z7Did not obtain expected values on process 4, expected `)	r   r(   rG   tensorr   r   r   r   r-   )r   r   Z	data_copyr   r   r   r   (test_split_between_processes_nested_dict  sH    
$

8& $$r   c                  C   s   t  } | jdkrtg dg dg| j}| |`}| jdkrht|tg d| jsJ n"t|tg d| jsJ W d    n1 s0    Y  | 	  d S )Nr'   r   )r            r   )
r   r(   rG   r   r   rY   r   r   r   r-   )r   r   r   r   r   r   #test_split_between_processes_tensor  s    

$@r   c                  C   s&  t  } | j}|jdkrtd t  |  |jtjkrDt	j
 }n|j}||jkr|jdkrftd t  |jdkr~td t  |jdkrtd t  |jdkrtd t  |jdkrtd t  |jdkrtd t  |jtjkrt  |jtjkrd S |jdkrtd	 t  d S )
Nr   z**Initialization**z
**Test process execution**z+
**Test split between processes as a list**z+
**Test split between processes as a dict**z-
**Test split between processes as a tensor**z1
**Test random number generator synchronization**z 
**DataLoader integration test**z
**Training integration test**)r   r   r   r   rF   r-   rK   r
   rL   rG   rH   Zdevice_countr(   r   rE   r   r   r   rN   rh   ZTPUrj   Z	DEEPSPEEDr   )r@   r   Znum_processes_per_noder   r   r   main  sF    







r   __main__)/r8   r6   r   r)   copyr   pathlibr   rG   Ztorch.utils.datar   Z
accelerater   Zaccelerate.data_loaderr   Zaccelerate.stater   Zaccelerate.test_utilsr   r	   Zaccelerate.utilsr
   r   r   r   r   r   r   r   r   r   r   r   r   r#   rE   rF   rN   rh   rj   r   r   r   r   r   r   __name__r   r   r   r   <module>   sB   (I=A (#2
