a
    Ad                     @   s   d Z ddlZddlZddlmZ ddlmZmZmZ i Z	dd Z
G dd dejZG dd	 d	ejZd
d Zdd ZG dd dejZG dd dejZdd Zdd Zdd ZdS )z}
Tiny AutoEncoder for Stable Diffusion
(DNN for encoding / decoding SD's latent space)

https://github.com/madebyollin/taesd
    N)devicespaths_internalsharedc                 K   s   t j| |dfddi|S )N   padding   )nnConv2d)n_inn_outkwargs r   </var/www/html/stable-diffusion-webui/modules/sd_vae_taesd.pyconv   s    r   c                   @   s   e Zd Zedd ZdS )Clampc                 C   s   t | d d S )Nr   )torchtanh)xr   r   r   forward   s    zClamp.forwardN)__name__
__module____qualname__staticmethodr   r   r   r   r   r      s   r   c                       s$   e Zd Z fddZdd Z  ZS )Blockc              	      sj   t    tt||t t||t t||| _||krRtj||dddnt | _t | _	d S )Nr   Fbias)
super__init__r   
Sequentialr   ReLUr	   Identityskipfuse)selfr
   r   	__class__r   r   r      s    
.$zBlock.__init__c                 C   s   |  | || | S )N)r"   r   r!   )r#   r   r   r   r   r   !   s    zBlock.forward)r   r   r   r   r   __classcell__r   r   r$   r   r      s   r   c                   C   s   t t tddt  tddtddtddt jddtddddtddtddtddt jddtddddtddtddtddt jddtddddtddtddS )N   @      )scale_factorFr   r   )r   r   r   r   r   r   Upsampler   r   r   r   decoder%   s    ...r,   c                   C   s   t tddtddtdddddtddtddtddtdddddtddtddtddtdddddtddtddtddtddS )Nr   r(   r)   F)strider   r'   )r   r   r   r   r   r   r   r   encoder/   s    &&&r.   c                       s&   e Zd ZdZdZd fdd	Z  ZS )TAESDDecoderr         ?taesd_decoder.pthc                    s<   t    t | _| jtj|tjjdkr.dndd dS zKInitialize pretrained TAESD on the given device from the given checkpoints.cudacpuN)map_location)	r   r   r,   load_state_dictr   loadr   devicetype)r#   Zdecoder_pathr$   r   r   r   =   s
    
zTAESDDecoder.__init__)r1   r   r   r   Zlatent_magnitudeZlatent_shiftr   r&   r   r   r$   r   r/   9   s   r/   c                       s&   e Zd ZdZdZd fdd	Z  ZS )TAESDEncoderr   r0   taesd_encoder.pthc                    s<   t    t | _| jtj|tjjdkr.dndd dS r2   )	r   r   r.   r6   r   r7   r   r8   r9   )r#   Zencoder_pathr$   r   r   r   I   s
    
zTAESDEncoder.__init__)r<   r:   r   r   r$   r   r;   E   s   r;   c                 C   sB   t j| s>t jt j| dd td|   tj||  d S )NT)exist_okzDownloading TAESD model to: )	ospathexistsmakedirsdirnameprintr   hubdownload_url_to_file)
model_path	model_urlr   r   r   download_modelQ   s    rH   c                  C   s   t tjddrdnd} t| }|d u rtjtj	d| }t
|d|   tj|r~t|}|  |tjtj |t| < ntd|jS )Nis_sdxlFztaesdxl_decoder.pthr1   	VAE-taesd.https://github.com/madebyollin/taesd/raw/main/TAESD model not found)getattrr   sd_modelsd_vae_taesd_modelsgetr>   r?   joinr   models_pathrH   r@   r/   evaltor   r8   dtypeFileNotFoundErrorr,   
model_nameloaded_modelrF   r   r   r   decoder_modelY   s    

rZ   c                  C   s   t tjddrdnd} t| }|d u rtjtj	d| }t
|d|   tj|r~t|}|  |tjtj |t| < ntd|jS )NrI   Fztaesdxl_encoder.pthr<   rJ   rK   rL   )rM   r   rN   rO   rP   r>   r?   rQ   r   rR   rH   r@   r;   rS   rT   r   r8   rU   rV   r.   rW   r   r   r   encoder_modell   s    

r[   )__doc__r>   r   torch.nnr   modulesr   r   r   rO   r   Moduler   r   r,   r.   r/   r;   rH   rZ   r[   r   r   r   r   <module>   s   

