a
    zd                     @   s   d dl Z d dlZd dlmZ dd Zdd Zeddd iedd	d iejej	d
ddZ
eddd ieddd iejej	d
ddZG dd de jjZejZdS )    Nc                 C   sP   | d8 } | | d? O } | | d? O } | | d? O } | | d? O } | | d? O } | d7 } | S )N                )nr   r   a/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/triton/ops/cross_entropy.pynext_power_of_2   s    r
   c                 C   s   | dk rdS | dk rdS dS )Ni   r   i    r   r   r   )Nr   r   r	   	num_warps   s
    r   c                 C   s   t | d S Nr   r   nargsr   r   r	   <lambda>       r   BLOCKc                 C   s   t | d S r   r
   r   r   r   r	   r      r   )r   c                 C   s   t d}t d|}t || }| ||  | } |||  | }	|||  | }
t j| ||k td d}|t j}|t |d }t t 	t 
|d| }t j|	|||k d t   t |
}t || | d S Nr   inf)maskother)r   )tl
program_idarangeloadfloattofloat32maxlogsumexpstoreZdebug_barrier)ZLOGITSPROBSIDXZLOSSr   r   rowcolsidxZ
WRIT_PROBSZ
READ_PROBSlogitsprobsr   r   r	   _forward   s    

r,   c                 C   s   t | d S r   r   r   r   r   r	   r   3   r   c                 C   s   t | d S r   r   r   r   r   r	   r   4   r   c                 C   s   t d}t d|}t || }| ||  | } t j| ||k tdd }t |t j}||k}	t || }
||	 |
 }t j| || j	j
||k d d S r   )r   r   r   r   r   r#   r   r   r$   dtypeZ
element_ty)r%   r&   ZDPROBSr   r   r'   r(   r)   r+   deltaZdoutZdinr   r   r	   	_backward3   s    
r/   c                   @   s$   e Zd Zedd Zedd ZdS )_cross_entropyc           	         s~   |j tjksJ d j j  }} jd tj|||d}tj ||d} fdd}t|  ||| ||| |S )Nz(Indices are expected to be of type long.)r-   devicec                    s       fS NZnumeloptr*   n_colsr   r	   r   R   r   z(_cross_entropy.forward.<locals>.<lambda>)r-   torchint64r2   shapeZ
empty_liker,   Zsave_for_backward)	clsctxr*   indicesr2   r-   resultneg_logprobsgridr   r7   r	   forwardH   s    
z_cross_entropy.forwardc                    s<   |j \}jd   fdd}t| ||  dfS )a  We know d(-log(p[i])/dlogit[k] = -id_mat[i,k] + p[k]
        so we initialize the gradient as neg_logprobs, so we can just exponentiate
        to get p[k], which is most of what we need...  neg_logprobs will be
        modified in place to become the gradient we want
        r1   c                    s       fS r3   r4   r5   r8   r@   r   r	   r   d   r   z)_cross_entropy.backward.<locals>.<lambda>N)Zsaved_tensorsr;   r/   )r<   r=   Zdneg_logprobsr>   rA   r   rC   r	   backwardX   s
    

z_cross_entropy.backwardN)__name__
__module____qualname__classmethodrB   rD   r   r   r   r	   r0   G   s   
r0   )r9   ZtritonZtriton.languagelanguager   r
   r   Z
heuristicsZjitZ	constexprr,   r/   ZautogradZFunctionr0   applyZcross_entropyr   r   r   r	   <module>   s   "