a
    
de.                     @   s  d Z ddlmZmZ ddlZddlmZmZ ddlmZ ddl	m
Z
mZmZmZmZ ddlmZmZmZ d	d
iZG dd d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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 )zh
    Partially ported from https://github.com/crowsonkb/k-diffusion/blob/master/k_diffusion/sampling.py
    )DictUnionN)
ListConfig	OmegaConf)tqdm   )get_ancestral_steplinear_multistep_coeffto_dto_neg_log_sigmato_sigma)append_dimsdefaultinstantiate_from_configtargetz3sgm.modules.diffusionmodules.guiders.IdentityGuiderc                   @   s\   e Zd Zdeeeef eedf eeeedf ee	dddZ
dddZd	d
 Zdd ZdS )BaseDiffusionSamplerNFcuda)discretization_config	num_stepsguider_configverbosedevicec                 C   s0   || _ t|| _tt|t| _|| _|| _d S N)r   r   discretizationr   DEFAULT_GUIDERguiderr   r   )selfr   r   r   r   r    r   l/var/www/html/stable-diffusion-webui/repositories/generative-models/sgm/modules/diffusionmodules/sampling.py__init__   s    
zBaseDiffusionSampler.__init__c                 C   sl   | j |d u r| jn|| jd}t||}|td|d d  9 }t|}||jd g}||||||fS )N)r         ?r          @)	r   r   r   r   torchsqrtlennew_onesshape)r   xconducr   sigmas
num_sigmass_inr   r   r   prepare_sampling_loop,   s    
z*BaseDiffusionSampler.prepare_sampling_loopc                 C   s&   || j |||| }|  ||}|S r   )r   prepare_inputs)r   r'   denoisersigmar(   r)   denoisedr   r   r   denoise9   s    zBaseDiffusionSampler.denoisec                 C   s|   t |d }| jrxtddd td| jj  td| jjj  td| jjj  t||d| jj d| d	d
}|S )N   z##############################z Sampling setting z	Sampler: zDiscretization: zGuider: zSampling with z for z steps)totaldesc)ranger   print	__class____name__r   r   r   )r   r+   Zsigma_generatorr   r   r   get_sigma_gen>   s    z"BaseDiffusionSampler.get_sigma_gen)NNFr   )NN)r9   
__module____qualname__r   r   r   r   intboolstrr   r-   r2   r:   r   r   r   r   r      s       

r   c                   @   s   e Zd Zdd Zdd ZdS )SingleStepDiffusionSamplerc           	      O   s   t d S r   )NotImplementedError)	r   r0   
next_sigmar/   r'   r(   r)   argskwargsr   r   r   sampler_stepN   s    z'SingleStepDiffusionSampler.sampler_stepc                 C   s   |||  S r   r   )r   r'   ddtr   r   r   
euler_stepQ   s    z%SingleStepDiffusionSampler.euler_stepN)r9   r;   r<   rE   rH   r   r   r   r   r@   M   s   r@   c                       s>   e Zd Zddeddf fdd	ZdddZdd	d
Z  ZS )
EDMSampler        infr    c                    s.   t  j|i | || _|| _|| _|| _d S r   )superr   s_churns_tmins_tmaxs_noise)r   rM   rN   rO   rP   rC   rD   r8   r   r   r   V   s
    zEDMSampler.__init__Nc              
   C   s   ||d  }|dkrHt || j }	||	t|d |d  |jd   }| |||||}
t|||
}t|| |j}| |||}| ||||||||}|S )Nr    r            ?)	r"   
randn_likerP   r   ndimr2   r
   rH   possible_correction_step)r   r0   rB   r/   r'   r(   r)   gammaZ	sigma_hatepsr1   rF   rG   rH   r   r   r   rE   `   s    $zEDMSampler.sampler_stepc              
   C   s   |  ||||\}}}}}}| |D ]f}	| j||	   krF| jkr^n nt| j|d  dnd}
| |||	  |||	d   |||||
}q&|S )Nr3   g4y?rJ   )r-   r:   rN   rO   minrM   rE   )r   r/   r'   r(   r)   r   r,   r*   r+   irW   r   r   r   __call__p   s(    


zEDMSampler.__call__)NrJ   )NN)r9   r;   r<   floatr   rE   r[   __classcell__r   r   rQ   r   rI   U   s   

rI   c                       s8   e Zd Zd fdd	Zdd Zdd Zdd	d
Z  ZS )AncestralSamplerr    c                    s,   t  j|i | || _|| _dd | _d S )Nc                 S   s
   t | S r   )r"   rT   r'   r   r   r   <lambda>       z+AncestralSampler.__init__.<locals>.<lambda>)rL   r   etarP   noise_sampler)r   rb   rP   rC   rD   rQ   r   r   r      s    zAncestralSampler.__init__c                 C   s*   t |||}t|| |j}| |||S r   )r
   r   rU   rH   )r   r'   r1   r0   
sigma_downrF   rG   r   r   r   ancestral_euler_step   s    z%AncestralSampler.ancestral_euler_stepc                 C   s:   t t||jdk|| || j t||j  |}|S )NrJ   )r"   wherer   rU   rc   rP   )r   r'   r0   rB   sigma_upr   r   r   ancestral_step   s    zAncestralSampler.ancestral_stepNc           
   	   C   sX   |  ||||\}}}}}}| |D ],}	| |||	  |||	d   ||||}q&|S )Nr3   r-   r:   rE   )
r   r/   r'   r(   r)   r   r,   r*   r+   rZ   r   r   r   r[      s    
	zAncestralSampler.__call__)r    r    )NN)r9   r;   r<   r   re   rh   r[   r]   r   r   rQ   r   r^      s   r^   c                       s(   e Zd Zd fdd	ZdddZ  ZS )	LinearMultistepSampler   c                    s   t  j|i | || _d S r   )rL   r   order)r   rl   rC   rD   rQ   r   r   r      s    zLinearMultistepSampler.__init__Nc                    s   |  ||||\}}}}	}}g }
|   | |	D ]||  }|| j||||i |}| ||}t|||}|
| t	|
| j
kr|
d td | j
  fddt D }|tdd t|t|
D  }q:|S )Nr   r3   c                    s   g | ]}t  |qS r   )r	   ).0jZ	cur_orderrZ   Z
sigmas_cpur   r   
<listcomp>   s   z3LinearMultistepSampler.__call__.<locals>.<listcomp>c                 s   s   | ]\}}|| V  qd S r   r   )rm   coeffrF   r   r   r   	<genexpr>   ra   z2LinearMultistepSampler.__call__.<locals>.<genexpr>)r-   detachcpunumpyr:   r   r.   r
   appendr$   rl   poprY   r6   sumzipreversed)r   r/   r'   r(   r)   r   rD   r,   r*   r+   dsr0   r1   rF   coeffsr   ro   r   r[      s.    

"zLinearMultistepSampler.__call__)rk   )NN)r9   r;   r<   r   r[   r]   r   r   rQ   r   rj      s    
rj   c                   @   s   e Zd Zdd ZdS )EulerEDMSamplerc	           	      C   s   |S r   r   )	r   rH   r'   rF   rG   rB   r/   r(   r)   r   r   r   rV      s    z(EulerEDMSampler.possible_correction_stepNr9   r;   r<   rV   r   r   r   r   r}      s   r}   c                   @   s   e Zd Zdd ZdS )HeunEDMSamplerc	                 C   sf   t |dk r|S | |||||}	t|||	}
||
 d }t t||jdk|||  |}|S d S )N+=r!   rJ   )r"   rx   r2   r
   rf   r   rU   )r   rH   r'   rF   rG   rB   r/   r(   r)   r1   Zd_newZd_primer   r   r   rV      s    z'HeunEDMSampler.possible_correction_stepNr~   r   r   r   r   r      s   r   c                   @   s   e Zd Zdd ZdS )EulerAncestralSamplerc           
      C   sJ   t ||| jd\}}| |||||}	| ||	||}| ||||}|S )Nrb   )r   rb   r2   re   rh   )
r   r0   rB   r/   r'   r(   r)   rd   rg   r1   r   r   r   rE      s
    z"EulerAncestralSampler.sampler_stepN)r9   r;   r<   rE   r   r   r   r   r      s   r   c                   @   s&   e Zd Zdd Zdd ZdddZdS )	DPMPP2SAncestralSamplerc                 C   s6   dd ||fD \}}|| }|d|  }||||fS )Nc                 S   s   g | ]}t |qS r   r   rm   sr   r   r   rp      ra   z9DPMPP2SAncestralSampler.get_variables.<locals>.<listcomp>rS   r   )r   r0   rd   tt_nexthr   r   r   r   get_variables   s    z%DPMPP2SAncestralSampler.get_variablesc           	      C   sB   t |t | }d|  }t |t | }|  }||||fS )Ng      ࿩r   expm1)	r   r   r   r   r   mult1mult2mult3mult4r   r   r   get_mult  s
    
z DPMPP2SAncestralSampler.get_multNc                    s   t ||| jd\}}	|  ||||}
|  |
||}t|dk rJ| n| ||\}}}} fdd| ||||D }|d   |d |
  }| ||t|||}|d   |d |  }t	t
| jd	k|| |  |||	  S )
Nr   r   c                    s   g | ]}t | jqS r   r   rU   rm   multr_   r   r   rp     s   z8DPMPP2SAncestralSampler.sampler_step.<locals>.<listcomp>r   r3   rR   r   rJ   )r   rb   r2   re   r"   rx   r   r   r   rf   r   rU   rh   )r   r0   rB   r/   r'   r(   r)   rD   rd   rg   r1   Zx_eulerr   r   r   r   r   x2Z	denoised2Z	x_dpmpp2sr   r_   r   rE   	  s    
z$DPMPP2SAncestralSampler.sampler_step)N)r9   r;   r<   r   r   rE   r   r   r   r   r      s   r   c                   @   s2   e Zd Zd
ddZdd ZdddZddd	ZdS )DPMPP2MSamplerNc           	      C   sV   dd ||fD \}}|| }|d urF|t | }|| }||||fS |d ||fS d S )Nc                 S   s   g | ]}t |qS r   r   r   r   r   r   rp   $  ra   z0DPMPP2MSampler.get_variables.<locals>.<listcomp>r   )	r   r0   rB   previous_sigmar   r   r   Zh_lastrr   r   r   r   #  s    zDPMPP2MSampler.get_variablesc           
      C   sV   t |t | }|  }|d urJddd|   }dd|  }	||||	fS ||fS d S )Nr3   rR   r   )
r   r   r   r   r   r   r   r   r   r   r   r   r   r   .  s    
zDPMPP2MSampler.get_multc	                    s   |   ||||}	| |||\}
}}} fdd| |
||||D }|d   |d |	  }|d u svt|dk r~||	fS |d |	 |d |  }|d   |d |  }tt| jdk||  |	fS )	Nc                    s   g | ]}t | jqS r   r   r   r_   r   r   rp   G  s   z/DPMPP2MSampler.sampler_step.<locals>.<listcomp>r   r3   r   rR   r   rJ   )r2   r   r   r"   rx   rf   r   rU   )r   old_denoisedr   r0   rB   r/   r'   r(   r)   r1   r   r   r   r   r   Z
x_standardZ
denoised_dZ
x_advancedr   r_   r   rE   9  s    
zDPMPP2MSampler.sampler_stepc                 K   s~   |  ||||\}}}}	}}d }
| |	D ]N}| j|
|dkr@d n|||d   |||  |||d   ||||d\}}
q*|S )Nr   r3   )r)   ri   )r   r/   r'   r(   r)   r   rD   r,   r*   r+   r   rZ   r   r   r   r[   [  s     
zDPMPP2MSampler.__call__)N)N)NN)r9   r;   r<   r   r   rE   r[   r   r   r   r   r   "  s
   
 
"r   )__doc__typingr   r   r"   	omegaconfr   r   r   Z'modules.diffusionmodules.sampling_utilsr   r	   r
   r   r   utilr   r   r   r   r   r@   rI   r^   rj   r}   r   r   r   r   r   r   r   r   <module>   s"   53(&
(