a
    AdO1                     @   s2  d dl Z d dlmZ d dlZd dlZd dlmZ d dlm	Z	m
Z
mZmZmZmZmZ d dlmZmZ d dlZedg dZG dd deZd)d	d
Zd ddddZd*ddZd+ddZdd Zd,ddZd-ddZd.ddZdd Zdd Z G dd  d e!Z"d!d" Z#e#  d#d$ Z$G d%d& d&Z%G d'd( d(Z&dS )/    N)
namedtuple)Image)devicesimagessd_vae_approxsd_samplerssd_vae_taesdshared	sd_models)optsstateSamplerData)nameconstructoraliasesoptionsc                   @   s   e Zd Zdd ZdS )r   c                 C   s   | j ddr|d }|S )Nsecond_orderF   )r   get)selfsteps r   B/var/www/html/stable-diffusion-webui/modules/sd_samplers_common.pytotal_steps   s    zSamplerData.total_stepsN)__name__
__module____qualname__r   r   r   r   r   r      s   c                 C   sf   t js|d urD|p| j}| jdkr6t|t| jd nd}|d }n| j}tt| jd| }||fS )Nr   g+?   )r   img2img_fix_stepsr   denoising_strengthintmin)pr   Zrequested_stepst_encr   r   r   setup_img2img_steps   s    
"
r$   r   r      )Fullz	Approx NNzApprox cheapTAESDc                 C   s
  |du st jjrPtjrPttjd}ddlm	} |dkrP|
t jrPt jjsPd}|dkrdt| }n|dkrt | tjtj }n||dkrt | tjtj }|d d }nJ|du rt j}t $ || |jj}W d   n1 s0    Y  |S )zhTransforms 4-channel latent space images into 3-channel RGB image tensors, with values in range [-1, 1].Nr   )lowvramr   r   r%   )r	   r   interruptedr   live_preview_fast_interruptapproximation_indexesr   show_progress_typemodulesr(   
is_enabledsd_modellive_preview_allow_lowvram_fullr   Zcheap_approximationmodeltor   devicedtypedetachr   Zdecoder_modelwithout_autocastdecode_first_stagefirst_stage_model)sampleapproximationr1   r(   x_sampler   r   r   samples_to_images_tensor%   s"    
2r<   c                 C   s\   t | d|d d d }tj|ddd}dt|  dd }|tj	}t
|S )Nr   g      ?              ?)r!   maxg     o@r   )r<   	unsqueezetorchclampnpmoveaxiscpunumpyastypeuint8r   	fromarray)r9   r:   r;   r   r   r   single_sample_to_image?   s
    rJ   c                 C   s&   | tj}ttjd}t||| S Nr   )r2   r   	dtype_vaer+   r   r   sd_vae_decode_methodr<   )r1   xZapprox_indexr   r   r   r7   I   s    r7   c                 C   s   t | | |S NrJ   )samplesindexr:   r   r   r   sample_to_imageO   s    rS   c                    s   t  fdd| D S )Nc                    s   g | ]}t | qS r   rP   ).0r9   r:   r   r   
<listcomp>T       z)samples_to_image_grid.<locals>.<listcomp>)r   
image_grid)rQ   r:   r   rU   r   samples_to_image_gridS   s    rY   c                    s   |du rt tjd}|dkr<| tjtj} t	 | }np du rJt
j  jtj | jt
jtjd} | d d } t| dkrt fdd| D }n  | }|S )	zimage[0, 1] -> latentNr   r%   )r4   r   r   c              
      s(   g | ] }   t|d d  qS )r   )get_first_stage_encodingencode_first_stagerA   r@   )rT   imgr1   r   r   rV   g   s   z,images_tensor_to_samples.<locals>.<listcomp>)r+   r   r   sd_vae_encode_methodr2   r   r3   r4   r   Zencoder_modelr	   r/   r8   rL   lenrA   stackrZ   r[   )imager:   r1   x_latentr   r]   r   images_tensor_to_samplesW   s     
rc   c                 C   sB   | t _tjr>tjdkr>tj jtj dkr>tjs>tj t	|  d S rK   )
r   current_latentr   live_previews_enableshow_progress_every_n_stepsr	   sampling_stepparallel_processing_allowedassign_current_imagerS   )decodedr   r   r   store_latents   s    "rk   c                 C   sl   t | j}| j}|du r,| jdur,| jj}|du rR|durR|jddrNdnd}|dkr^dS |jddS )zTreturns whether sampler from config will use eta noise seed delta for image creationNZdefault_eta_is_0Fr   r>   	uses_ensd)r   find_sampler_configsampler_nameetasamplerr   r   )r"   sampler_configro   r   r   r   %is_sampler_using_eta_noise_seed_delta{   s    rr   c                   @   s   e Zd ZdS )InterruptedExceptionN)r   r   r   r   r   r   r   rs      s   rs   c                  C   s   dd l } dd }|| jj_d S )Nr   c                 S   s   t || j||dS )N)r3   r4   )r   Zrandn_localr2   )sizer4   r3   seedr   r   r   torchsde_randn   s    z1replace_torchsde_browinan.<locals>.torchsde_randn)Z$torchsde._brownian.brownian_interval	_brownianbrownian_interval_randn)torchsderv   r   r   r   replace_torchsde_browinan   s    r{   c                 C   s  | j | j }| jj}| jj}|d ur0||k r0dS |d u sDtjj|krHdS t| jddr| jj	}t
jdkrp|rpdS t
jdkr|sdS t
jdkrt
j| jjd< |j| jjd< || jjd< t  tj|d W d    n1 s0    Y  t  | j  |   d	S )
NF	enable_hrz
first passzsecond passzHires refinerZRefinerzRefiner switch at)infoT)stepr   r"   refiner_switch_atrefiner_checkpoint_infor	   r/   sd_checkpoint_infogetattr
is_hr_passr   hires_fix_refiner_passextra_generation_paramsshort_titler
   SkipWritingToConfigreload_model_weightsr   torch_gcsetup_condsZupdate_inner_model)Zcfg_denoiserZcompleted_ratior   r   Zis_second_passr   r   r   apply_refiner   s.    

*
r   c                   @   s(   e Zd ZdZdd Zdd Zdd ZdS )	TorchHijacka;  This is here to replace torch.randn_like of k-diffusion.

    k-diffusion has random_sampler argument for most samplers, but not for all, so
    this is needed to properly replace every use of torch.randn_like.

    We need to replace to make images generated in batches to be same as images generated individually.c                 C   s   |j | _ d S rO   )rngr   r"   r   r   r   __init__   s    zTorchHijack.__init__c                 C   sB   |dkr| j S tt|r"tt|S tdt| j d| dd S )N
randn_like'z' object has no attribute ')r   hasattrrA   r   AttributeErrortyper   )r   itemr   r   r   __getattr__   s
    

zTorchHijack.__getattr__c                 C   s
   | j  S rO   )r   next)r   rN   r   r   r   r      s    zTorchHijack.randn_likeN)r   r   r   __doc__r   r   r   r   r   r   r   r      s   	r   c                   @   sV   e Zd Zdd Zdd Zdd Zdd Zed	d
dZdd Z	dddZ
dddZdS )Samplerc                 C   s   || _ || _g | _d | _d | _d | _d | _d | _d | _d| _	d| _
td| _d| _d| _d| _d| _tjjj| _d | _d | _d | _i | _d S )Nr=   infr>   eta_ancestralEta)funcnamefuncextra_paramsZsampler_noisesstop_atro   configlast_latents_min_unconds_churns_tminfloats_tmaxs_noiseeta_option_fieldeta_infotext_fieldeta_defaultr	   r/   r1   conditioning_keyr"   model_wrap_cfgsampler_extra_argsr   )r   r   r   r   r   r      s*    
zSampler.__init__c                 C   s4   |d }| j d ur || j kr t|t_tj  d S )Ni)r   rs   r   rg   r	   
total_tqdmupdate)r   dr~   r   r   r   callback_state   s
    zSampler.callback_statec                 C   sh   || j _| j|| j _|t_dt_z| W S  tyL   td | j	 Y S  t
yb   | j	 Y S 0 d S )Nr   zEncountered RecursionError during sampling, returning last latent. rho >5 with a polyexponential scheduler may cause this error. You should try to use a smaller rho value instead.)r   r   r   r   r   sampling_stepsrg   RecursionErrorprintr   rs   )r   r   r   r   r   r   launch_sampling   s    
zSampler.launch_samplingc                 C   s   |j S rO   )r   r   r   r   r   number_of_needed_noises  s    zSampler.number_of_needed_noises)returnc                 C   s  || _ || j_ t|dr|jnd | j_t|dr6|jnd | j_d| j_t|dd | j_|jd urf|jntt	| j
d| _t|dd| _t|tj_i }| jD ].}t||r|t| jjv rt||||< qdt| jjv r| j| jkr| j|j| j< | j|d< t| jdkr
tt	d|j}tt	d	|j}tt	d
|jpB| j}tt	d|j}d|v r|| jkr||d< ||_||jd< d	|v r|| jkr||d	< ||_||jd< d
|v r|| jkr||d
< ||_||jd< d|v r
|| jkr
||d< ||_||jd< |S )Nmasknmaskr   image_cfg_scaler=   r   ro   r   r   r   r   zSigma churnz
Sigma tminz
Sigma tmaxzSigma noise)r"   r   r   r   r   r~   r   r   ro   r   r   r   r   k_diffusionsamplingrA   r   inspect	signaturer   
parametersr   r   r   r_   r   r   r   r   )r   r"   extra_params_kwargs
param_namer   r   r   r   r   r   r   
initialize  sN     





zSampler.initializec                 C   sd   t jjrdS ddlm} ||dk  |  }}|j|j|j	 |jd |j	  }|||||dS )ziFor DPM++ SDE: manually create noise sampler to enable deterministic results across different batch sizesNr   )BrownianTreeNoiseSamplerr   )ru   )
r	   r   no_dpmpp_sde_batch_determinismk_diffusion.samplingr   r!   r?   	all_seeds	iteration
batch_size)r   rN   sigmasr"   r   	sigma_min	sigma_maxZcurrent_iter_seedsr   r   r   create_noise_samplerC  s    "zSampler.create_noise_samplerNc                 C   s
   t  d S rO   NotImplementedError)r   r"   rN   conditioningunconditional_conditioningr   image_conditioningr   r   r   r9   M  s    zSampler.samplec                 C   s
   t  d S rO   r   )r   r"   rN   noiser   r   r   r   r   r   r   sample_img2imgP  s    zSampler.sample_img2img)NN)NN)r   r   r   r   r   r   r   dictr   r   r9   r   r   r   r   r   r      s   	0

r   )N)NN)N)r   N)N)NN)'r   collectionsr   rF   rC   rA   PILr   r-   r   r   r   r   r   r	   r
   modules.sharedr   r   r   r   ZSamplerDataTupler   r$   r+   r<   rJ   r7   rS   rY   rc   rk   rr   BaseExceptionrs   r{   r   r   r   r   r   r   r   <module>   s2   $






	$