a
    d=1                     @   s   d dl Z d dlZd dlZd dlZd dlZd dlZd dlZd dlm	Z	 d dl
mZ ejejejeZG dd dZG dd dejZG dd	 d	ejZdS )
    N)load_file_from_url)
functionalc                   @   sV   e Zd ZdZdddZdd
dZdd Zdd Zdd Zdd Z	e
 dddZdS )RealESRGANerar  A helper class for upsampling images with RealESRGAN.

    Args:
        scale (int): Upsampling scale factor used in the networks. It is usually 2 or 4.
        model_path (str): The path to the pretrained model. It can be urls (will first download it automatically).
        model (nn.Module): The defined network. Default: None.
        tile (int): As too large images result in the out of GPU memory issue, so this tile option will first crop
            input images into tiles, and then process each of them. Finally, they will be merged into one image.
            0 denotes for do not use tile. Default: 0.
        tile_pad (int): The pad size for each tile, to remove border artifacts. Default: 10.
        pre_pad (int): Pad the input images to avoid border artifacts. Default: 10.
        half (float): Whether to use half precision during inference. Default: False.
    Nr   
   Fc                 C   s@  || _ || _|| _|| _d | _|| _|
rV|	d u rNttj	 rHd|
 ndn|	| _n&|	d u rvttj	 rpdndn|	| _t
|trt|t|ksJ d| |d |d |}n8|drt|tjtdd	d d
}tj|tdd}d|v rd}nd}|j|| d	d |  || j| _| jr<| j | _d S )Nzcuda:cpucudaz6model_path and dni_weight should have the save length.r      zhttps://weightsT)urlZ	model_dirprogress	file_nameZmap_locationZ
params_emaparams)strict)scale	tile_sizetile_padpre_pad	mod_scalehalftorchdevicer   Zis_available
isinstancelistlendni
startswithr   ospathjoinROOT_DIRloadZload_state_dictevaltomodel)selfr   Z
model_path
dni_weightr$   Ztiler   r   r   r   Zgpu_idZloadnetkeyname r(   Y/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/realesrgan/utils.py__init__   s<    &

zRealESRGANer.__init__r   r   c                 C   sj   t j|t |d}t j|t |d}||  D ]0\}}|d | |d || |   || |< q4|S )z|Deep network interpolation.

        ``Paper: Deep Network Interpolation for Continuous Imagery Effect Transition``
        r   r   r   )r   r!   r   items)r%   Znet_aZnet_br&   keylockZv_ar(   r(   r)   r   M   s
    *zRealESRGANer.dnic                 C   s  t t|d }|d| j| _| j	r<| j	 | _| j
dkrdt| jd| j
d| j
fd| _| jdkrvd| _n| jdkrd| _| jdurd\| _| _| j \}}}}|| j dkr| j|| j  | _|| j dkr| j|| j  | _t| jd| jd| jfd| _dS )	zVPre-process, such as pre-pad and mod pad, so that the images can be divisible
        )   r   r   r   Zreflectr/   r      N)r   r   )r   Z
from_numpynp	transposefloatZ	unsqueezer#   r   imgr   r   Fpadr   r   	mod_pad_h	mod_pad_wsize)r%   r4   _hwr(   r(   r)   pre_processX   s$    


zRealESRGANer.pre_processc                 C   s   |  | j| _d S N)r$   r4   outputr%   r(   r(   r)   processq   s    zRealESRGANer.processc           $      C   s`  | j j\}}}}|| j }|| j }||||f}| j || _t|| j }t|| j }	t|	D ]}
t|D ]}|| j }|
| j }|}t	|| j |}|}t	|| j |}t
|| j d}t	|| j |}t
|| j d}t	|| j |}|| }|| }|
| | d }| j dddd||||f }z8t  | |}W d   n1 sd0    Y  W n0 ty } ztd| W Y d}~n
d}~0 0 td| d||	   || j }|| j }|| j }|| j }|| | j } | || j  }!|| | j }"|"|| j  }#|dddd|"|#| |!f | jdddd||||f< qtqfdS )zIt will first crop input images to tiles, and then process each tile.
        Finally, all the processed tiles are merged into one images.

        Modified from: https://github.com/ata4/esrgan-launcher
        r   r   NErrorz	Tile /)r4   shaper   Z	new_zerosr?   mathceilr   rangeminmaxr   r   no_gradr$   RuntimeErrorprint)$r%   batchZchannelheightwidthZoutput_heightZoutput_widthZoutput_shapeZtiles_xZtiles_yyxZofs_xZofs_yZinput_start_xZinput_end_xZinput_start_yZinput_end_yZinput_start_x_padZinput_end_x_padZinput_start_y_padZinput_end_y_padZinput_tile_widthZinput_tile_heightZtile_idxZ
input_tileZoutput_tileerrorZoutput_start_xZoutput_end_xZoutput_start_yZoutput_end_yZoutput_start_x_tileZoutput_end_x_tileZoutput_start_y_tileZoutput_end_y_tiler(   r(   r)   tile_processu   sV    



"
. 



zRealESRGANer.tile_processc                 C   s   | j d urX| j \}}}}| jd d d d d|| j| j  d|| j| j  f | _| jdkr| j \}}}}| jd d d d d|| j| j  d|| j| j  f | _| jS )Nr   )r   r?   r9   r7   r   r8   r   )r%   r:   r;   r<   r(   r(   r)   post_process   s    
<
<zRealESRGANer.post_process
realesrganc                 C   s  |j dd \}}|tj}t|dkr:d}td nd}|| }t|j dkrhd}t|tj	}nz|j d dkrd	}|d d d d d
f }|d d d d dd
f }t|tj
}|dkrt|tj	}nd}t|tj
}| | | jdkr|   n|   |  }	|	j   dd }	t|	g dd d d d f d}	|dkrjt|	tj}	|d	krV|dkr| | | jdkr|   n|   |  }
|
j   dd }
t|
g dd d d d f d}
t|
tj}
n4|j dd \}}tj||| j || j ftjd}
t|	tj}	|
|	d d d d d
f< |dkrv|	d  tj}n|	d  tj}|d ur|t| jkrtj|t|| t|| ftj d}||fS )Nr   r/      i  z	Input is a 16-bit image   Lr0   ZRGBA   rU   ZRGBr   )r/   r   r   )r   r/   r   )interpolationg    @g     o@)!rD   Zastyper1   float32rI   rL   r   cv2ZcvtColorZCOLOR_GRAY2RGBZCOLOR_BGR2RGBr=   r   rS   rA   rT   dataZsqueezer3   r   Zclamp_numpyr2   ZCOLOR_BGR2GRAYresizer   ZINTER_LINEARZCOLOR_BGR2BGRAroundZuint16Zuint8intZINTER_LANCZOS4)r%   r4   ZoutscaleZalpha_upsamplerZh_inputZw_inputZ	max_rangeZimg_modealphaZ
output_imgZoutput_alphar;   r<   r?   r(   r(   r)   enhance   sl    


"




""


zRealESRGANer.enhance)NNr   r   r   FNN)r   r   )NrU   )__name__
__module____qualname____doc__r*   r   r=   rA   rS   rT   r   rJ   rc   r(   r(   r(   r)   r      s"           
0
Ar   c                       s8   e Zd ZdZ fddZdd Zdd Zdd	 Z  ZS )
PrefetchReaderzPrefetch images.

    Args:
        img_list (list[str]): A image list of image paths to be read.
        num_prefetch_queue (int): Number of prefetch queue.
    c                    s    t    t|| _|| _d S r>   )superr*   queueQueuequeimg_list)r%   rm   Znum_prefetch_queue	__class__r(   r)   r*     s    
zPrefetchReader.__init__c                 C   s6   | j D ]}t|tj}| j| q| jd  d S r>   )rm   r\   ZimreadZIMREAD_UNCHANGEDrl   put)r%   Zimg_pathr4   r(   r(   r)   run  s    
zPrefetchReader.runc                 C   s   | j  }|d u rt|S r>   )rl   getStopIteration)r%   Z	next_itemr(   r(   r)   __next__  s    
zPrefetchReader.__next__c                 C   s   | S r>   r(   r@   r(   r(   r)   __iter__$  s    zPrefetchReader.__iter__)	rd   re   rf   rg   r*   rq   rt   ru   __classcell__r(   r(   rn   r)   rh   
  s
   rh   c                       s$   e Zd Z fddZdd Z  ZS )
IOConsumerc                    s    t    || _|| _|| _d S r>   )ri   r*   _queueqidopt)r%   rz   rl   ry   rn   r(   r)   r*   *  s    
zIOConsumer.__init__c                 C   sR   | j  }t|tr|dkrq<|d }|d }t|| q td| j d d S )Nquitr?   	save_pathz
IO worker z	 is done.)rx   rr   r   strr\   ZimwriterL   ry   )r%   msgr?   r|   r(   r(   r)   rq   0  s    
zIOConsumer.run)rd   re   rf   r*   rq   rv   r(   r(   rn   r)   rw   (  s   rw   )r\   rE   r^   r1   r   rj   	threadingr   Zbasicsr.utils.download_utilr   Ztorch.nnr   r5   r   dirnameabspath__file__r    r   Threadrh   rw   r(   r(   r(   r)   <module>   s    }