a
    
dG                     @   s   d dl Z d dlZd dlmZ ddlmZmZ G dd dej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G dd deZdS )    N)nn   )samplingutilsc                       sH   e Zd ZdZ fddZdd Zdd Zdd	 Zd
d Zdd Z	  Z
S )	VDenoiserz4A v-diffusion-pytorch model wrapper for k-diffusion.c                    s   t    || _d| _d S )N      ?super__init__inner_model
sigma_data)selfr   	__class__ U/var/www/html/stable-diffusion-webui/repositories/k-diffusion/k_diffusion/external.pyr
      s    
zVDenoiser.__init__c                 C   sb   | j d |d | j d   }| | j  |d | j d  d  }d|d | j d  d  }|||fS N         ?r   r   r   sigmac_skipc_outc_inr   r   r   get_scalings   s    "zVDenoiser.get_scalingsc                 C   s   |  tj d S Nr   )atanmathpi)r   r   r   r   r   
sigma_to_t   s    zVDenoiser.sigma_to_tc                 C   s   |t j d  S r   )r   r   tan)r   tr   r   r   
t_to_sigma   s    zVDenoiser.t_to_sigmac                    s|    fdd|  |D \}}} |t| j  }| j|| | |fi |}	 ||  | }
|	|
 dddS )Nc                    s   g | ]}t | jqS r   r   append_dimsndim.0xinputr   r   
<listcomp>       z"VDenoiser.loss.<locals>.<listcomp>r   r   )	r   r   r%   r&   r   r    powflattenmeanr   r+   noiser   kwargsr   r   r   noised_inputmodel_outputtargetr   r*   r   loss   s
    zVDenoiser.lossc                    sH    fdd|  |D \}}}| j | | |fi ||  |  S )Nc                    s   g | ]}t | jqS r   r$   r'   r*   r   r   r,   %   r-   z%VDenoiser.forward.<locals>.<listcomp>)r   r   r    r   r+   r   r3   r   r   r   r   r*   r   forward$   s    zVDenoiser.forward)__name__
__module____qualname____doc__r
   r   r    r#   r7   r9   __classcell__r   r   r   r   r   	   s   r   c                       sT   e Zd ZdZ fddZedd Zedd Zdd	d
ZdddZ	dd Z
  ZS )DiscreteSchedulez[A mapping between continuous noise levels (sigmas) and a list of discrete noise
    levels.c                    s0   t    | d| | d|  || _d S )Nsigmas
log_sigmas)r	   r
   register_bufferlogquantize)r   r@   rD   r   r   r   r
   -   s    
zDiscreteSchedule.__init__c                 C   s
   | j d S )Nr   r@   r   r   r   r   	sigma_min3   s    zDiscreteSchedule.sigma_minc                 C   s
   | j d S )NrE   rF   r   r   r   	sigma_max7   s    zDiscreteSchedule.sigma_maxNc                 C   sN   |d u rt | jdS t| jd }tj|d|| jjd}t | |S )Nr   r   )device)	r   append_zeror@   fliplentorchlinspacerJ   r#   )r   nt_maxr"   r   r   r   
get_sigmas;   s
    zDiscreteSchedule.get_sigmasc                 C   s   |d u r| j n|}| }|| jd d d f  }|rL| jdd|jS |djddj	ddj
| jjd d d}|d }| j| | j|  }}|| ||  }	|	
dd}	d|	 | |	|  }
|
|jS )Nr   dimr   )maxr   )rD   rC   rA   absargminviewshapegecumsumargmaxclamp)r   r   rD   	log_sigmadistslow_idxhigh_idxlowhighwr"   r   r   r   r    B   s    .zDiscreteSchedule.sigma_to_tc                 C   sT   |  }|  |  |   }}}d| | j|  || j|   }| S )Nr   )floatfloorlongceilfracrA   exp)r   r"   r`   ra   rd   r^   r   r   r   r#   P   s    $ zDiscreteSchedule.t_to_sigma)N)N)r:   r;   r<   r=   r
   propertyrG   rI   rR   r    r#   r>   r   r   r   r   r?   )   s   



r?   c                       s@   e Zd ZdZ fddZdd Zdd Zdd	 Zd
d Z  Z	S )DiscreteEpsDDPMDenoiserzVA wrapper for discrete schedule DDPM models that output eps (the predicted
    noise).c                    s*   t  d| | d | || _d| _d S Nr   r   r   r   r   modelalphas_cumprodrD   r   r   r   r
   [   s    z DiscreteEpsDDPMDenoiser.__init__c                 C   s(   | }d|d | j d  d  }||fS )Nr   r   r   r   )r   r   r   r   r   r   r   r   `   s    z$DiscreteEpsDDPMDenoiser.get_scalingsc                 O   s   | j |i |S Nr   r   argsr3   r   r   r   get_epse   s    zDiscreteEpsDDPMDenoiser.get_epsc           	         sj    fdd|  |D \}} |t| j  }| j|| | |fi |}|| dddS )Nc                    s   g | ]}t | jqS r   r$   r'   r*   r   r   r,   i   r-   z0DiscreteEpsDDPMDenoiser.loss.<locals>.<listcomp>r   r   )	r   r   r%   r&   ru   r    r.   r/   r0   )	r   r+   r2   r   r3   r   r   r4   epsr   r*   r   r7   h   s    zDiscreteEpsDDPMDenoiser.lossc                    sF    fdd|  |D \}}| j | | |fi |} ||  S )Nc                    s   g | ]}t | jqS r   r$   r'   r*   r   r   r,   o   r-   z3DiscreteEpsDDPMDenoiser.forward.<locals>.<listcomp>)r   ru   r    )r   r+   r   r3   r   r   rv   r   r*   r   r9   n   s    zDiscreteEpsDDPMDenoiser.forward)
r:   r;   r<   r=   r
   r   ru   r7   r9   r>   r   r   r   r   rl   W   s   rl   c                       s*   e Zd ZdZd	 fdd	Zdd Z  ZS )
OpenAIDenoiserz&A wrapper for OpenAI diffusion models.FTcpuc                    s0   t j|j|t jd}t j|||d || _d S )N)rJ   dtyperD   )rN   tensorrp   float32r	   r
   has_learned_sigmas)r   ro   	diffusionrD   r}   rJ   rp   r   r   r   r
   w   s    zOpenAIDenoiser.__init__c                 O   s,   | j |i |}| jr(|jdddd S |S )Nr   r   rS   r   )r   r}   chunk)r   rt   r3   r5   r   r   r   ru   |   s    zOpenAIDenoiser.get_eps)FTrx   r:   r;   r<   r=   r
   ru   r>   r   r   r   r   rw   t   s   rw   c                       s*   e Zd ZdZd fdd	Zdd Z  ZS )	CompVisDenoiserz'A wrapper for CompVis diffusion models.Frx   c                    s   t  j||j|d d S Nrz   r	   r
   rp   r   ro   rD   rJ   r   r   r   r
      s    zCompVisDenoiser.__init__c                 O   s   | j j|i |S rq   r   apply_modelrs   r   r   r   ru      s    zCompVisDenoiser.get_eps)Frx   r   r   r   r   r   r      s   r   c                       s@   e Zd ZdZ fddZdd Zdd Zdd	 Zd
d Z  Z	S )DiscreteVDDPMDenoiserz:A wrapper for discrete schedule DDPM models that output v.c                    s*   t  d| | d | || _d| _d S rm   r   rn   r   r   r   r
      s    zDiscreteVDDPMDenoiser.__init__c                 C   sb   | j d |d | j d   }| | j  |d | j d  d  }d|d | j d  d  }|||fS r   r   r   r   r   r   r      s    "z"DiscreteVDDPMDenoiser.get_scalingsc                 O   s   | j |i |S rq   rr   rs   r   r   r   get_v   s    zDiscreteVDDPMDenoiser.get_vc                    s|    fdd|  |D \}}} |t| j  }| j|| | |fi |}	 ||  | }
|	|
 dddS )Nc                    s   g | ]}t | jqS r   r$   r'   r*   r   r   r,      r-   z.DiscreteVDDPMDenoiser.loss.<locals>.<listcomp>r   r   )	r   r   r%   r&   r   r    r.   r/   r0   r1   r   r*   r   r7      s
    zDiscreteVDDPMDenoiser.lossc                    sH    fdd|  |D \}}}| j | | |fi ||  |  S )Nc                    s   g | ]}t | jqS r   r$   r'   r*   r   r   r,      r-   z1DiscreteVDDPMDenoiser.forward.<locals>.<listcomp>)r   r   r    r8   r   r*   r   r9      s    zDiscreteVDDPMDenoiser.forward)
r:   r;   r<   r=   r
   r   r   r7   r9   r>   r   r   r   r   r      s   r   c                       s*   e Zd ZdZd fdd	Zdd Z  ZS )	CompVisVDenoiserz5A wrapper for CompVis diffusion models that output v.Frx   c                    s   t  j||j|d d S r   r   r   r   r   r   r
      s    zCompVisVDenoiser.__init__c                 K   s   | j |||S rq   r   )r   r)   r"   condr3   r   r   r   r      s    zCompVisVDenoiser.get_v)Frx   )r:   r;   r<   r=   r
   r   r>   r   r   r   r   r      s   r   )r   rN   r    r   r   Moduler   r?   rl   rw   r   r   r   r   r   r   r   <module>   s    .
