a
    
d                     @   sJ   d dl Z d dlZG dd dZG dd deZG dd deZdd	 ZdS )
    Nc                   @   s   e Zd Zdd Zdd ZdS )AbstractDistributionc                 C   s
   t  d S NNotImplementedErrorself r   z/var/www/html/stable-diffusion-webui/repositories/stable-diffusion-stability-ai/ldm/modules/distributions/distributions.pysample   s    zAbstractDistribution.samplec                 C   s
   t  d S r   r   r   r   r   r	   mode	   s    zAbstractDistribution.modeN)__name__
__module____qualname__r
   r   r   r   r   r	   r      s   r   c                   @   s$   e Zd Zdd Zdd Zdd ZdS )DiracDistributionc                 C   s
   || _ d S r   value)r   r   r   r   r	   __init__   s    zDiracDistribution.__init__c                 C   s   | j S r   r   r   r   r   r	   r
      s    zDiracDistribution.samplec                 C   s   | j S r   r   r   r   r   r	   r      s    zDiracDistribution.modeN)r   r   r   r   r
   r   r   r   r   r	   r      s   r   c                   @   s@   e Zd ZdddZdd ZdddZg d	fd
dZdd ZdS )DiagonalGaussianDistributionFc                 C   s   || _ tj|ddd\| _| _t| jdd| _|| _td| j | _t| j| _	| jr|t
| jj| j jd | _	| _d S )N      dimg      >g      4@      ?device)
parameterstorchchunkmeanlogvarclampdeterministicexpstdvar
zeros_liketor   )r   r   r!   r   r   r	   r      s    z%DiagonalGaussianDistribution.__init__c                 C   s*   | j | jt| j jj| jjd  }|S )Nr   )r   r#   r   randnshaper&   r   r   )r   xr   r   r	   r
   #   s    &z#DiagonalGaussianDistribution.sampleNc                 C   s   | j rtdgS |d u rJdtjt| jd| j d | j g dd S dtjt| j|j d|j | j|j  d | j |j g dd S d S )N        r   r   g      ?r   r      r   )r!   r   Tensorsumpowr   r$   r   )r   otherr   r   r	   kl'   s0    
zDiagonalGaussianDistribution.klr+   c                 C   sR   | j rtdgS tdtj }dtj|| j t|| j	 d| j
  |d S )Nr*   g       @r   r   r   )r!   r   r-   nplogpir.   r   r/   r   r$   )r   r
   dimslogtwopir   r   r	   nll5   s     z DiagonalGaussianDistribution.nllc                 C   s   | j S r   )r   r   r   r   r	   r   =   s    z!DiagonalGaussianDistribution.mode)F)N)r   r   r   r   r
   r1   r7   r   r   r   r   r	   r      s
   


r   c                    s   d | |||fD ]}t |tjr|  q*q dus:J d fdd||fD \}}dd| | t||  | | d t|    S )a*  
    source: https://github.com/openai/guided-diffusion/blob/27c20a8fab9cb472df5d6bdd6c8d11c8f430b924/guided_diffusion/losses.py#L12
    Compute the KL divergence between two gaussians.
    Shapes are automatically broadcasted, so batches can be compared to
    scalars, among other use cases.
    Nz&at least one argument must be a Tensorc                    s,   g | ]$}t |tjr|nt| qS r   )
isinstancer   r-   tensorr&   ).0r)   r9   r   r	   
<listcomp>Q   s   znormal_kl.<locals>.<listcomp>r   g      r   )r8   r   r-   r"   )mean1logvar1mean2logvar2objr   r;   r	   	normal_klA   s(    

rB   )r   numpyr2   r   r   objectr   rB   r   r   r   r	   <module>   s
   )