a
    Ad                     @   s  d dl mZ d dlZd dlmZ d dlmZ d dlZd dlmZm	Z	m
Z
 eeeeedddZG d	d
 d
e	ZG dd dZG dd dZeeeeedddZeeeeeedddZeeeeedddZG dd de	Zdeeeee ee dddZdS )    )partialN)Tensor)
checkpoint)Optional
NamedTupleList)inputdimstartlengthreturnc                 C   s0   t | ||| j| || kr |n| j| | S N)torchnarrowshape)r   r	   r
   r    r   G/var/www/html/stable-diffusion-webui/modules/sub_quadratic_attention.pynarrow_trunc   s    r   c                   @   s&   e Zd ZU eed< eed< eed< dS )	AttnChunk
exp_valuesZexp_weights_sum	max_scoreN)__name__
__module____qualname__r   __annotations__r   r   r   r   r      s   
r   c                   @   s$   e Zd ZeeeeedddZdS )SummarizeChunkquerykeyvaluer   c                 C   s   d S r   r   r   r   r   r   r   r   __call__%   s    zSummarizeChunk.__call__N)r   r   r   staticmethodr   r   r!   r   r   r   r   r   $   s   r   c                   @   s$   e Zd ZeeeeedddZdS )ComputeQueryChunkAttnr   c                 C   s   d S r   r   r    r   r   r   r!   .   s    zComputeQueryChunkAttn.__call__N)r   r   r   r"   r   r!   r   r   r   r   r#   -   s   r#   )r   r   r   scaler   c           	      C   s   t jt jddd| j| jd| |dd|dd}t j|ddd\}}| }t || }| jj	d	krrt 
||nt 
|||j|j}|d}t||jdd
|S )N   devicedtype   r   alphabetaTkeepdimmpsr	   )r   baddbmmzerosr'   r(   	transposemaxdetachexptypebmmtosqueezer   sum)	r   r   r   r$   attn_weightsr   _Zexp_weightsr   r   r   r   _summarize_chunk6   s    
4
r?   )r   r   r   summarize_chunkkv_chunk_sizer   c                    s   j \}}}j \}}}	ttdfdd  fddtd|D }
tttjt|
  }|\}}}tj|ddd\}}t	|| }|t
|d	9 }||9 }|jdd
}t
|d	jdd
}|| S )N	chunk_idxr   c                    s(   t  d| }t d| }||S Nr%   )r   )rC   Z	key_chunkZvalue_chunk)r   rA   r   r@   r   r   r   chunk_scannerU   s    z-_query_chunk_attention.<locals>.chunk_scannerc                    s   g | ]} |qS r   r   ).0chunk)rE   r   r   
<listcomp>d   s   z*_query_chunk_attention.<locals>.<listcomp>r   Tr.   r-   r1   )r   intr   r   arangemapstackzipr5   r7   	unsqueezer<   )r   r   r   r@   rA   batch_x_headsk_tokensZk_channels_per_headr>   Zv_channels_per_headchunksZ	acc_chunkZchunk_valuesZchunk_weightsZ	chunk_maxZ
global_maxZ	max_diffs
all_valuesall_weightsr   )rE   r   rA   r   r@   r   r   _query_chunk_attentionK   s    

rT   c                 C   sv   t jt jddd| j| jd| |dd|dd}|jdd}~| jjdkrVt ||nt ||	|j	|j}|S )	Nr%   r&   r)   r   r*   r-   r1   r0   )
r   r2   r3   r'   r(   r4   softmaxr8   r9   r:   )r   r   r   r$   attn_scores
attn_probsZhidden_states_slicer   r   r   $_get_attention_scores_no_kv_chunkingu   s    
4rX   c                   @   s   e Zd ZU eed< eed< dS )ScannedChunkrC   Z
attn_chunkN)r   r   r   rI   r   r   r   r   r   r   rY      s   
rY      T)r   r   r   rA   kv_chunk_size_minc                    s   j \} }|j \}	}
}	|d }t|p2tt|
|
}|durJt||}ttd fdd}tt|d}|r|tt	|n|}|
|krtt
|dntt||d} kr|||dS t}tt  D ]F}||| ||d}||dd| | |j d	  ddf< q|S )
a  Computes efficient dot-product attention given query, key, and value.
      This is efficient version of attention presented in
      https://arxiv.org/abs/2112.05682v2 which comes with O(sqrt(n)) memory requirements.
      Args:
        query: queries for calculating attention with shape of
          `[batch * num_heads, tokens, channels_per_head]`.
        key: keys for calculating attention with shape of
          `[batch * num_heads, tokens, channels_per_head]`.
        value: values to be used in attention with shape of
          `[batch * num_heads, tokens, channels_per_head]`.
        query_chunk_size: int: query chunks size
        kv_chunk_size: Optional[int]: key/value chunks size. if None: defaults to sqrt(key_tokens)
        kv_chunk_size_min: Optional[int]: key/value minimum chunk size. only considered when kv_chunk_size is None. changes `sqrt(key_tokens)` into `max(sqrt(key_tokens), kv_chunk_size_min)`, to ensure our chunk sizes don't get too small (smaller chunks = more chunks = less concurrent work done).
        use_checkpoint: bool: whether to use checkpointing (recommended True for training, False for inference)
      Returns:
        Output of shape `[batch * num_heads, query_tokens, channels_per_head]`.
      g      NrB   c                    s   t d| t S rD   )r   min)rC   q_tokensr   query_chunk_sizer   r   get_query_chunk   s    z8efficient_dot_product_attention.<locals>.get_query_chunk)r$   )rA   r@   r    r%   )r   r\   rI   mathsqrtr5   r   r   r?   r   rX   rT   r   
zeros_likerangeceil)r   r   r   r_   rA   r[   use_checkpointrO   Zq_channels_per_headr>   rP   r$   r`   r@   Zcompute_query_chunk_attnresirV   r   r]   r   efficient_dot_product_attention   sF    


.ri   )rZ   NNT)	functoolsr   r   r   torch.utils.checkpointr   ra   typingr   r   r   rI   r   r   r   r#   floatr?   rT   rX   rY   ri   r   r   r   r   <module>   sZ   		
+	    