a
    AduH                  
   @   s0  d dl mZ d dlZd dlZd dlmZ d dlm  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d5ddZG dd dej	ZG dd dej	Zdd Zd6ddZd7ddZdd  Zd8d#d$ZG d%d& d&ej	Zd'd( Zd)d* Zd+d, ZG d-d. d.ej	Zd/d0 Zd9d3d4ZdS ):    )OrderedDictNc                       s(   e Zd Zd fd
d	ZdddZ  ZS )RRDBNet          N	leakyreluCNAupconvConv2DFc              	      s~  t t|   tt|d}|dkr*d}d| _|d dkrDd| _n|dkr^|d dkr^d| _t|dd d d} fdd	t|D }tdd |
d
}|dkrt	n|dkrt
ntd| d|dkrd d}n fdd	t|D }tdd  d}t|dd d d}|rDt|nd }t|ttg ||R  g||||R  | _d S )N   r      r      r   )kernel_size	norm_typeact_typeconvtypec                    s.   g | ]&}t d dddd ddqS )r   r   r   zeror   )r   gcstridebiaspad_typer   r   moder   gaussian_noiseplus)RRDB.0_)r   r   r   nfr   nrr    A/var/www/html/stable-diffusion-webui/modules/esrgan_model_arch.py
<listcomp>   s
   z$RRDBNet.__init__.<locals>.<listcomp>)r   r   r   r   r   r	   pixelshufflezupsample mode [] is not foundr   r   c                    s   g | ]} d qS )r%   r    r   )r   r   r   upsample_blockr    r!   r"   ,       )superr   __init__intmathlogresrgan_scale
conv_blockrangeupconv_blockpixelshuffle_blockNotImplementedErroract
sequentialShortcutBlockmodel)selfin_ncout_ncr   nbr   r   upscaler   r   r   Zupsample_moder   Zfinalactr   r   Z	n_upscaleZfea_convZ	rb_blocksZLR_conv	upsamplerZHR_conv0ZHR_conv1Zoutact	__class__)r   r   r   r   r   r   r   r&   r!   r)      sD    zRRDBNet.__init__c                 C   s>   | j dkrt|dd}n| j dkr0t|dd}n|}| |S )Nr   r   )scaler   )r-   pixel_unshuffler6   )r7   xZoutmfeatr    r    r!   forward5   s    

zRRDBNet.forward)r   r   r   Nr   r   r	   r
   NFF)N__name__
__module____qualname__r)   rC   __classcell__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 )r   zr
    Residual in Residual Dense Block
    (ESRGAN: Enhanced Super-Resolution Generative Adversarial Networks)
    r   r   r   r   Nr   r   r
   Fc                    s   t t|   |dkrt	 
d| _t	 
d| _t	 
d| _n: 	
fddt|D }tj	| | _
d S )Nr   spectral_normr   r   c                    s.   g | ]&}t 	 
d qS )rI   )ResidualDenseBlock_5Cr   r   r   r   r   r   r   r   r   r   r   r   rJ   r   r    r!   r"   V   s
   
z!RRDB.__init__.<locals>.<listcomp>)r(   r   r)   rK   RDB1RDB2RDB3r/   nn
SequentialRDBs)r7   r   r   r   r   r   r   r   r   r   r   r   rJ   r   r   ZRDB_listr=   rL   r!   r)   F   s$    



"zRRDB.__init__c                 C   s@   t | dr*| |}| |}| |}n
| |}|d | S )NrM   皙?)hasattrrM   rN   rO   rR   )r7   rA   outr    r    r!   rC   [   s    



zRRDB.forward)r   r   r   r   r   r   Nr   r   r
   FFFrE   rF   rG   __doc__r)   rC   rH   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 )rK   a  
    Residual Dense Block
    The core module of paper: (Residual Dense Network for Image Super-Resolution, CVPR 18)
    Modified options that can be used:
        - "Partial Convolution based Padding" arXiv:1811.11718
        - "Spectral normalization" arXiv:1802.05957
        - "ICASSP 2020 - ESRGAN+ : Further Improving ESRGAN" N. C.
            {Rakotonirina} and A. {Rasoanaivo}
    @   r   r   r   r   Nr   r   r
   Fc                    s  t t|   |rt nd | _|r,t||nd | _t|||||||||	|
|d| _t|| ||||||||	|
|d| _t|d|  ||||||||	|
|d| _	t|d|  ||||||||	|
|d| _
|	dkrd }n|}t|d|  |d||||||	|
|d| _d S )N)r   r   r   r   r   r   rJ   r   r   r   r   )r(   rK   r)   GaussianNoisenoiseconv1x1r.   conv1conv2conv3conv4conv5)r7   r   r   r   r   r   r   r   r   r   r   rJ   r   r   Zlast_actr=   r    r!   r)   p   s4    zResidualDenseBlock_5C.__init__c              	   C   s   |  |}| t||fd}| jr4|| | }| t|||fd}| t||||fd}| jrt|| }| t|||||fd}| jr| |	d| S |d | S d S )Nr   rS   )
r\   r]   torchcatr[   r^   r_   r`   rZ   mul)r7   rA   x1x2x3x4x5r    r    r!   rC      s    
zResidualDenseBlock_5C.forward)rX   r   r   r   r   r   Nr   r   r
   FFFrV   r    r    r=   r!   rK   e   s   
   rK   c                       s&   e Zd Zd fdd	Zdd Z  ZS )rY   皙?Fc                    s,   t    || _|| _tjdtjd| _d S )Nr   )dtype)r(   r)   sigmais_relative_detachra   tensorfloatrZ   )r7   rk   rl   r=   r    r!   r)      s    
zGaussianNoise.__init__c                 C   sb   | j r^| jdkr^| j|j| _| jr4| j|  n| j| }| jj|  	 | }|| }|S )Nr   )
trainingrk   rZ   todevicerl   detachrepeatsizenormal_)r7   rA   r?   Zsampled_noiser    r    r!   rC      s    zGaussianNoise.forward)ri   FrD   r    r    r=   r!   rY      s   rY   r   c                 C   s   t j| |d|ddS )Nr   F)r   r   r   )rP   Conv2d)	in_planes
out_planesr   r    r    r!   r[      s    r[   c                       s*   e Zd ZdZd fdd	Zd	d
 Z  ZS )SRVGGNetCompactzA compact VGG-style network structure for super-resolution.
    This class is copied from https://github.com/xinntao/Real-ESRGAN
    r   rX   r   r   preluc           	   
      sN  t t|   || _|| _|| _|| _|| _|| _t	
 | _| jt	||ddd |dkrlt	jdd}n,|dkrt	j|d}n|dkrt	jd	dd
}| j| t|D ]n}| jt	||ddd |dkrt	jdd}n.|dkrt	j|d}n|dkrt	jd	dd
}| j| q| jt	||| | ddd t	|| _d S )Nr   r   reluTinplacerz   )num_parametersr   ri   )negative_sloper}   )r(   ry   r)   	num_in_ch
num_out_chnum_featnum_convr;   r   rP   
ModuleListbodyappendrv   ReLUPReLU	LeakyReLUr/   PixelShuffler<   )	r7   r   r   r   r   r;   r   
activationr   r=   r    r!   r)      s6    

"zSRVGGNetCompact.__init__c                 C   sP   |}t dt| jD ]}| j| |}q| |}tj|| jdd}||7 }|S )Nr   nearestscale_factorr   )r/   lenr   r<   Finterpolater;   )r7   rA   rU   ibaser    r    r!   rC      s    
zSRVGGNetCompact.forward)r   r   rX   r   r   rz   rV   r    r    r=   r!   ry      s   &ry   c                       s2   e Zd ZdZd
 fdd	Zdd Zdd	 Z  ZS )UpsamplezUpsamples a given multi-channel 1D (temporal), 2D (spatial) or 3D (volumetric) data.
    The input data is assumed to be of the form
    `minibatch x channels x [optional depth] x [optional height] x width`.
    Nr   c                    sV   t t|   t|tr.tdd |D | _n|r:t|nd | _|| _|| _|| _	d S )Nc                 s   s   | ]}t |V  qd S N)rn   )r   factorr    r    r!   	<genexpr>   r'   z$Upsample.__init__.<locals>.<genexpr>)
r(   r   r)   
isinstancetupler   rn   r   rt   align_corners)r7   rt   r   r   r   r=   r    r!   r)      s    
zUpsample.__init__c                 C   s   t jj|| j| j| j| jdS )N)rt   r   r   r   )rP   
functionalr   rt   r   r   r   )r7   rA   r    r    r!   rC     s    zUpsample.forwardc                 C   s8   | j d urd| j  }nd| j }|d| j 7 }|S )Nzscale_factor=zsize=z, mode=)r   rt   r   )r7   infor    r    r!   
extra_repr  s
    
zUpsample.extra_repr)NNr   N)rE   rF   rG   rW   r)   rC   r   rH   r    r    r=   r!   r      s   
r   c           
      C   s|   |   \}}}}||d  }|| dkr4|| dks8J || }|| }| ||||||}	|	dddddd||||S )z Pixel unshuffle.
    Args:
        x (Tensor): Input feature with shape (b, c, hh, hw).
        scale (int): Downsample ratio.
    Returns:
        Tensor: the pixel unshuffled feature.
    r   r   r   r      r   )rt   viewpermutereshape)
rA   r?   bchhhwout_channelhwx_viewr    r    r!   r@     s    r@   r   r   Tr   r{   r
   c
                 C   s\   t | ||d  ||||dd|	d	}
t|}|r:t||nd}|rJt|nd}t|
|||S )z
    Pixel shuffle layer
    (Real-Time Single Image and Video Super-Resolution Using an Efficient Sub-Pixel Convolutional
    Neural Network, CVPR17)
    r   Nr   r   r   r   r   )r.   rP   r   normr3   r4   )r8   r9   upscale_factorr   r   r   r   r   r   r   convpixel_shufflenar    r    r!   r1     s    
r1   r   c                 C   sF   |
dkrd||fn|}t ||	d}t| ||||||||
d	}t||S )z Upconv layer Conv3Dr   r   r   )r   r.   r4   )r8   r9   r   r   r   r   r   r   r   r   r   upsampler   r    r    r!   r0   /  s    r0   c                 K   s0   g }t |D ]}|| f i | qtj| S )a  Make layers by stacking the same blocks.
    Args:
        basic_block (nn.module): nn.module class for basic block. (block)
        num_basic_block (int): number of blocks. (n_layers)
    Returns:
        nn.Sequential: Stacked blocks in nn.Sequential.
    )r/   r   rP   rQ   )basic_blocknum_basic_blockkwarglayersr   r    r    r!   
make_layerD  s    r   rS         ?c                 C   s   |   } | dkrt|}nb| dv r2t||}nL| dkrJtj||d}n4| dkr\t }n"| dkrnt }ntd|  d|S )	z activation helper r{   )r   lrelurz   )r~   inittanhsigmoidzactivation layer [r$   )lowerrP   r   r   r   TanhSigmoidr2   )r   r}   Z	neg_slopeZn_prelubetalayerr    r    r!   r3   R  s    

r3   c                       s$   e Zd Z fddZdd Z  ZS )Identityc                    s   t t|   d S r   )r(   r   r)   )r7   kwargsr=   r    r!   r)   e  s    zIdentity.__init__c                 G   s   |S r   r    )r7   rA   r   r    r    r!   rC   h  s    zIdentity.forwardrD   r    r    r=   r!   r   d  s   r   c                 C   s^   |   } | dkr tj|dd}n:| dkr8tj|dd}n"| dkrJdd }ntd	|  d
|S )z Return a normalization layer batchT)affineinstanceFnonec                 S   s   t  S r   )r   )rA   r    r    r!   
norm_layert  r'   znorm.<locals>.norm_layerznormalization layer [r$   )r   rP   BatchNorm2dInstanceNorm2dr2   )r   ncr   r   r    r    r!   r   l  s    
r   c                 C   sd   |   } |dkrdS | dkr(t|}n8| dkr<t|}n$| dkrPt|}ntd|  d|S )z padding layer helper r   Nreflect	replicater   zpadding layer [z] is not implemented)r   rP   ReflectionPad2dReplicationPad2d	ZeroPad2dr2   )r   paddingr   r    r    r!   padz  s    r   c                 C   s$   | | d |d   } | d d }|S )Nr   r   r    )r   dilationr   r    r    r!   get_valid_padding  s    r   c                       s0   e Zd ZdZ fddZdd Zdd Z  ZS )r5   z8 Elementwise sum the output of a submodule to its input c                    s   t t|   || _d S r   )r(   r5   r)   sub)r7   	submoduler=   r    r!   r)     s    zShortcutBlock.__init__c                 C   s   ||  | }|S r   )r   )r7   rA   outputr    r    r!   rC     s    zShortcutBlock.forwardc                 C   s   d| j  dd S )NzIdentity + 
|
z
|)r   __repr__replace)r7   r    r    r!   r     s    zShortcutBlock.__repr__)rE   rF   rG   rW   r)   rC   r   rH   r    r    r=   r!   r5     s   r5   c                  G   s~   t | dkr*t| d tr"td| d S g }| D ]@}t|tjr\| D ]}|| qJq2t|tjr2|| q2tj| S )z/ Flatten Sequential. It unwraps nn.Sequential. r   r   z.sequential does not support OrderedDict input.)	r   r   r   r2   rP   rQ   childrenr   Module)argsmodulesmoduler   r    r    r!   r4     s    r4   r   Fc              
   C   s  |
dv sJ d|
 dt ||}|r8|dkr8t||nd}|dkrH|nd}|dkrzddlm} || |||||||d	}nl|d
krddlm} || |||||||d	}n>|dkrtj| |||||||d	}ntj| |||||||d	}|rtj	|}|	rt
|	nd}d|
v r4|r"t||nd}t||||S |
dkr|du r^|	dur^t
|	dd}|rnt|| nd}t||||S dS )z4 Conv layer with padding, normalization, activation )r   NACZCNACzWrong conv mode []r   Nr   ZPartialConv2D)PartialConv2d)r   r   r   r   r   groupsZDeformConv2D)DeformConv2dr   r   r   Fr|   )r   r   Ztorchvision.opsr   r   rP   Conv3drv   utilsrJ   r3   r   r4   )r8   r9   r   r   r   r   r   r   r   r   r   r   rJ   r   pr   r   r   r   r   r    r    r!   r.     s@    


r.   )r   )r   r   r   Tr   Nr{   r
   )	r   r   r   Tr   Nr{   r   r
   )TrS   r   r   )
r   r   r   Tr   Nr{   r   r
   F)collectionsr   r+   ra   torch.nnrP   Ztorch.nn.functionalr   r   r   r   r   rK   rY   r[   ry   r   r@   r1   r0   r   r3   r   r   r   r   r5   r4   r.   r    r    r    r!   <module>   s<   2%;
;  
  

   