a
    
d"w                     @   s  d dl Z d dlmZ d dlZd dlmZ d dlmZ d dlZd dlm	Z	m
Z
 ddlmZ dd	 ZdOddZdPddZdQddZdRddZdd ZdSddZdd ZG dd dZG d d! d!Ze dddd"d"ed#dfd$d%Ze dTd&d'Ze dddd"d"ed#dfd(d)Ze dddd"d"ed#dfd*d+Ze dUd,d-Zd.d/ Ze dVd1d2Ze dWd4d5Z G d6d7 d7Z!G d8d9 d9ej"Z#e dXd:d;Z$e dYdAdBZ%e dZdCdDZ&e d[dFdGZ'e d\dHdIZ(e d]dKdLZ)e d^dMdNZ*dS )_    N)	integrate)nn)odeint)trangetqdm   )utilsc                 C   s   t | | dggS Nr   )torchcat	new_zerosx r   U/var/www/html/stable-diffusion-webui/repositories/k-diffusion/k_diffusion/sampling.pyappend_zero   s    r         @cpuc           	      C   sH   t dd| }|d|  }|d|  }||||   | }t||S )z6Constructs the noise schedule of Karras et al. (2022).r   r   )r
   linspacer   to)	n	sigma_min	sigma_maxrhodevicerampmin_inv_rhomax_inv_rhosigmasr   r   r   get_sigmas_karras   s
    r   c                 C   s*   t jt|t|| |d }t|S )z)Constructs an exponential noise schedule.r   )r
   r   mathlogexpr   )r   r   r   r   r   r   r   r   get_sigmas_exponential   s    "r$         ?c                 C   sF   t jdd| |d| }t |t|t|  t| }t|S )z5Constructs an polynomial in log sigma noise schedule.r   r   r    )r
   r   r#   r!   r"   r   )r   r   r   r   r   r   r   r   r   r   get_sigmas_polyexponential    s    (r&   fffff3@皙?MbP?c                 C   sB   t jd|| |d}t t ||d  d ||  d }t|S )z*Constructs a continuous VP noise schedule.r   r       )r
   r   sqrtr#   r   )r   Zbeta_dZbeta_minZeps_sr   tr   r   r   r   get_sigmas_vp'   s    (r-   c                 C   s   | | t || j S )z6Converts a denoiser output to a Karras ODE derivative.)r   append_dimsndim)r   sigmadenoisedr   r   r   to_d.   s    r2   c                 C   sV   |s|dfS t |||d | d |d   | d  d  }|d |d  d }||fS )zCalculates the noise level (sigma_down) to step down to and the amount
    of noise to add (sigma_up) when doing an ancestral sampling step.        r*         ?)min)
sigma_fromsigma_toetasigma_up
sigma_downr   r   r   get_ancestral_step3   s
    .r;   c                    s    fddS )Nc                    s
   t  S N)r
   
randn_like)r0   
sigma_nextr   r   r   <lambda>>       z'default_noise_sampler.<locals>.<lambda>r   r   r   r   r   default_noise_sampler=   s    rA   c                   @   s.   e Zd ZdZd	ddZedd Zdd ZdS )
BatchedBrownianTreezGA wrapper around torchsde.BrownianTree that enables batches of entropy.Nc                    s   |  \| _ dt||d u r@tddg  }d| _z"t||j	d ks^J d W n t
y   |g}d| _Y n0  fdd|D | _d S )Nw0r   l    TFc                    s&   g | ]}t jfd |i qS )entropy)torchsdeZBrownianTree).0skwargst0t1rC   r   r   
<listcomp>P   r@   z0BatchedBrownianTree.__init__.<locals>.<listcomp>)sortsigngetr
   
zeros_likerandintitembatchedlenshape	TypeErrortrees)selfr   rJ   rK   seedrI   r   rH   r   __init__D   s    zBatchedBrownianTree.__init__c                 C   s   | |k r| |dfS || dfS )Nr   r   )abr   r   r   rM   R   s    zBatchedBrownianTree.sortc                    sJ   |   \ }t fdd| jD | j|  }| jrB|S |d S )Nc                    s   g | ]}| qS r   r   )rF   treerJ   rK   r   r   rL   X   r@   z0BatchedBrownianTree.__call__.<locals>.<listcomp>r   )rM   r
   stackrW   rN   rS   )rX   rJ   rK   rN   wr   r_   r   __call__V   s    &zBatchedBrownianTree.__call__)N)__name__
__module____qualname____doc__rZ   staticmethodrM   rb   r   r   r   r   rB   A   s
   

rB   c                   @   s*   e Zd ZdZddd fddZdd ZdS )	BrownianTreeNoiseSampleras  A noise sampler backed by a torchsde.BrownianTree.

    Args:
        x (Tensor): The tensor whose shape, device and dtype to use to generate
            random samples.
        sigma_min (float): The low end of the valid interval.
        sigma_max (float): The high end of the valid interval.
        seed (int or List[int]): The random seed. If a list of seeds is
            supplied instead of a single integer, then the noise sampler will
            use one BrownianTree per batch item, each with its own seed.
        transform (callable): A function that maps sigma to the sampler's
            internal timestep.
    Nc                 C   s   | S r<   r   r   r   r   r   r?   k   r@   z!BrownianTreeNoiseSampler.<lambda>c                 C   s<   || _ |  t||  t| }}t||||| _d S r<   )	transformr
   	as_tensorrB   r^   )rX   r   r   r   rY   ri   rJ   rK   r   r   r   rZ   k   s    "z!BrownianTreeNoiseSampler.__init__c                 C   s>   |  t||  t| }}| ||||    S r<   )ri   r
   rj   r^   absr+   )rX   r0   r>   rJ   rK   r   r   r   rb   p   s    "z!BrownianTreeNoiseSampler.__call__)rc   rd   re   rf   rZ   rb   r   r   r   r   rh   \   s   rh   r3   infc
                 C   s  |du ri n|}| |jd g}
tt|d |dD ]}|||   krR|krln nt|t|d  dnd}t||	 }|| |d  }|dkr|||d || d  d   }| |||
 fi |}t|||}|dur||||| ||d	 ||d  | }|||  }q6|S )
z?Implements Algorithm 2 (Euler steps) from Karras et al. (2022).Nr   r   disable4y?r3   r*   r4   r   ir0   	sigma_hatr1   new_onesrU   r   rT   r5   r
   r=   r2   )modelr   r   
extra_argscallbackrn   s_churns_tmins_tmaxs_noises_inrq   gammaepsrr   r1   ddtr   r   r   sample_euleru   s    6 r   c	                 C   s
  |du ri n|}|du r t |n|}||jd g}	tt|d |dD ]}
| |||
 |	 fi |}t||
 ||
d  |d\}}|dur|||
||
 ||
 |d t|||
 |}|||
  }|||  }||
d  dkrJ||||
 ||
d  | |  }qJ|S )z+Ancestral sampling with Euler method steps.Nr   r   rm   r8   rp   )rA   rt   rU   r   rT   r;   r2   )ru   r   r   rv   rw   rn   r8   r{   noise_samplerr|   rq   r1   r:   r9   r   r   r   r   r   sample_euler_ancestral   s    $r   c
                 C   s  |du ri n|}| |jd g}
tt|d |dD ]H}|||   krT|krnn nt|t|d  dnd}t||	 }|| |d  }|dkr|||d || d  d   }| |||
 fi |}t|||}|dur||||| ||d	 ||d  | }||d  dkr*|||  }q6|||  }| |||d  |
 fi |}t|||d  |}|| d }|||  }q6|S )
z>Implements Algorithm 2 (Heun steps) from Karras et al. (2022).Nr   r   rm   ro   r3   r*   r4   rp   rs   )ru   r   r   rv   rw   rn   rx   ry   rz   r{   r|   rq   r}   r~   rr   r1   r   r   x_2
denoised_2d_2d_primer   r   r   sample_heun   s*    6 r   c
                 C   s  |du ri n|}| |jd g}
tt|d |dD ]d}|||   krT|krnn nt|t|d  dnd}t||	 }|| |d  }|dkr|||d || d  d   }| |||
 fi |}t|||}|dur||||| ||d	 ||d  dkr*||d  | }|||  }q6| 	||d   d
 }|| }||d  | }|||  }| |||
 fi |}t|||}|||  }q6|S )
zMA sampler inspired by DPM-Solver-2 and Algorithm 2 from Karras et al. (2022).Nr   r   rm   ro   r3   r*   r4   rp   )rt   rU   r   rT   r5   r
   r=   r2   r"   lerpr#   )ru   r   r   rv   rw   rn   rx   ry   rz   r{   r|   rq   r}   r~   rr   r1   r   r   	sigma_middt_1dt_2r   r   r   r   r   r   sample_dpm_2   s.    6  r   c	                 C   st  |du ri n|}|du r t |n|}||jd g}	tt|d |dD ]"}
| |||
 |	 fi |}t||
 ||
d  |d\}}|dur|||
||
 ||
 |d t|||
 |}|dkr|||
  }|||  }qJ||
  | d	 }|||
  }|||
  }|||  }| |||	 fi |}t|||}|||  }||||
 ||
d  | |  }qJ|S )z6Ancestral sampling with DPM-Solver second-order steps.Nr   r   rm   r   rp   r4   )
rA   rt   rU   r   rT   r;   r2   r"   r   r#   )ru   r   r   rv   rw   rn   r8   r{   r   r|   rq   r1   r:   r9   r   r   r   r   r   r   r   r   r   r   r   sample_dpm_2_ancestral   s*    $r   c                    sT   d  kr t d d   fdd}tj|   d  ddd S )	Nr   zOrder z too high for step c                    sL   d}t D ]:}|krq||  |       |    9 }q|S )Nr%   )range)tauprodkrq   jorderr,   r   r   fn   s    .z"linear_multistep_coeff.<locals>.fn-C6?)epsrelr   )
ValueErrorr   quad)r   r,   rq   r   r   r   r   r   linear_multistep_coeff   s    r      c                    s
  |d u ri n|}| |jd g}|   g }tt|d |dD ]| || | fi |}	t|| |	}
||
 t||kr|	d |d ur||| | |	d t
d |  fddt D }|tdd t|t|D  }qJ|S )	Nr   r   rm   rp   c                    s   g | ]}t  |qS r   )r   )rF   r   	cur_orderrq   
sigmas_cpur   r   rL     r@   zsample_lms.<locals>.<listcomp>c                 s   s   | ]\}}|| V  qd S r<   r   )rF   coeffr   r   r   r   	<genexpr>  r@   zsample_lms.<locals>.<genexpr>)rt   rU   detachr   numpyr   rT   r2   appendpopr5   r   sumzipreversed)ru   r   r   rv   rw   rn   r   r|   dsr1   r   coeffsr   r   r   
sample_lms  s     

"r   r   c                    s    d u ri n  | |jd gt|dd d d fdd}|||jd gf}|||g}	t|||	||dd}
|
d d |
d d  }}tjd|	|
dd}|| d	ifS )
Nr   r*   r   c                    s   t  x |d   }||  fi  }t|| |}d7 t j|  |d }| dd}W d    n1 s0    Y  | |fS )Nr   r   )	r
   enable_gradr   requires_grad_r2   autogradgradr   flatten)r0   r   r1   r   r   Zd_llrv   fevalsru   r|   vr   r   ode_fn  s    
2zlog_likelihood.<locals>.ode_fndopri5)atolrtolmethodr[   r   )rt   rU   r
   randint_liker   
new_tensorr   distributionsNormallog_probr   r   )ru   r   r   r   rv   r   r   r   x_minr,   sollatentZdelta_llZll_priorr   r   r   log_likelihood  s    
 r   c                   @   s*   e Zd ZdZdddZdd Zd	d
 ZdS )PIDStepSizeControllerz4A PID controller for ODE adaptive step size control.r   Q?:0yE>c                 C   sL   || _ || | | | _|d|   | | _|| | _|| _|| _g | _d S )Nr*   )hb1b2b3accept_safetyr~   errs)rX   r   pcoefficoeffdcoeffr   r   r~   r   r   r   rZ   2  s    
zPIDStepSizeController.__init__c                 C   s   dt |d  S r	   )r!   atan)rX   r   r   r   r   limiter;  s    zPIDStepSizeController.limiterc                 C   s   dt || j  }| js$|||g| _|| jd< | jd | j | jd | j  | jd | j  }| |}|| jk}|r| jd | jd< | jd | jd< |  j|9  _|S )Nr   r   r*   )	floatr~   r   r   r   r   r   r   r   )rX   error	inv_errorfactoracceptr   r   r   propose_step>  s    
0

z"PIDStepSizeController.propose_stepN)r   r   r   )rc   rd   re   rf   rZ   r   r   r   r   r   r   r   0  s   
	r   c                       sl   e Zd ZdZd fdd	Zdd Zdd Zd	d
 ZdddZd ddZ	d!ddZ
d"ddZd#ddZ  ZS )$	DPMSolverz1DPM-Solver. See https://arxiv.org/abs/2206.00927.Nc                    s2   t    || _|d u ri n|| _|| _|| _d S r<   )superrZ   ru   rv   eps_callbackinfo_callback)rX   ru   rv   r   r   	__class__r   r   rZ   P  s
    
zDPMSolver.__init__c                 C   s
   |   S r<   )r"   )rX   r0   r   r   r   r,   W  s    zDPMSolver.tc                 C   s   |   S r<   negr#   )rX   r,   r   r   r   r0   Z  s    zDPMSolver.sigmac           	      O   s   ||v r|| |fS |  |||jd g }|| j||g|R i | j| |  | }| jd urp|   |||i|fS )Nr   )r0   rt   rU   ru   rv   r   )	rX   	eps_cachekeyr   r,   argsrI   r0   r~   r   r   r   r~   ]  s    .
zDPMSolver.epsc                 C   sN   |d u ri n|}|| }|  |d||\}}|| ||  |  }||fS )Nr~   r~   r0   expm1)rX   r   r,   t_nextr   r   r~   x_1r   r   r   dpm_solver_1_stepf  s
    zDPMSolver.dpm_solver_1_stepr4   c                 C   s   |d u ri n|}|| }|  |d||\}}|||  }|| |||   |  }	|  |d|	|\}
}|| ||  |  | |d|  |  |
|   }||fS )Nr~   eps_r1r*   r   )rX   r   r,   r   r1r   r   r~   s1u1r   r   r   r   r   dpm_solver_2_stepm  s    <zDPMSolver.dpm_solver_2_stepUUUUUU?UUUUUU?c                 C   s"  |d u ri n|}|| }|  |d||\}}|||  }	|||  }
|| |	||   |  }|  |d||	\}}|| |
||   |  | |
||  ||  ||  d  ||   }|  |d||
\}}|| ||  |  | || | | d  ||   }||fS )Nr~   r   r   eps_r2r   )rX   r   r,   r   r   r2r   r   r~   r   s2r   r   u2r   Zx_3r   r   r   dpm_solver_3_stepw  s    P@zDPMSolver.dpm_solver_3_stepr3   r%   c              	   C   s  |d u rt |n|}||ks(|r(tdt|d d }tj|||d |jd}	|d dkrvdg|d  ddg }
ndg|d  |d g }
tt|
D ]N}i }|	| |	|d   }}|rt	| 
|| 
||\}}t|| |}| 
|d | 
|d  d }n
|d }}| |d	||\}}|| 
||  }| jd urf| |||	| ||d
 |
| dkr| j||||d\}}n<|
| dkr| j||||d\}}n| j||||d\}}||| || 
|| 
|  }q|S )N"eta must be 0 for reverse sampling   r   r    r   r*   r4   r3   r~   )r   rq   r,   t_upr1   r   )rA   r   r!   floorr
   r   r   r   rT   r;   r0   minimumr,   r~   r   r   r   r   )rX   r   t_startt_endnfer8   r{   r   mtsordersrq   r   r,   r   sdsut_next_r~   r1   r   r   r   dpm_solver_fast  s6    "
$zDPMSolver.dpm_solver_fastr   皙?q?r   c               
   C   s  |d u rt |n|}|dvr$td||k}|s<|r<tdt||rJdnd }t|}t|}|}|}d}t|||	|
|rdn||}ddddd	}|r||d
 k rn||d
 kri }|rt|||j nt|||j }|r<t	| 
|| 
||\}}t|| |}| 
|d | 
|d  d }n
|d }}| |d||\}}|| 
||  }|dkr| j||||d\}}| j||||d\}}n.| j|||d|d\}}| j||||d\}}t||t| |  }tj|| | | d  }||}|r^|}||| || 
|| 
|  }|}|d  d7  < n|d  d7  < |d  |7  < |d  d7  < | jd ur| ||d d |||||jd| q||fS )N>   r*   r   zorder should be 2 or 3r   r   r[   Tg      ?r   )stepsr   n_acceptn_rejectgh㈵>r*   r4   r3   r~   r   r   )r   r   r  r  r   r  )r   rq   r,   r   r1   r   r   )rA   r   rk   r
   tensorr   r   r   maximumr;   r0   r,   r~   r   r   r   linalgnormnumelr   r   ) rX   r   r   r   r   r   r   h_initr   r   r   r   r8   r{   r   forwardrG   x_prevr   pidinfor   r,   r   r   t_r~   r1   x_lowZx_highdeltar   r   r   r   dpm_solver_adaptive  sV    

"("

  
"
*zDPMSolver.dpm_solver_adaptive)NNN)N)r4   N)r   r   N)r3   r%   N)r   r  r  r  r3   r%   r3   r   r3   r%   N)rc   rd   re   rf   rZ   r,   r0   r~   r   r   r   r  r  __classcell__r   r   r   r   r   M  s   	




'r   c              
      s   |dks|dkrt dt||df}t| ||jd durN fdd_|t|t||||	|
W  d   S 1 s0    Y  dS )zHDPM-Solver-Fast (fixed step size). See https://arxiv.org/abs/2206.00927.r   %sigma_min and sigma_max must not be 0)totalrn   r   Nc                    s&     | d  | d d| S Nr,   r   )r0   rr   r0   r  rw   Z
dpm_solverr   r   r?     r@   z!sample_dpm_fast.<locals>.<lambda>)	r   r   r   updater   r  r,   r
   r  )ru   r   r   r   r   rv   rw   rn   r8   r{   r   pbarr   r  r   sample_dpm_fast  s    r   r   r  r  r   Fc                    s   |dks|dkrt dt|dx}t| ||jd durL fdd_|t|t||||	|
|||||||\}}W d   n1 s0    Y  |r||fS |S )zPDPM-Solver-12 and 23 (adaptive step size). See https://arxiv.org/abs/2206.00927.r   r  rm   r  Nc                    s&     | d  | d d| S r  r  r  r  r   r   r?     r@   z%sample_dpm_adaptive.<locals>.<lambda>)	r   r   r   r  r   r  r,   r
   r  )ru   r   r   r   rv   rw   rn   r   r   r   r  r   r   r   r   r8   r{   r   Zreturn_infor  r  r   r  r   sample_dpm_adaptive  s    ^r!  c	                 C   s  |du ri n|}|du r t |n|}||jd g}	dd }
dd }tt|d |dD ]T}| ||| |	 fi |}t|| ||d  |d\}}|dur||||| || |d	 |dkrt||| |}|||  }|||  }n||| || }}d
}|| }|||  }|
||
| | | |  |  }| ||
||	 fi |}|
||
| | |  |  }||d  dkrZ|||| ||d  | |  }qZ|S )z<Ancestral sampling with DPM-Solver++(2S) second-order steps.Nr   c                 S   s   |    S r<   r   r,   r   r   r   r?     r@   z+sample_dpmpp_2s_ancestral.<locals>.<lambda>c                 S   s   |    S r<   r"   r   r  r   r   r   r?     r@   r   rm   r   rp   r4   )rA   rt   rU   r   rT   r;   r2   r   )ru   r   r   rv   rw   rn   r8   r{   r   r|   sigma_fnt_fnrq   r1   r:   r9   r   r   r,   r   rr   rG   r   r   r   r   r   sample_dpmpp_2s_ancestral  s0    &"$r'  r4   c
                 C   s6  ||dk   |  }
}|du r.t||
|n|}|du r>i n|}||jd g}dd }dd }tt|d |dD ]}| ||| | fi |}|dur||||| || |d ||d  dkr t||| |}||d  ||  }|||  }qx||| |||d   }}|| }|||	  }dd	|	  }t|||||\}}||}|||| | || 	 |  }||||||| |  }| |||| fi |}t|||||\}}||}d| | ||  }|||| | || 	 |  }||||||| |  }qx|S )
zDPM-Solver++ (stochastic).r   Nc                 S   s   |    S r<   r   r"  r   r   r   r?   %  r@   z"sample_dpmpp_sde.<locals>.<lambda>c                 S   s   |    S r<   r#  r  r   r   r   r?   &  r@   r   rm   rp   r*   )
r5   maxrh   rt   rU   r   rT   r2   r;   r   )ru   r   r   rv   rw   rn   r8   r{   r   r&  r   r   r|   r$  r%  rq   r1   r   r   r,   r   r   rG   facr   r   s_r   r   r  
denoised_dr   r   r   sample_dpmpp_sde  s:    $$ r,  c                 C   s\  |du ri n|}| |jd g}dd }dd }d}	tt|d |dD ]
}
| |||
 | fi |}|dur|||
||
 ||
 |d |||
 |||
d   }}|| }|	du s||
d  dkr|||| | |  |  }nb||||
d   }|| }ddd	|   | dd	|  |	  }|||| | |  |  }|}	qJ|S )
zDPM-Solver++(2M).Nr   c                 S   s   |    S r<   r   r"  r   r   r   r?   M  r@   z!sample_dpmpp_2m.<locals>.<lambda>c                 S   s   |    S r<   r#  r  r   r   r   r?   N  r@   r   rm   rp   r*   )rt   rU   r   rT   r   )ru   r   r   rv   rw   rn   r|   r$  r%  old_denoisedrq   r1   r,   r   r   h_lastr&  r+  r   r   r   sample_dpmpp_2mH  s&    $$"r/  midpointc
                 C   s  |	dvrt d||dk  |  }
}|du r>t||
|n|}|du rNi n|}||jd g}d}d}tt|d |dD ]}| ||| | fi |}|dur||||| || |d ||d  dkr|}n4||   ||d     }}|| }|| }||d  ||  | 	  | | | 
  |  }|dur|| }|	dkr|| | 
  | |  d d|  ||   }n4|	d	kr|d
| | 
   d|  ||   }|r|||| ||d  ||d   d| 
    |  }|}|}q|S )zDPM-Solver++(2M) SDE.>   r0  heunz(solver_type must be 'heun' or 'midpoint'r   Nr   rm   rp   r1  r0  r4   )r   r5   r(  rh   rt   rU   r   rT   r"   r#   r   r   r+   )ru   r   r   rv   rw   rn   r8   r{   r   solver_typer   r   r|   r-  r.  rq   r1   r,   rG   r   Zeta_hr&  r   r   r   sample_dpmpp_2m_sdeb  s:    "8

6
*>r4  c	                  C   s\  ||dk   |  }	}
|du r.t||	|
n|}|du r>i n|}||jd g}d\}}d\}}tt|d |dD ]}| ||| | fi |}|dur||||| || |d ||d  dkr|}nn||   ||d     }}|| }||d  }t	| | | 
  |  }|dur|| }|| }|| | }|| | }||| | ||   }|| ||  }| 
 | d }|| d }|||  ||  }n>|dur|| }|| | }| 
 | d }|||  }|rB|||| ||d  ||d   d| | 
    |  }|| }}|| }}qx|S )	zDPM-Solver++(3M) SDE.r   N)NNr   rm   rp   r4   r2  )r5   r(  rh   rt   rU   r   rT   r"   r
   r#   r   r   r+   ) ru   r   r   rv   rw   rn   r8   r{   r   r   r   r|   Z
denoised_1r   Zh_1Zh_2rq   r1   r,   rG   r   Zh_etar0r   Zd1_0Zd1_1d1d2Zphi_2phi_3r&  r   r   r   r   sample_dpmpp_3m_sde  sH    ""

B
r9  )r   r   )r   )r%   r   )r'   r(   r)   r   )r%   )NNNr%   r%   N)NNNr%   r%   N)NNNr   )Nr   r   )NNNr3   r%   N)NNNr   r  r  r  r3   r%   r3   r   r3   r%   NF)NNNr%   r%   N)NNNr%   r%   Nr4   )NNN)NNNr%   r%   Nr0  )NNNr%   r%   N)+r!   scipyr   r
   r   Ztorchdiffeqr   rE   	tqdm.autor   r    r   r   r   r$   r&   r-   r2   r;   rA   rB   rh   no_gradr   r   r   r   r   r   r   r   r   r   Moduler   r   r!  r'  r,  r/  r4  r9  r   r   r   r   <module>   sb   
	




 !),