a
    þdt?  ã                   @   s”   d dl Z d dl mZ d dlmZ d dlmZ ddlmZ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e ¡ G dd„ dejƒƒZdS )é    N)Únn)Ú
functional)ÚARCH_REGISTRYé   )Ú	DCNv2PackÚResidualBlockNoBNÚ
make_layerc                       s*   e Zd ZdZd‡ fdd„	Zdd„ Z‡  ZS )	ÚPCDAlignmentaV  Alignment module using Pyramid, Cascading and Deformable convolution
    (PCD). It is used in EDVR.

    Ref:
        EDVR: Video Restoration with Enhanced Deformable Convolutional Networks

    Args:
        num_feat (int): Channel number of middle features. Default: 64.
        deformable_groups (int): Deformable groups. Defaults: 8.
    é@   é   c                    sp  t t| ƒ ¡  t ¡ | _t ¡ | _t ¡ | _t ¡ | _t ¡ | _	t
dddƒD ]¼}d|› }t |d |ddd¡| j|< |dkr˜t ||ddd¡| j|< n4t |d |ddd¡| j|< t ||ddd¡| j|< t||dd|d| j|< |dk rLt |d |ddd¡| j	|< qLt |d |ddd¡| _t ||ddd¡| _t||dd|d| _tjddd	d
| _tjddd| _d S )Né   r   éÿÿÿÿÚlé   r   )ÚpaddingÚdeformable_groupsÚbilinearF©Zscale_factorÚmodeZalign_cornersçš™™™™™¹?T©Znegative_slopeZinplace)Úsuperr	   Ú__init__r   Z
ModuleDictÚoffset_conv1Úoffset_conv2Úoffset_conv3Údcn_packÚ	feat_convÚrangeÚConv2dr   Úcas_offset_conv1Úcas_offset_conv2Úcas_dcnpackÚUpsampleÚupsampleÚ	LeakyReLUÚlrelu)ÚselfÚnum_featr   ÚiÚlevel©Ú	__class__© ú`/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/basicsr/archs/edvr_arch.pyr      s*    





zPCDAlignment.__init__c           	   	   C   sf  d\}}t dddƒD ]}d|› }tj||d  ||d  gdd}|  | j| |ƒ¡}|dkrx|  | j| |ƒ¡}n6|  | j| tj||gddƒ¡}|  | j| |ƒ¡}| j| ||d  |ƒ}|dk rê| j| tj||gddƒ}|dkrü|  |¡}|dkr|  	|¡d }|  	|¡}qtj||d gdd}|  |  
|  |  |¡¡¡¡}|  |  ||¡¡}|S )	aë  Align neighboring frame features to the reference frame features.

        Args:
            nbr_feat_l (list[Tensor]): Neighboring feature list. It
                contains three pyramid levels (L1, L2, L3),
                each with shape (b, c, h, w).
            ref_feat_l (list[Tensor]): Reference feature list. It
                contains three pyramid levels (L1, L2, L3),
                each with shape (b, c, h, w).

        Returns:
            Tensor: Aligned features.
        )NNr   r   r   r   r   ©Zdimr   )r   ÚtorchÚcatr&   r   r   r   r   r   r$   r!   r    r"   )	r'   Ú
nbr_feat_lÚ
ref_feat_lZupsampled_offsetZupsampled_featr)   r*   ÚoffsetÚfeatr-   r-   r.   Úforward8   s*    
""
zPCDAlignment.forward)r
   r   ©Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r6   Ú__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 )
Ú	TSAFusiona  Temporal Spatial Attention (TSA) fusion module.

    Temporal: Calculate the correlation between center frame and
        neighboring frames;
    Spatial: It has 3 pyramid levels, the attention is similar to SFT.
        (SFT: Recovering realistic texture in image super-resolution by deep
            spatial feature transform.)

    Args:
        num_feat (int): Channel number of middle features. Default: 64.
        num_frame (int): Number of frames. Default: 5.
        center_frame_idx (int): The index of center frame. Default: 2.
    r
   é   r   c                    sX  t t| ƒ ¡  || _t ||ddd¡| _t ||ddd¡| _t || |dd¡| _tj	dddd| _
tjdddd| _t || |d¡| _t |d |d¡| _t ||ddd¡| _t ||d¡| _t ||ddd¡| _t ||d¡| _t |d |ddd¡| _t ||ddd¡| _t ||d¡| _t ||d¡| _tjddd| _tjddd	d
| _d S )Nr   r   r   )Zstrider   r   Tr   r   Fr   )r   r=   r   Úcenter_frame_idxr   r   Útemporal_attn1Útemporal_attn2Úfeat_fusionZ	MaxPool2dÚmax_poolZ	AvgPool2dÚavg_poolÚspatial_attn1Úspatial_attn2Úspatial_attn3Úspatial_attn4Úspatial_attn5Úspatial_attn_l1Úspatial_attn_l2Úspatial_attn_l3Úspatial_attn_add1Úspatial_attn_add2r%   r&   r#   r$   )r'   r(   Ú	num_framer?   r+   r-   r.   r   t   s&    zTSAFusion.__init__c              	   C   s*  |  ¡ \}}}}}|  |dd…| jdd…dd…dd…f  ¡ ¡}|  | d|||¡¡}| ||d||¡}g }	t|ƒD ]F}
|dd…|
dd…dd…dd…f }t || d¡}|	 	| 
d¡¡ qtt tj|	dd¡}| 
d¡ |||||¡}| ¡  |d||¡}| |d||¡| }|  |  |¡¡}|  |  |¡¡}|  |¡}|  |¡}|  |  tj||gdd¡¡}|  |  |¡¡}|  |¡}|  |¡}|  |  tj||gdd¡¡}|  |  |¡¡}|  |¡}|  |  |¡¡| }|  |  |¡¡}|  |¡}|  |¡}|  |  |  |¡¡¡}t |¡}|| d | }|S )z½
        Args:
            aligned_feat (Tensor): Aligned features with shape (b, t, c, h, w).

        Returns:
            Tensor: Features after TSA with the shape (b, c, h, w).
        Nr   r   r/   r   )Úsizer@   r?   ÚclonerA   Úviewr   r0   ÚsumÚappendZ	unsqueezeZsigmoidr1   ÚexpandÚ
contiguousr&   rB   rE   rC   rD   rF   rJ   rK   rL   r$   rG   rH   rI   rN   rM   )r'   Úaligned_featÚbÚtÚcÚhÚwZembedding_refZ	embeddingZcorr_lr)   Zemb_neighborZcorrZ	corr_probr5   ZattnZattn_maxZattn_avgZ
attn_levelZattn_addr-   r-   r.   r6      s@    ."







zTSAFusion.forward)r
   r>   r   r7   r-   r-   r+   r.   r=   e   s   r=   c                       s*   e Zd ZdZd	‡ fdd„	Zdd„ Z‡  ZS )
ÚPredeblurModulea  Pre-dublur module.

    Args:
        num_in_ch (int): Channel number of input image. Default: 3.
        num_feat (int): Channel number of intermediate features. Default: 64.
        hr_in (bool): Whether the input has high resolution. Default: False.
    r   r
   Fc                    sæ   t t| ƒ ¡  || _t |ˆ ddd¡| _| jrVt ˆ ˆ ddd¡| _t ˆ ˆ ddd¡| _t ˆ ˆ ddd¡| _	t ˆ ˆ ddd¡| _
tˆ d| _tˆ d| _tˆ d| _t ‡ fdd„tdƒD ƒ¡| _tjddd	d
| _tjddd| _d S )Nr   r   r   ©r(   c                    s   g | ]}t ˆ d ‘qS )r^   )r   )Ú.0r)   r^   r-   r.   Ú
<listcomp>Û   ó    z,PredeblurModule.__init__.<locals>.<listcomp>r>   r   Fr   r   Tr   )r   r]   r   Úhr_inr   r   Ú
conv_firstÚstride_conv_hr1Ústride_conv_hr2Ústride_conv_l2Ústride_conv_l3r   Úresblock_l3Úresblock_l2_1Úresblock_l2_2Z
ModuleListr   Úresblock_l1r#   r$   r%   r&   )r'   Ú	num_in_chr(   rb   r+   r^   r.   r   Ê   s    zPredeblurModule.__init__c                 C   sÊ   |   |  |¡¡}| jr6|   |  |¡¡}|   |  |¡¡}|   |  |¡¡}|   |  |¡¡}|  |  |¡¡}|  	|¡| }|  |  
|¡¡}tdƒD ]}| j| |ƒ}qŒ|| }tddƒD ]}| j| |ƒ}q²|S )Nr   r>   )r&   rc   rb   rd   re   rf   rg   r$   rh   ri   rj   r   rk   )r'   ÚxÚfeat_l1Úfeat_l2Úfeat_l3r)   r-   r-   r.   r6   à   s    zPredeblurModule.forward)r   r
   Fr7   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 )ÚEDVRaI  EDVR network structure for video super-resolution.

    Now only support X4 upsampling factor.
    Paper:
        EDVR: Video Restoration with Enhanced Deformable Convolutional Networks

    Args:
        num_in_ch (int): Channel number of input image. Default: 3.
        num_out_ch (int): Channel number of output image. Default: 3.
        num_feat (int): Channel number of intermediate features. Default: 64.
        num_frame (int): Number of input frames. Default: 5.
        deformable_groups (int): Deformable groups. Defaults: 8.
        num_extract_block (int): Number of blocks for feature extraction.
            Default: 5.
        num_reconstruct_block (int): Number of blocks for reconstruction.
            Default: 10.
        center_frame_idx (int): The index of center frame. Frame counting from
            0. Default: Middle of input frames.
        hr_in (bool): Whether the input has high resolution. Default: False.
        with_predeblur (bool): Whether has predeblur module.
            Default: False.
        with_tsa (bool): Whether has TSA module. Default: True.
    r   r
   r>   r   é
   NFTc                    sœ  t t| ƒ ¡  |d u r"|d | _n|| _|	| _|
| _|| _| jrdt|| jd| _t	 
||dd¡| _nt	 
||ddd¡| _tt||d| _t	 
||ddd¡| _t	 
||ddd¡| _t	 
||ddd¡| _t	 
||ddd¡| _t||d| _| jrt||| jd| _nt	 
|| |dd¡| _tt||d| _t	 
||d ddd¡| _t	 
|d	ddd¡| _t	 d¡| _t	 
d
d
ddd¡| _t	 
d
dddd¡| _t	jddd| _ d S )Nr   )r(   rb   r   r   r^   )r(   r   )r(   rO   r?   é   é   r
   r   Tr   )!r   rq   r   r?   rb   Úwith_predeblurÚwith_tsar]   Ú	predeblurr   r   Úconv_1x1rc   r   r   Úfeature_extractionÚ	conv_l2_1Ú	conv_l2_2Ú	conv_l3_1Ú	conv_l3_2r	   Ú	pcd_alignr=   ÚfusionÚreconstructionÚupconv1Úupconv2ZPixelShuffleÚpixel_shuffleÚconv_hrÚ	conv_lastr%   r&   )r'   rl   Z
num_out_chr(   rO   r   Znum_extract_blockZnum_reconstruct_blockr?   rb   ru   rv   r+   r-   r.   r     s6    zEDVR.__init__c              
   C   s"  |  ¡ \}}}}}| jr:|d dkr0|d dksZJ dƒ‚n |d dkrR|d dksZJ dƒ‚|d d …| jd d …d d …d d …f  ¡ }| jr¾|  |  | d|||¡¡¡}| jrÚ|d |d  }}n|  |  	| d|||¡¡¡}|  
|¡}|  |  |¡¡}	|  |  |	¡¡}	|  |  |	¡¡}
|  |  |
¡¡}
| ||d||¡}|	 ||d|d |d ¡}	|
 ||d|d |d ¡}
|d d …| jd d …d d …d d …f  ¡ |	d d …| jd d …d d …d d …f  ¡ |
d d …| jd d …d d …d d …f  ¡ g}g }t|ƒD ]ˆ}|d d …|d d …d d …d d …f  ¡ |	d d …|d d …d d …d d …f  ¡ |
d d …|d d …d d …d d …f  ¡ g}| |  ||¡¡ qìtj|dd	}| jsœ| |d||¡}|  |¡}|  |¡}|  |  |  |¡¡¡}|  |  |  |¡¡¡}|  |  |¡¡}|  |¡}| jr|}ntj|dd
dd}||7 }|S )Né   r   z,The height and width must be multiple of 16.rs   z+The height and width must be multiple of 4.r   r   r   r/   r   Fr   )rP   rb   r?   rV   ru   rx   rw   rR   r&   rc   ry   rz   r{   r|   r}   rQ   r   rT   r~   r0   Ústackrv   r   r€   rƒ   r   r‚   r„   r…   ÚFZinterpolate)r'   rm   rX   rY   rZ   r[   r\   Zx_centerrn   ro   rp   r3   rW   r)   r2   r5   ÚoutÚbaser-   r-   r.   r6   F  sP    " (
L&þlÿ


zEDVR.forward)r   r   r
   r>   r   r>   rr   NFFTr7   r-   r-   r+   r.   rq   ö   s              õ6rq   )r0   r   Ztorch.nnr   rˆ   Zbasicsr.utils.registryr   Z	arch_utilr   r   r   ÚModuler	   r=   r]   Úregisterrq   r-   r-   r-   r.   Ú<module>   s   \\5