a
    d                     @   sd   d dl Z d dlmZmZ de jedddZdd Zde jeeeeee	e j
eeef d		d
dZdS )    N)TupleCallablexmodec                 C   s   | S )N r   r   r   U/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/tomesd/merge.py
do_nothing   s    r	   c                 C   sP   | j d dkr>t| d|dk r*|d n||ddS t| ||S d S )N   r   )shapetorchgatherZ	unsqueezeZsqueeze)inputdimindexr   r   r   mps_gather_workaround	   s    r   F)	metricwhsxsyrno_rand	generatorreturnc                    s  | j \ }dkrttfS | jjdkr,tntjt & || ||  }	}
|rntj|	|
d| jtj	d}n&tj
|| |	|
df|j|d| j}tj|	|
|| | jtj	d}|jd|tj||jd d ||	|
||dd|	| |
| }|	| |k s|
| |k rFtj||| jtj	d}||d	|	| d	|
| f< n|}|dd
djdd}~~|	|
 |d	d	d	d	d	f |d	d	d	d	d	f  fdd| | jd
dd } | \}}||d
d }t|j d |jd
d\}}|jd
ddd }|dd	d	d	f 
|dd	d	d	f 	|d d	dW d	   n1 sj0    Y  dtjtjd	
fdd}tjtjd 	
f
dd}||fS )ab  
    Partitions the tokens into src and dst and merges r tokens from src to dst.
    Dst tokens are partitioned by choosing one randomy in each (sx, sy) region.

    Args:
     - metric [B, N, C]: metric to use for similarity
     - w: image width in tokens
     - h: image height in tokens
     - sx: stride in the x dimension for dst, must divide w
     - sy: stride in the y dimension for dst, must divide h
     - r: number of tokens to remove (by merging)
     - no_rand: if true, disable randomness (use top left corner only)
     - rand_seed: if no_rand is false, and if not None, sets random seed.
    r   Zmpsr   devicedtype)sizer   r      )r   r   r   srcNr
   r   c                    sF   | j d }| d  |d}| d |d}||fS )Nr
   r   r   r   )r   expand)r   Cr"   dst)BNa_idxb_idxr   num_dstr   r   splitM   s    
z/bipartite_soft_matching_random2d.<locals>.splitT)r   Zkeepdim)r   Z
descending).N.r$   mean)r   r   c                    sz   | \}}|j \}}}|d|| |d}|d||d}|jd ||||d}tj||gddS )Nr.   r$   )reducer   r#   )r   r%   Zscatter_reducer   cat)r   r   r"   r'   nt1cunm)dst_idxr   r   r-   src_idxunm_idxr   r   mergec   s    z/bipartite_soft_matching_random2d.<locals>.mergec                    s   	j d }| dd |d d f | d|d d d f  }}|j \}}}|d |d}tj || j| jd}|jd ||d |jd j d dd	d |||d |jd j d ddd ||d |S )Nr   .r.   r$   r   r!   )r   r%   r   zerosr   r   scatter_)r   Zunm_lenr5   r'   _r4   r"   out)
r(   r)   r*   r+   r6   r   r,   r   r7   r8   r   r   unmergem   s    
.44z1bipartite_soft_matching_random2d.<locals>.unmerge)r/   )r   r	   r   typer   r   r   Zno_gradr:   int64randinttor;   Z	ones_liker   viewZ	transposeZreshapeZargsortZnormminmaxTensor)r   r   r   r   r   r   r   r   r<   ZhsyZwsxZrand_idxZidx_buffer_viewZ
idx_bufferabZscoresZnode_maxZnode_idxZedge_idxr9   r>   r   )r(   r)   r*   r+   r6   r   r,   r   r-   r7   r8   r    bipartite_soft_matching_random2d   sD    &(2$
*rI   )N)FN)r   typingr   r   rF   strr	   r   intbool	GeneratorrI   r   r   r   r   <module>   s     
