CoolFace
Apppublic

declare-lab/tango2

sourceHugging Faceupdated 2y agoView on Hugging Face
92likes
vae.cpython-39.pyc97 linesDownload Raw Back to __pycache__
1a

2��'d�6�@s�ddlmZddlmZddlZddlZddlmZddl	m3Z4mZddlm
Z
mZmZeGdd	�d	e5��ZGd6d�dej�ZGdd
�d
ej�ZGdd�dej�ZGdd�de�ZdS)�)�	dataclass)�OptionalN�)�7BaseOutput�randn_tensor�)�UNetMidBlock2D�get_down_block�get_up_blockc@seZdZUdZejed<dS)�
DecoderOutputz�8    Output of decoding method.9 10    Args:11        sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):12            Decoded output sample of the model. Output of the last layer of the model.13    �sampleN)�__name__�14__module__�__qualname__�__doc__�torch�FloatTensor�__annotations__�rr�I/home/deep/Projects/audio_diffusion/diffusers/src/diffusers/models/vae.pyrs15rcs&eZdZd�fdd	�	Zd16d�Z�ZS)
�Encoder���DownEncoderBlock2D��@r� �siluTc	st���||_tjj||ddddd�|_d|_t�g�|_	|d}	t17|�D]R\}18}|	}||19}	|20t|�dk}
t||j||	|
dd||ddd�}|j	�
|�qNt|dd|ddd|dd	�|_tj|d|dd21�|_t��|_|r�d|n|}tj|d|ddd�|_d
|_dS)Nrrr��kernel_size�stride�padding�����ư>)2223num_layers�in_channels�out_channelsZadd_downsample�24resnet_epsZdownsample_padding�
resnet_act_fn�
resnet_groups�attn_num_head_channels�
temb_channels������default�r$r&r'Zoutput_scale_factorZresnet_time_scale_shiftr)r(r*��num_channels�25num_groups�epsr�r!F)�super�__init__�layers_per_blockr�nn�Conv2d�conv_in�	mid_block�26ModuleList�down_blocks�	enumerate�lenr	�appendr�	GroupNorm�
conv_norm_out�SiLU�conv_act�conv_out�gradient_checkpointing)�selfr$r%�down_block_types�block_out_channelsr5�norm_num_groups�act_fn�double_z�output_channel�iZdown_block_typeZ
input_channel�is_final_block�27down_blockZconv_out_channels��	__class__rrr4'sZ28��
�29zEncoder.__init__cCs�|}|�|�}|jrZ|jrZdd�}|jD]}tjj�||�|�}q(tjj�||j�|�}n|jD]}||�}q`|�|�}|�|�}|�	|�}|�30|�}|S)Ncs�fdd�}|S)Ncs�|�S�Nr��inputs��modulerr�custom_forwardrszFEncoder.forward.<locals>.create_custom_forward.<locals>.custom_forwardr�rUrVrrTr�create_custom_forwardqsz.Encoder.forward.<locals>.create_custom_forward)r8�trainingrDr;r�utils�31checkpointr9r@rBrC)rE�xrrXrNrrr�forwardks3233343536373839zEncoder.forward)rrrrrrrT�r
rrr4r]�
__classcell__rrrOrr&s�Drcs&eZdZd�fdd�	Zd	d40�Z�ZS)�Decoderr��UpDecoderBlock2Drrrrcst���||_tj||ddddd�|_d|_t�g�|_t	|dd|ddd|dd�|_t41t|��}|d}	t|�D]Z\}42}|	}||43}	|44t
|�dk}
t||jd||	d|
d||ddd	�}|j�|�|	}qvtj|d|dd45�|_t��|_tj|d|ddd�|_d|_dS)
Nr+rrrr"r,r-r)46r#r$r%�prev_output_channelZadd_upsampler&r'r(r)r*r.r2F)r3r4r5r6r7r8r9r:�	up_blocksr�list�reversedr<r=r47r>r?r@rArBrCrD)rEr$r%�up_block_typesrGr5rHrIZreversed_block_out_channelsrKrLZ
up_block_typercrM�up_blockrOrrr4�s\48 49���
50zDecoder.__init__cCs�|}|�|�}|jrZ|jrZdd�}tjj�||j�|�}|jD]}tjj�||�|�}q>n|�|�}|jD]}||�}qj|�|�}|�	|�}|�51|�}|S)Ncs�fdd�}|S)Ncs�|�SrQrrRrTrrrV�szFDecoder.forward.<locals>.create_custom_forward.<locals>.custom_forwardrrWrrTrrX�sz.Decoder.forward.<locals>.create_custom_forward)r8rYrDrrZr[r9rdr@rBrC)rE�zrrXrhrrrr]�s5253545556575859zDecoder.forward)rrrarrrrr^rrrOrr`�s�Dr`csBeZdZdZd�fdd�	Zdd	�Zd60d�Zdd
�Zdd�Z�Z	S)�VectorQuantizerz�61    Improved version over VectorQuantizer, can be used as a drop-in replacement. Mostly avoids costly matrix62    multiplications and allows for post-hoc remapping of indices.63    N�randomFTcs�t���||_||_||_||_t�|j|j�|_|jj	j64�d|jd|j�||_|jdur�|�
dt�t�|j���|jjd|_||_|jdkr�|j|_|jd|_td|j�d|j�d	|j�d65��n||_||_dS)Ng���?�usedr�extrarz66Remapping z indices to z indices. Using z for unknown indices.)r3r4�n_e�vq_embed_dim�beta�legacyr6�	Embedding�	embedding�weight�data�uniform_�remap�register_bufferr�tensor�np�loadrm�shape�re_embed�
unknown_index�print�sane_index_shape)rErorprqrxrr�rrrOrrr4�s,676869��zVectorQuantizer.__init__cCs�|j}t|�dksJ�|�|dd�}|j�|�}|dd�dd�df|dk��}|�d�}|�d�dk}|jdkr�t	j70d|j||jd�j|jd�||<n71|j||<|�|�S)	Nrrr+)NN.rrk)�size)�device)
r}r=�reshaperm�to�long�argmax�sumrr�randintr~r�)rE�inds�ishaperm�match�new�unknownrrr�
remap_to_useds"7273(74zVectorQuantizer.remap_to_usedcCs�|j}t|�dksJ�|�|dd�}|j�|�}|j|jjdkrXd|||jjdk<t�|ddd�f|jddgdd�fd|�}|�|�S)Nrrr+)r}r=r�rmr�r~r�gather)rEr�r�rm�backrrr�unmap_to_all)s2zVectorQuantizer.unmap_to_allcCsR|�dddd���}|�d|j�}tjt�||jj�dd�}|�|��|j	�}d}d}|j75s�|jt�|�
�|d�t�||�
�d�}n2t�|�
�|d�|jt�||�
�d�}|||�
�}|�dddd���}|jdu�r|�|j	dd�}|�|�}|�dd�}|j�rB|�|j	d|j	d|j	d�}|||||ffS)Nrrrrr+��dim)�permute�76contiguous�viewrpr�argmin�cdistrtrur}rrrq�mean�detachrxr�r�r�)rEriZz_flattenedZmin_encoding_indices�z_q�77perplexityZ
min_encodings�lossrrrr]3s$4278 zVectorQuantizer.forwardcCsb|jdur.|�|dd�}|�|�}|�d�}|�|�}|dur^|�|�}|�dddd���}|S)Nrr+rrr)rxr�r�rtr�r�r�)rE�indicesr}r�rrr�get_codebook_entryUs7980818283z"VectorQuantizer.get_codebook_entry)NrkFT)84r
rrrr4r�r�r]r�r_rrrOrrj�s	�85"rjc@sReZdZddd�Zdeejejd�dd�Zddd	�Z	gd86�fdd�Z87d
d�ZdS)�DiagonalGaussianDistributionFcCs�||_tj|ddd�\|_|_t�|jdd�|_||_t�d|j�|_t�|j�|_	|jr~tj88|j|jj|jjd�|_	|_dS)Nrrr�g>�g4@��?)r��dtype)
�89parametersr�chunkr��logvar�clamp�
deterministic�exp�std�var�90zeros_liker�r�)rEr�r�rrrr4hs�z%DiagonalGaussianDistribution.__init__N)�	generator�returncCs0t|jj||jj|jjd�}|j|j|}|S)N)r�r�r�)rr�r}r�r�r�r�)rEr�rr\rrrrts91�z#DiagonalGaussianDistribution.samplecCs�|jrt�dg�S|durJdtjt�|jd�|jd|jgd�d�Sdtjt�|j|jd�|j|j|jd|j|jgd�d�SdS)N�r�rrl�rrrr�)r�r�Tensorr��powr�r�r�)rE�otherrrr�kl|s 092�����zDiagonalGaussianDistribution.klr�cCsR|jrt�dg�St�dtj�}dtj||jt�||j	d�|j93|d�S)Nr�g@r�rr�)r�rr�r{�log�pir�r�r�r�r�)rEr�dimsZlogtwopirrr�nll�sz DiagonalGaussianDistribution.nllcCs|jSrQ)r�)rErrr�mode�sz!DiagonalGaussianDistribution.mode)F)N)N)r
rrr4rr�	Generatorrrr�r�r�rrrrr�gs949596r�)�dataclassesr�typingr�numpyr{r�torch.nnr6rZrrZunet_2d_blocksrr	r97r�Modulerr`rj�objectr�rrrr�<module>shgr