a
    
d,                     @   s   d dl Z d dlmZ d dlmZ d dlmZmZmZm	Z	 d dl
Zd dlZd dlmZ d dlmZ d dlmZ dd	lmZmZ dd
lmZ ddlmZ ddlmZmZmZ G dd dejZ G dd de Z!G dd de!Z"G dd de"Z#G dd de Z$dS )    N)abstractmethod)contextmanager)AnyDictTupleUnion)
ListConfig)version)	load_file   )DecoderEncoder)DiagonalGaussianDistribution)LitEma)defaultget_obj_from_strinstantiate_from_configc                	       s   e Zd ZdZdedef edef eedef eeee	f d fddZ
e feeeee	f ddd	d
ZeedddZdd ZedddZeejdddZeejdddZdd ZedddZ  ZS )AbstractAutoencodera   
    This is the base class for all autoencoders, including image autoencoders, image autoencoders with discriminators,
    unCLIP models, etc. Hence, it is fairly general, and specific features
    (e.g. discriminator training, encoding, decoding) must be implemented in subclasses.
    Njpg )	ema_decaymonitor	input_key	ckpt_pathignore_keysc                    s   t    || _|d u| _|d ur(|| _| jrZt| |d| _tdtt	| j
  d |d urp| j||d ttjtdkrd| _d S )N)decayzKeeping EMAs of .r   z2.0.0F)super__init__r   use_emar   r   	model_emaprintlenlistbuffersinit_from_ckptr	   parsetorch__version__automatic_optimization)selfr   r   r   r   r   	__class__r   ]/var/www/html/stable-diffusion-webui/repositories/generative-models/sgm/models/autoencoder.pyr      s    

zAbstractAutoencoder.__init__)pathr   returnc           	      C   s   | drtj|ddd }n| dr2t|}ntt| }|D ].}|D ]$}t||rNt	d
| ||= qNqF| j|dd\}}t	d	| d
t| dt| d t|dkrt	d|  t|dkrt	d|  d S )Nckptcpu)map_location
state_dictsafetensorsz Deleting key {} from state_dict.F)strictzRestored from z with z missing and z unexpected keysr   zMissing Keys: zUnexpected Keys: )endswithr(   loadload_safetensorsNotImplementedErrorr$   keysrematchr"   formatload_state_dictr#   )	r+   r/   r   sdr;   kikmissing
unexpectedr   r   r.   r&   1   s&    



z"AbstractAutoencoder.init_from_ckptr0   c                 C   s
   t  d S Nr:   r+   batchr   r   r.   	get_inputJ   s    zAbstractAutoencoder.get_inputc                 O   s   | j r| |  d S rF   )r    r!   r+   argskwargsr   r   r.   on_train_batch_endN   s    z&AbstractAutoencoder.on_train_batch_endc              
   c   s   | j r8| j|   | j|  |d ur8t| d z6d V  W | j r| j|   |d urt| d n.| j r| j|   |d urt| d 0 d S )Nz: Switched to EMA weightsz: Restored training weights)r    r!   store
parameterscopy_tor"   restore)r+   contextr   r   r.   	ema_scopeS   s    zAbstractAutoencoder.ema_scopec                 O   s   t dd S )Nz-encode()-method of abstract base class calledrG   rK   r   r   r.   encodeb   s    zAbstractAutoencoder.encodec                 O   s   t dd S )Nz-decode()-method of abstract base class calledrG   rK   r   r   r.   decodef   s    zAbstractAutoencoder.decodec                 C   s:   t d|d  d t|d |fd|i|dt S )Nzloading >>> targetz <<< optimizer from configlrparams)r"   r   getdict)r+   rY   rX   cfgr   r   r.   !instantiate_optimizer_from_configj   s    
z5AbstractAutoencoder.instantiate_optimizer_from_configc                 C   s
   t  d S rF   rG   r+   r   r   r.   configure_optimizersp   s    z(AbstractAutoencoder.configure_optimizers)NNr   Nr   )N)__name__
__module____qualname____doc__r   floatstrr   r$   r   r   tupler&   r   r   rJ   rN   r   rT   r(   TensorrU   rV   r]   r_   __classcell__r   r   r,   r.   r      s:        


r   c                       s  e Zd ZdZdddeeeeeedf ed fddZeej	dd	d
Z
edddZedddZdd Zd'eeedddZeej	dddZeeej	ej	ej	f dddZedddZedddZd(edd!d"Zedd#d$Ze eedd%d&Z  ZS ))AutoencodingEnginez
    Base class for all image autoencoders that we train, like VQGAN or AutoencoderKL
    (we also restore them explicitly as special cases for legacy reasons).
    Regularizations such as KL or VQ are moved to the regularizer class.
    N      ?)optimizer_configlr_g_factor)encoder_configdecoder_configloss_configregularizer_configrk   rl   c          	         sT   t  j|i | t|| _t|| _t|| _t|| _t|ddi| _|| _	d S )NrW   ztorch.optim.Adam)
r   r   r   encoderdecoderlossregularizationr   rk   rl   )	r+   rm   rn   ro   rp   rk   rl   rL   rM   r,   r   r.   r   {   s    



zAutoencodingEngine.__init__)rI   r0   c                 C   s
   || j  S rF   )r   rH   r   r   r.   rJ      s    zAutoencodingEngine.get_inputrE   c                 C   s<   t | j t | j  t | j  t | j  }|S rF   )r$   rq   rP   rr   rt   get_trainable_parametersrs   Z$get_trainable_autoencoder_parametersr+   rY   r   r   r.   get_autoencoder_params   s    z)AutoencodingEngine.get_autoencoder_paramsc                 C   s   t | j }|S rF   )r$   rs   ru   rv   r   r   r.   get_discriminator_params   s    z+AutoencodingEngine.get_discriminator_paramsc                 C   s
   | j  S rF   )rr   get_last_layerr^   r   r   r.   ry      s    z!AutoencodingEngine.get_last_layerF)xreturn_reg_logr0   c                 C   s(   |  |}| |\}}|r$||fS |S rF   )rq   rt   )r+   rz   r{   zreg_logr   r   r.   rU      s
    
zAutoencodingEngine.encode)r|   r0   c                 C   s   |  |}|S rF   )rr   )r+   r|   rz   r   r   r.   rV      s    
zAutoencodingEngine.decoderz   r0   c                 C   s&   | j |dd\}}| |}|||fS )NT)r{   )rU   rV   )r+   rz   r|   r}   decr   r   r.   forward   s    
zAutoencodingEngine.forwardc              	   C   s   |  |}| |\}}}|dkrZ| j||||| j|  dd\}}	| j|	ddddd |S |dkr| j||||| j|  dd\}
}| j|ddddd |
S d S )Nr   trainZ
last_layersplitFT)prog_barloggeron_stepon_epoch   )rJ   rs   global_stepry   log_dict)r+   rI   	batch_idxoptimizer_idxrz   r|   xrecregularization_logaelosslog_dict_aedisclosslog_dict_discr   r   r.   training_step   s<    




	
z AutoencodingEngine.training_stepc                 C   sR   |  ||}|  * | j ||dd}|| W d    n1 sD0    Y  |S )NZ_ema)postfix)_validation_steprT   update)r+   rI   r   r   Zlog_dict_emar   r   r.   validation_step   s
    
(z"AutoencodingEngine.validation_step c              	   C   s   |  |}| |\}}}| j|||d| j|  d| d\}}	| j|||d| j|  d| d\}
}| d| d|	d| d  |	| | |	 |	S )Nr   valr   r   z	/rec_loss)rJ   rs   r   ry   logr   r   )r+   rI   r   r   rz   r|   r   r   r   r   r   r   r   r   r.   r      s0    



	 

z#AutoencodingEngine._validation_stepc                 C   sL   |   }|  }| |t| jd| j | j}| || j| j}||gg fS )Nrj   )rw   rx   r]   r   rl   learning_raterk   )r+   Z	ae_paramsZdisc_paramsZopt_aeZopt_discr   r   r.   r_      s    
z'AutoencodingEngine.configure_optimizersc                 K   sp   t  }| |}| |\}}}||d< ||d< |  & | |\}}}||d< W d    n1 sb0    Y  |S )NinputsZreconstructionsZreconstructions_ema)r[   rJ   rT   )r+   rI   rM   r   rz   _r   Zxrec_emar   r   r.   
log_images  s    

&zAutoencodingEngine.log_images)F)r   )r`   ra   rb   rc   r   r   rd   r   r(   rg   rJ   r$   rw   rx   ry   r   boolrU   rV   r   r   r   r   r   r_   no_gradr   rh   r   r   r,   r.   ri   t   s0   
	 %ri   c                       s2   e Zd Zed fddZdd Zdd Z  ZS )AutoencoderKL)	embed_dimc                    s   | d}| dd }| dd}t jf ddiddiddi| dd| |d	 s\J tf i || _tf i || _tj	d
|d  d
| d| _
tj	||d d| _|| _|d ur| j||d d S )Nddconfigr   r   r   rW   ztorch.nn.IdentityZ
lossconfig)rm   rn   rp   ro   Zdouble_zr   Z
z_channelsr   r   )popr   r   r   rq   r   rr   r(   nnConv2d
quant_convpost_quant_convr   r&   )r+   r   rM   r   r   r   r,   r   r.   r     s&    
zAutoencoderKL.__init__c                 C   s8   | j rJ | jj d| |}| |}t|}|S )Nz" only supports inference currently)trainingr-   r`   rq   r   r   )r+   rz   hZmomentsZ	posteriorr   r   r.   rU   0  s    

zAutoencoderKL.encodec                 K   s    |  |}| j|fi |}|S rF   )r   rr   )r+   r|   Zdecoder_kwargsr   r   r   r.   rV   9  s    
zAutoencoderKL.decode)r`   ra   rb   intr   rU   rV   rh   r   r   r,   r.   r     s   	r   c                       s   e Zd Z fddZ  ZS )AutoencoderKLInferenceWrapperc                    s   t  | S rF   )r   rU   sampler+   rz   r,   r   r.   rU   @  s    z$AutoencoderKLInferenceWrapper.encode)r`   ra   rb   rU   rh   r   r   r,   r.   r   ?  s   r   c                       sL   e Zd Z fddZeedddZeedddZeeddd	Z  ZS )
IdentityFirstStagec                    s   t  j|i | d S rF   )r   r   rK   r,   r   r.   r   E  s    zIdentityFirstStage.__init__r~   c                 C   s   |S rF   r   r   r   r   r.   rJ   H  s    zIdentityFirstStage.get_inputc                 O   s   |S rF   r   r+   rz   rL   rM   r   r   r.   rU   K  s    zIdentityFirstStage.encodec                 O   s   |S rF   r   r   r   r   r.   rV   N  s    zIdentityFirstStage.decode)	r`   ra   rb   r   r   rJ   rU   rV   rh   r   r   r,   r.   r   D  s   r   )%r<   abcr   
contextlibr   typingr   r   r   r   pytorch_lightningplr(   	omegaconfr   	packagingr	   safetensors.torchr
   r9   Zmodules.diffusionmodules.modelr   r   Z#modules.distributions.distributionsr   Zmodules.emar   utilr   r   r   LightningModuler   ri   r   r   r   r   r   r   r.   <module>   s$   b '%