a
    d                     @   s   d dl Z d dlmZ d dlmZ ddlmZmZ G dd dejZ	G dd dej
ZG d	d
 d
ej
ZG dd dej
ZG dd dej
Ze G dd dej
ZdS )    N)ARCH_REGISTRY   )ResidualBlockNoBN
make_layerc                       s"   e Zd ZdZd fdd	Z  ZS )	MeanShifta   Data normalization with mean and std.

    Args:
        rgb_range (int): Maximum value of RGB.
        rgb_mean (list[float]): Mean for RGB channels.
        rgb_std (list[float]): Std for RGB channels.
        sign (int): For subtraction, sign is -1, for addition, sign is 1.
            Default: -1.
        requires_grad (bool): Whether to update the self.weight and self.bias.
            Default: True.
    Tc                    s   t t| jdddd t|}tddddd| j_| jj	|dddd || t| | j
_| j
j	| || _d S )N   r   )kernel_size)superr   __init__torchZTensorZeyeviewZweightdataZdiv_Zbiasrequires_grad)selfZ	rgb_rangergb_meanrgb_stdsignr   Zstd	__class__ b/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/basicsr/archs/ridnet_arch.pyr      s    
zMeanShift.__init__)r   T)__name__
__module____qualname____doc__r   __classcell__r   r   r   r   r      s   r   c                       s(   e Zd ZdZ fddZdd Z  ZS )EResidualBlockNoBNzEnhanced Residual block without BN.

    There are three convolution layers in residual branch.

    It has a style of:
        ---Conv-ReLU-Conv-ReLU-Conv-+-ReLU-
         |__________________________|
    c                    sn   t t|   tt||dddtjddt||dddtjddt||ddd| _tjdd| _d S )Nr   r   TZinplacer   )	r
   r   r   nn
SequentialConv2dReLUbodyrelu)r   in_channelsout_channelsr   r   r   r   )   s    

zEResidualBlockNoBN.__init__c                 C   s   |  |}| || }|S N)r#   r$   r   xoutr   r   r   forward5   s    
zEResidualBlockNoBN.forwardr   r   r   r   r   r+   r   r   r   r   r   r      s   	r   c                       s*   e Zd ZdZd fdd	Zdd Z  ZS )	MergeRuna   Merge-and-run unit.

    This unit contains two branches with different dilated convolutions,
    followed by a convolution to process the concatenated features.

    Paper: Real Image Denoising with Feature Attention
    Ref git repo: https://github.com/saeed-anwar/RIDNet
    r   r   c                    s   t t|   tt|||||tjddt||||ddtjdd| _tt||||ddtjddt||||ddtjdd| _tt|d ||||tjdd| _	d S )NTr      r      )
r
   r-   r   r   r    r!   r"   	dilation1	dilation2aggregation)r   r%   r&   r	   Zstridepaddingr   r   r   r   E   s    zMergeRun.__init__c                 C   s<   |  |}| |}tj||gdd}| |}|| }|S )Nr   )Zdim)r0   r1   r   catr2   )r   r)   r0   r1   r*   r   r   r   r+   R   s    


zMergeRun.forward)r   r   r   r,   r   r   r   r   r-   ;   s   	r-   c                       s*   e Zd ZdZd fdd	Zdd Z  ZS )ChannelAttentionzChannel attention.

    Args:
        num_feat (int): Channel number of intermediate features.
        squeeze_factor (int): Channel squeeze factor. Default:
       c                    s\   t t|   ttdtj||| dddtjddtj|| |dddt | _	d S )Nr   r   )r3   Tr   )
r
   r5   r   r   r    ZAdaptiveAvgPool2dr!   r"   ZSigmoid	attention)r   mid_channelsZsqueeze_factorr   r   r   r   c   s
    $zChannelAttention.__init__c                 C   s   |  |}|| S r'   )r7   )r   r)   yr   r   r   r+   i   s    
zChannelAttention.forward)r6   r,   r   r   r   r   r5   [   s   r5   c                       s(   e Zd ZdZ fddZdd Z  ZS )EAMak  Enhancement attention modules (EAM) in RIDNet.

    This module contains a merge-and-run unit, a residual block,
    an enhanced residual block and a feature attention unit.

    Attributes:
        merge: The merge-and-run unit.
        block1: The residual block.
        block2: The enhanced residual block.
        ca: The feature/channel attention unit.
    c                    sL   t t|   t||| _t|| _t||| _t	|| _
tjdd| _d S )NTr   )r
   r:   r   r-   merger   block1r   block2r5   car   r"   r$   )r   r%   r8   r&   r   r   r   r   {   s    

zEAM.__init__c                 C   s2   |  |}| | |}| |}| |}|S r'   )r;   r$   r<   r=   r>   r(   r   r   r   r+      s
    


zEAM.forwardr,   r   r   r   r   r:   n   s   
r:   c                       s*   e Zd ZdZd
 fdd	Zdd	 Z  ZS )RIDNeta0  RIDNet: Real Image Denoising with Feature Attention.

    Ref git repo: https://github.com/saeed-anwar/RIDNet

    Args:
        in_channels (int): Channel number of inputs.
        mid_channels (int): Channel number of EAM modules.
            Default: 64.
        out_channels (int): Channel number of outputs.
        num_block (int): Number of EAM. Default: 4.
        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.
    r/        o@gw#?g8EGr?gB`"?      ?rC   rC   c                    sz   t t|   t|||| _t|||d| _t||ddd| _t	t
||||d| _t||ddd| _tjdd| _d S )Nr   r   )r%   r8   r&   Tr   )r
   r?   r   r   sub_meanadd_meanr   r!   headr   r:   r#   tailr"   r$   )r   r%   r8   r&   Z	num_blockZ	img_ranger   r   r   r   r   r      s    
zRIDNet.__init__c              	   C   s<   |  |}| | | | |}| |}|| }|S r'   )rD   rG   r#   r$   rF   rE   )r   r)   resr*   r   r   r   r+      s
    

zRIDNet.forward)r/   r@   rA   rB   r,   r   r   r   r   r?      s       r?   )r   Ztorch.nnr   Zbasicsr.utils.registryr   Z	arch_utilr   r   r!   r   Moduler   r-   r5   r:   registerr?   r   r   r   r   <module>   s    