a
    yþdZC  ã                   @   sr  d dl Z d dlZd dlZd dlZd dlmZ d dlZd dlm  m	  m
Z ddlmZ z d dlm  m	  mZ dZW n ey’   dZdZY n0 d dl
Z
d6dd„Zd	d
„ Zdddddejfdd„Zdd„ Zd7dd„Zd8dd„Zd9dd„Zdd„ Zd:dd „ZG d!d"„ d"ƒZG d#d$„ d$ƒZd%d&„ Zd;d'd(„Z d<ej!d)œd*d+„Z"d,d-„ Z#d.d/„ Z$ed=d2d3„ƒZ%d>ej!d)œd4d5„Z&dS )?é    N)Úcontextmanageré   )ÚOutOfResourcesTFc              
   C   sH   z
| ƒ }W n8 t yB } z |r,| t|ƒ¡ W Y d }~d S d }~0 0 |S ©N)r   ÚskipÚstr)ÚkernelZpytest_handleÚresÚe© r   úW/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/triton/testing.pyÚ	catch_oor   s    
r   c                 C   sš   t j|  d¡| ¡ ||f| j| jd}tt|jddŽ ƒD ]Z\}\}}}| d d …||| |d | …|| |d | …f |d d …|d d …d d …f< q:|S )Nr   ©ÚdtypeÚdeviceT©Úas_tupler   )	ÚtorchÚemptyÚsizeÚsumr   r   Ú	enumerateÚzipÚnonzero)ÚxÚmaskÚblockÚretÚidxÚhÚiÚjr   r   r   Úsparsify_tensor!   s    &"Nr"   Úcudaç{®Gáz„?g        c           	      C   sn   |d u rt j| t jd|d}|}|| | }| ¡  |¡}|rJ| ¡  ¡ }| ¡  ¡ }| ¡  ¡  ¡ }||fS )NT)r   Zrequires_gradr   )	r   ZrandnÚfloat32ZhalfÚtoÚtZrequires_grad_ÚdetachÚclone)	Úshaper   ÚalphaÚbetaZtransÚdatar   Zref_retZtri_retr   r   r   Ú	make_pair(   s    r.   c                 C   s  t d u rtdƒ‚| jd |jd  }}| jd |jd  }}||ksHJ ‚| j|jksXJ ‚| j|jkshJ ‚tj||fd|f| j| jd}t| jƒ d¡d }t  	|  
¡ | 
¡ | 
¡ ||||  d¡|  d¡| d¡| d¡| d¡| d¡|||| jjtj | j¡j¡ |S )NzCannot find cutlass libraryr   r   r   Ú.éÿÿÿÿ)Ú_cutlassÚRuntimeErrorr*   r   r   r   Zempty_stridedr   ÚsplitÚmatmulZdata_ptrZstrideÚindexr#   Zcurrent_streamZcuda_stream)ÚaÚbÚMÚNZKaZKbÚcr   r   r   r   Úcutlass_matmul5   s$    úr;   c              	   C   s`   |   ¡ }t|dkjddŽ D ]>\}}}||d d …||| |d | …|| |d | …f< q|S )Nr   Tr   r   )r)   r   r   )r   r   r   Úvaluer   r   r    r!   r   r   r   Úmask_tensorL   s    6r=   é   Ú c                 C   s‚   dd l m} t| tjƒr<| jtjkr,|  ¡ } |  ¡  	¡  
¡ } t|tjƒrl|jtjkr\| ¡ }| ¡  	¡  
¡ }|j| |||d d S )Nr   )Úerr_msgÚdecimal)Znumpy.testingZtestingÚ
isinstancer   ZTensorr   Úbfloat16ÚfloatÚcpur(   ÚnumpyZassert_array_almost_equal)r   ÚyrA   r@   Znptr   r   r   Úassert_almost_equalS   s    rH   c                 C   s¾   | j |j kr"t| j › d| j › ƒ‚| j|jkrDt| j› d|j› ƒ‚| j tjkrbt | |A ¡dkS | j tjtjtjtj	fv r€d}t
| | ƒ}t | ¡}t |¡}t |¡t ||¡ }||kS )Nz did not match with r   )r   r2   r*   r   Úboolr   Úint8Úint16Úint32Úint64ÚabsÚmax)r   rG   ZtolÚdiffZx_maxÚy_maxÚerrr   r   r   Úallclose`   s    

rS   c                 C   sL   d  | ¡} dddd|  dg}t |¡}| tjj¡ d¡}dd„ |D ƒ}|S )	Nú,ú
nvidia-smiú-iÚ0ú--query-gpu=ú--format=csv,noheader,nounitsc                 S   s   g | ]}t |ƒ‘qS r   ©Úint©Ú.0r   r   r   r   Ú
<listcomp>u   ó    znvsmi.<locals>.<listcomp>©ÚjoinÚ
subprocessÚcheck_outputÚdecodeÚsysÚstdoutÚencodingr3   ©ÚattrsÚcmdÚoutr   r   r   r   Únvsmip   s    

rl   é   éd   ©g      à?gš™™™™™É?gš™™™™™é?c                 C   s   | ƒ  t j ¡  t jjdd}t jjdd}| ¡  tdƒD ]
}	| ƒ  q<| ¡  t j ¡  | |¡d }
tdt||
 ƒƒ}tdt||
 ƒƒ}dd„ t|ƒD ƒ}dd„ t|ƒD ƒ}|rÌt j	tdƒt jd	d
}nt j	tdƒt j
d	d
}t|ƒD ]
}	| ƒ  qêt|ƒD ]F}|dur|D ]}d|_q| ¡  ||  ¡  | ƒ  ||  ¡  qþt j ¡  t  dd„ t||ƒD ƒ¡}|rŽt  |t  |¡¡ ¡ }t|ƒS t  |¡ ¡ S dS )a³  
    Benchmark the runtime of the provided function. By default, return the median runtime of :code:`fn` along with
    the 20-th and 80-th performance percentile.

    :param fn: Function to benchmark
    :type fn: Callable
    :param warmup: Warmup time (in ms)
    :type warmup: int
    :param rep: Repetition time (in ms)
    :type rep: int
    :param grad_to_none: Reset the gradient of the provided tensor to None
    :type grad_to_none: torch.tensor, optional
    :param percentiles: Performance percentile to return in addition to the median.
    :type percentiles: list[float]
    :param fast_flush: Use faster kernel to flush L2 between measurements
    :type fast_flush: bool
    T©Zenable_timingé   r   c                 S   s   g | ]}t jjd d‘qS ©Trp   ©r   r#   ÚEvent©r]   r    r   r   r   r^   Ÿ   r_   zdo_bench.<locals>.<listcomp>c                 S   s   g | ]}t jjd d‘qS rr   rs   ru   r   r   r   r^       r_   g    €„ŽAr#   r   g    €„®ANc                 S   s   g | ]\}}|  |¡‘qS r   )Úelapsed_time)r]   Úsr
   r   r   r   r^   ¸   r_   )r   r#   Zsynchronizert   ÚrecordÚrangerv   rO   r[   r   rJ   ZgradZzero_Ztensorr   ZquantileÚtolistÚtupleÚmeanÚitem)ÚfnZwarmupÚrepZgrad_to_noneZpercentilesZrecord_clocksZ
fast_flushZstart_eventZ	end_eventÚ_Zestimate_msZn_warmupZn_repeatÚcacher    r   Útimesr   r   r   Údo_benchy   sB    




rƒ   c                   @   s   e Zd ZdZddd„ZdS )Ú	Benchmarkzk
    This class is used by the :code:`perf_report` function to generate line plots with a concise API.
    r?   FNc                 C   sL   || _ || _|
| _|| _|| _|| _|| _|| _|| _|	| _	|| _
|| _dS )a  
        Constructor

        :param x_names: Name of the arguments that should appear on the x axis of the plot. If the list contains more than one element, all the arguments are assumed to have the same value.
        :type x_names: List[str]
        :param x_vals: List of values to use for the arguments in :code:`x_names`.
        :type x_vals: List[Any]
        :param line_arg: Argument name for which different values correspond to different lines in the plot.
        :type line_arg: str
        :param line_vals: List of values to use for the arguments in :code:`line_arg`.
        :type line_vals: List[str]
        :param line_names: Label names for the different lines.
        :type line_names: List[str]
        :param plot_name: Name of the plot.
        :type plot_name: str
        :param args: List of arguments to remain fixed throughout the benchmark.
        :type args: List[str]
        :param xlabel: Label for the x axis of the plot.
        :type xlabel: str, optional
        :param ylabel: Label for the y axis of the plot.
        :type ylabel: str, optional
        :param x_log: Whether the x axis should be log scale.
        :type x_log: bool, optional
        :param y_log: Whether the y axis should be log scale.
        :type y_log: bool, optional
        N)Úx_namesÚx_valsÚx_logÚline_argÚ	line_valsÚ
line_namesÚy_logÚstylesÚxlabelÚylabelÚ	plot_nameÚargs)Úselfr…   r†   rˆ   r‰   rŠ   r   r   r   rŽ   r‡   r‹   ÚcolorrŒ   r   r   r   Ú__init__Å   s    *zBenchmark.__init__)r?   r?   FFNN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r“   r   r   r   r   r„   À   s         òr„   c                   @   s&   e Zd Zdd„ Zdd„ Zd
dd„Zd	S )ÚMarkc                 C   s   || _ || _d S r   )r~   Ú
benchmarks)r‘   r~   r™   r   r   r   r“   ÿ   s    zMark.__init__c              
      sê  dd l }dd lm} dd l}|j}dd„ |jD ƒ}	dd„ |jD ƒ}
|j|jd g| |	 |
 d}|jD ]À‰ ‡ fdd„|jD ƒ}g g g   }}}|jD ]t}| j	f i |¤|j
|i¤|j¤Ž}z|\}}	}
W n" tyê   |d d   }}	}
Y n0 ||g7 }||	g7 }||
g7 }q–ˆ g| | | |jt|ƒ< qh|jrŒ| ¡  | ¡ }|jd ‰ t|jƒD ] \}}||d  ||d	   }	}
|jrŽ|j| d nd }|jr¨|j| d
 nd }|j|ˆ  || |||d |	d urV|
d urV|j|ˆ  |	|
d|d qV| ¡  |jr|jn
d |j¡}| |¡ | |j¡ | |jr@dnd¡ | |jrVdnd¡ |rj|  ¡  |rŒ| !|j" ||j› d¡¡ ||jd g|j  }|r¾t#|jd ƒ t#|ƒ |ræ|j$|j" ||j› d¡ddd d S )Nr   c                 S   s   g | ]}|› d ‘qS )ú-minr   r\   r   r   r   r^   	  r_   zMark._run.<locals>.<listcomp>c                 S   s   g | ]}|› d ‘qS )ú-maxr   r\   r   r   r   r^   
  r_   )Úcolumnsc                    s   i | ]
}|ˆ “qS r   r   )r]   Zx_name©r   r   r   Ú
<dictcomp>  r_   zMark._run.<locals>.<dictcomp>rš   r›   r   )Úlabelr’   Zlsg333333Ã?)r+   r’   z = ÚlogZlinearz.pngú:z.csvz%.1fF)Zfloat_formatr5   )%ÚosZmatplotlib.pyplotZpyplotZpandasrŠ   Z	DataFramer…   r†   r‰   r~   rˆ   r   Ú	TypeErrorÚlocÚlenr   ÚfigureZsubplotr   rŒ   ZplotZfill_betweenZlegendr   ra   Z
set_xlabelZ
set_ylabelrŽ   Z
set_xscaler‡   Z
set_yscaler‹   ÚshowZsavefigÚpathÚprintZto_csv)r‘   ÚbenchÚ	save_pathÚ
show_plotsÚ
print_datar¢   ZpltÚpdZy_meanZy_minrQ   ZdfZx_argsZrow_meanZrow_minZrow_maxrG   r   Úaxr    ÚcolZstyr   r   r   r   Ú_run  s^     

 



z	Mark._runFr?   c                 C   s„   t | jtƒ}|r| jgn| j}|r@ttj |d¡dƒ}| d¡ |D ],}|  ||||¡ |rD| d|j	› d¡ qD|r€| d¡ d S )Nzresults.htmlÚwz<html><body>
z<image src="z.png"/>
z</body></html>
)
rB   r™   r„   Úopenr¢   r¨   ra   Úwriter±   r   )r‘   r¬   r­   r«   Zhas_single_benchr™   Úhtmlrª   r   r   r   Úrun6  s    
zMark.runN)FFr?   )r”   r•   r–   r“   r±   r¶   r   r   r   r   r˜   þ   s   3r˜   c                    s   ‡ fdd„}|S )zê
    Mark a function for benchmarking. The benchmark can then be executed by using the :code:`.run` method on the return value.

    :param benchmarks: Benchmarking configurations.
    :type benchmarks: List of :class:`Benchmark`
    c                    s
   t | ˆ ƒS r   )r˜   )r~   ©r™   r   r   Ú<lambda>K  r_   zperf_report.<locals>.<lambda>r   )r™   Úwrapperr   r·   r   Úperf_reportD  s    rº   c                 C   sX   | st jjj} |stj ¡ }tjj	 
|¡d }tjj	 
|¡d }|| d d d }|S )z return DRAM bandwidth in GB/s Zmem_clock_rateZmem_bus_widthr>   g    €„.Aé   )Ú_tritonÚruntimeÚbackendÚCUDAr   r#   Úcurrent_deviceÚtritonÚcompilerÚ
cuda_utilsÚget_device_properties)r¾   r   Zmem_clock_khzZ	bus_widthZbw_gbpsr   r   r   Úget_dram_gbpsO  s    

rÅ   )r   c                 C   sÐ   |st jjj}|stj ¡ }tj 	¡  tjj
 |¡d d }|sRtjj
 |¡d }tj |¡}|d dk r~| tjksxJ ‚d}n>| tjkrŽd}n.| tjtjfv r¤d}n| tjkr´d}ntd	ƒ‚|| | d
 }|S )NZmultiprocessor_counté   Zsm_clock_rater   r»   é   i   i   údtype not supportedç•Ö&è.>)r¼   r½   r¾   r¿   r   r#   rÀ   rÁ   rÂ   Zinit_cuda_utilsrÃ   rÄ   Zget_device_capabilityÚfloat16r%   rC   rJ   r2   )r   r¾   r   Ú
clock_rateÚnum_subcoresZ
capabilityÚops_per_sub_coreÚtflopsr   r   r   Úget_max_tensorcore_tflops\  s*    




rÏ   c                     s   ‡ fdd„}|S )Nc                    s   t  ˆ ¡‡‡ fdd„ƒ}|S )Nc            
         sÞ   dd l }| t ¡ ¡ ¡ }ˆ  ¡ | ¡ k}|rÌ|dkrÌtj ˆjd ¡}tj	d ddœ}d|v shJ dƒ‚|d j
jj}|› d	ˆj› d
|› d}tjddd|gd|d}	|	jdks¸J dƒ‚dt|	jƒv sÚJ ‚nˆ| i |¤Ž d S )Nr   zcuda-memcheckÚ__file__ÚPATHÚ1)rÑ   ZPYTORCH_NO_CUDA_MEMORY_CACHINGÚrequestz@memcheck'ed test must have a (possibly unused) `request` fixturez::ú[ú]Zpytestz-vsT)Úcapture_outputÚenvz7cuda-memcheck returned an error: bounds checking failedzERROR SUMMARY: 0 errors)ÚpsutilÚProcessr¢   ÚgetppidÚnameÚitemsr¨   ÚrealpathÚ__globals__ÚenvironÚnodeZcallspecÚidr”   rb   r¶   Ú
returncoder   rf   )
r   ÚkwargsrØ   Z	ppid_nameZrun_cuda_memcheckr¨   r×   Ztest_idrj   rk   )Útarget_kwargsÚtest_fnr   r   r¹   |  s    z1cuda_memcheck.<locals>.decorator.<locals>.wrapper)Ú	functoolsÚwraps)rå   r¹   ©rä   )rå   r   Ú	decorator{  s    z cuda_memcheck.<locals>.decoratorr   )rä   ré   r   rè   r   Úcuda_memcheckz  s    rê   c                 C   sL   d  | ¡} dddd|  dg}t |¡}| tjj¡ d¡}dd„ |D ƒ}|S )	NrT   rU   rV   rW   rX   rY   c                 S   s   g | ]}t |ƒ‘qS r   rZ   r\   r   r   r   r^   œ  r_   znvsmi_attr.<locals>.<listcomp>r`   rh   r   r   r   Ú
nvsmi_attr‘  s    
û
rë   éF  é¿  c              
   c   s$  zòt  g d¢¡ t  dddd| › d| › g¡ t  dddd|› d|› g¡ tdgƒd	 }td
gƒd	 }t||  ƒdk sˆJ d| › dƒ‚t|| ƒdk s¨J d|› dƒ‚d|  }d| d }||fV  W t  g d¢¡ t  g d¢¡ t  g d¢¡ n,t  g d¢¡ t  g d¢¡ t  g d¢¡ 0 d S )N)rU   rV   rW   ú-pmrÒ   rU   rV   rW   z--lock-gpu-clocks=rT   z--lock-memory-clocks=zclocks.current.smr   zclocks.current.memoryé
   zGPU SMs must run at z MHzgÞ 3ßÁOÌ?i   gü©ñÒMbP?)rU   rV   rW   rî   rW   )rU   rV   rW   z-rgc)rU   rV   rW   z-rmc)rb   rc   rë   rN   )Zref_sm_clockZref_mem_clockZcur_sm_clockZcur_mem_clockrÎ   Zgbpsr   r   r   Úset_gpu_clock   s:    üÿüÿ  þrð   c                 C   s¾   |st jjj}|stj ¡ }t j ||¡d }t j ||¡}t j 	||¡}|dk r|| tj
krbd}qª| tjkrrd}qªtdƒ‚n.| tj
krŒd}n| tjtjfv r¢d}ntdƒ‚|| | d }|S )NrÆ   éP   é    é@   rÈ   rÉ   )r¼   r½   r¾   r¿   r   r#   rÀ   Znum_smrË   Úccr%   rÊ   r2   rC   )r   r¾   r   rÌ   rË   rô   rÍ   rÎ   r   r   r   Úget_max_simd_tflopsÁ  s(    





rõ   )N)r   )r>   r?   )r$   )rm   rn   Nro   FF)NN)NNN)rì   rí   )NN)'ræ   r¢   rb   re   Ú
contextlibr   r   Ztriton._C.libtriton.tritonZ_CZ	libtritonrÁ   r¼   rÂ   r   Ztriton._C.libtriton.cutlassZcutlassr1   Zhas_cutlassÚImportErrorr   r"   r%   r.   r;   r=   rH   rS   rl   rƒ   r„   r˜   rº   rÅ   r   rÏ   rê   rë   rð   rõ   r   r   r   r   Ú<module>   sH   





	   þ
G>F
 