a
    AþdO  ã                   @   s>   d dl Z d dlmZ d dlZd dlmZ G dd„ dejƒZdS )é    N)Ú	rearrangec                       sD   e Zd ZdZd‡ fdd„	Zdd	„ Zd
d„ Zddd„Zdd„ Z‡  Z	S )ÚVectorQuantizer2z´
    Improved version over VectorQuantizer, can be used as a drop-in replacement. Mostly
    avoids costly matrix multiplications and allows for post-hoc remapping of indices.
    NÚrandomFTc                    sâ   t ƒ  ¡  || _|| _|| _|| _t | j| j¡| _| jj	j
 d| j d| j ¡ || _| jd urÒ|  dt t | j¡¡¡ | jjd | _|| _| jdkr®| j| _| jd | _td| j› d| j› d	| j› d
ƒ n|| _|| _d S )Ng      ð¿ç      ð?Úusedr   Úextraé   z
Remapping z indices to z indices. Using z for unknown indices.)ÚsuperÚ__init__Ún_eÚe_dimÚbetaÚlegacyÚnnÚ	EmbeddingÚ	embeddingÚweightÚdataÚuniform_ÚremapÚregister_bufferÚtorchÚtensorÚnpÚloadr   ÚshapeÚre_embedÚunknown_indexÚprintÚsane_index_shape)Úselfr   r   r   r   r   r   r   ©Ú	__class__© úN/var/www/html/stable-diffusion-webui/extensions-builtin/LDSR/vqvae_quantize.pyr
   '   s(    


ÿzVectorQuantizer2.__init__c                 C   s²   |j }t|ƒdksJ ‚| |d d¡}| j |¡}|d d …d d …d f |d k ¡ }| d¡}| d¡dk }| jdkržt	j
d| j|| j dj|jd||< n
| j||< | |¡S )	Nr   r   éÿÿÿÿ)NN.é   r   )Úsize)Údevice)r   ÚlenÚreshaper   ÚtoÚlongÚargmaxÚsumr   r   Úrandintr   r(   )r    ÚindsÚishaper   ÚmatchÚnewÚunknownr#   r#   r$   Úremap_to_usedA   s    "

(
zVectorQuantizer2.remap_to_usedc                 C   s”   |j }t|ƒdksJ ‚| |d d¡}| j |¡}| j| jj d krXd||| jj d k< t |d d d …f |j d dg d d …f d|¡}| |¡S )Nr   r   r%   )r   r)   r*   r   r+   r   r   Úgather)r    r0   r1   r   Úbackr#   r#   r$   Úunmap_to_allO   s    2zVectorQuantizer2.unmap_to_allc              	   C   s¸  |d u s|dksJ dƒ‚|du s(J dƒ‚|du s8J dƒ‚t |dƒ ¡ }| d| j¡}tj|d ddd	tj| jjd dd
 dt d|t | jjdƒ¡  }tj	|dd
}|  |¡ |j
¡}d }	d }
| jsü| jt | ¡ | d ¡ t || ¡  d ¡ }n2t | ¡ | d ¡| jt || ¡  d ¡  }|||  ¡  }t |dƒ ¡ }| jd ur€| |j
d d¡}|  |¡}| dd¡}| jr¨| |j
d |j
d |j
d ¡}|||	|
|ffS )Nr   z)Only for interface compatible with GumbelFzb c h w -> b h w cr%   r&   r   T)ÚdimÚkeepdim)r9   z	bd,dn->bnz
n d -> d nzb h w c -> b c h wr   é   )r   Ú
contiguousÚviewr   r   r.   r   r   ÚeinsumÚargminr   r   r   ÚmeanÚdetachr   r*   r5   r   )r    ÚzÚtempZrescale_logitsÚreturn_logitsÚz_flattenedÚdÚmin_encoding_indicesÚz_qÚ
perplexityÚmin_encodingsÚlossr#   r#   r$   ÚforwardY   sD    ÿÿÿÿÿ
ÿzVectorQuantizer2.forwardc                 C   sb   | j d ur.| |d d¡}|  |¡}| d¡}|  |¡}|d ur^| |¡}| dddd¡ ¡ }|S )Nr   r%   r;   r   r&   )r   r*   r8   r   r=   Úpermuter<   )r    Úindicesr   rH   r#   r#   r$   Úget_codebook_entry„   s    




z#VectorQuantizer2.get_codebook_entry)Nr   FT)NFF)
Ú__name__Ú
__module__Ú__qualname__Ú__doc__r
   r5   r8   rL   rO   Ú__classcell__r#   r#   r!   r$   r      s     ÿ

+r   )	r   Útorch.nnr   Únumpyr   Úeinopsr   ÚModuler   r#   r#   r#   r$   Ú<module>   s   