a
    dr                     @   sR   d dl Z d dl mZ d dlmZmZmZ d dlmZ e G dd dej	Z
dS )    N)nn)ResidualBlockNoBNUpsample
make_layer)ARCH_REGISTRYc                       s*   e Zd ZdZd fdd		Zd
d Z  ZS )EDSRa4  EDSR network structure.

    Paper: Enhanced Deep Residual Networks for Single Image Super-Resolution.
    Ref git repo: https://github.com/thstkdgus35/EDSR-PyTorch

    Args:
        num_in_ch (int): Channel number of inputs.
        num_out_ch (int): Channel number of outputs.
        num_feat (int): Channel number of intermediate features.
            Default: 64.
        num_block (int): Block number in the trunk network. Default: 16.
        upscale (int): Upsampling factor. Support 2^n and 3.
            Default: 4.
        res_scale (float): Used to scale the residual in residual block.
            Default: 1.
        img_range (float): Image range. Default: 255.
        rgb_mean (tuple[float]): Image mean in RGB orders.
            Default: (0.4488, 0.4371, 0.4040), calculated from DIV2K dataset.
    @                 o@gw#?g8EGr?gB`"?c	           	         s   t t|   || _t|dddd| _t	||ddd| _
tt|||dd| _t	||ddd| _t||| _t	||ddd| _d S )Nr      T)num_feat	res_scaleZpytorch_init)superr   __init__	img_rangetorchZTensorviewmeanr   ZConv2d
conv_firstr   r   bodyconv_after_bodyr   upsample	conv_last)	selfZ	num_in_chZ
num_out_chr   Z	num_blockZupscaler   r   Zrgb_mean	__class__ `/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/basicsr/archs/edsr_arch.pyr      s    	zEDSR.__init__c                 C   sd   | j || _ || j  | j }| |}| | |}||7 }| | |}|| j | j  }|S )N)r   Ztype_asr   r   r   r   r   r   )r   xresr   r   r    forward2   s    
zEDSR.forward)r   r	   r
   r   r   r   )__name__
__module____qualname____doc__r   r#   __classcell__r   r   r   r    r      s         r   )r   r   Zbasicsr.archs.arch_utilr   r   r   Zbasicsr.utils.registryr   registerModuler   r   r   r   r    <module>   s
   