a
    d(                     @   s   d dl Zd dlZd dlmZ d dlm  mZ d dlm	Z	 d dl
mZ ddlmZmZmZmZmZ ddlmZ G dd dejZe G d	d
 d
ejZdS )    N)spectral_norm)ARCH_REGISTRY   )AttentionBlockBlurMSDilationBlock
UpResBlockadaptive_instance_normalization)VGGFeatureExtractorc                       s*   e Zd ZdZd fdd	Zdd Z  ZS )	
SFTUpBlocka0  Spatial feature transform (SFT) with upsampling block.

    Args:
        in_channel (int): Number of input channels.
        out_channel (int): Number of output channels.
        kernel_size (int): Kernel size in convolutions. Default: 3.
        padding (int): Padding in convolutions. Default: 1.
       r   c                    s   t t|   tt|ttj||||dtdd| _	ttj
ddddttj||||dtdd| _ttt||d	d
d
tddtt||d	d
d
| _ttt||d	d
d
tddtt||d	d
d
t | _d S )N)paddingg{Gz?T   bilinearF)Zscale_factormodealign_corners皙?r   r   )superr   __init__nn
Sequentialr   r   Conv2d	LeakyReLUconv1ZUpsampleconvupscale_blockZSigmoidshift_block)selfZ
in_channelZout_channelkernel_sizer   	__class__ b/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/basicsr/archs/dfdnet_arch.pyr      s&    

zSFTUpBlock.__init__c                 C   s8   |  |}| |}| |}|| | }| |}|S )N)r   r   r   r   )r   xupdated_featoutZscaleshiftr!   r!   r"   forward-   s    



zSFTUpBlock.forward)r   r   )__name__
__module____qualname____doc__r   r'   __classcell__r!   r!   r   r"   r      s   	r   c                       s8   e Zd ZdZ fddZdd Zdd Zdd	 Z  ZS )
DFDNetzDFDNet: Deep Face Dictionary Network.

    It only processes faces with 512x512 size.

    Args:
        num_feat (int): Number of feature channels.
        dict_path (str): Path to the facial component dictionary.
    c                    sV  t    g d| _g d}tg d| _g d| _d| _t	|| _
t| jddddd| _t | _t| jD ]0\}}| jD ] }t|| | j| d	| < qqrt|d
 g dd| _t|d
 |d
 | _t|d
 |d | _t|d |d | _t|d || _ttt||dddtddt|t|tj|dddddt | _d S )N)Zleft_eyeZ	right_eyeZnoseZmouth)         r0   )r/   r.   @       )Zrelu2_2Zrelu3_4Zrelu4_4conv5_4FZvgg19T)Zlayer_name_listZvgg_typeZuse_input_normZ
range_normZrequires_grad_   )   r   r   r   )Zdilationr6   r   r   r   r   )r   Zstrider   ) r   r   partsnparrayfeature_sizes
vgg_layersflag_dict_devicetorchloaddictr
   vgg_extractorr   Z
ModuleDictattn_blocks	enumerater   r   multi_scale_dilationr   	upsample0	upsample1	upsample2	upsample3r   r   r   r   r   ZTanh	upsample4)r   Znum_featZ	dict_pathZchannel_sizesidxZ	feat_sizenamer   r!   r"   r   C   s6    




 $zDFDNet.__init__c                 C   s
  |dddd|d |d |d |d f   }tj|| dd ddd	}t||}t||}	tj|	d
dd}	t	|	}
t||
|
d  | dd }| j
| dt|  || }|| }|| |dddd|d |d |d |d f< |S )z&swap the features from the dictionary.Nr   r   r   r   r6   r   F)r   r   )Zdimr4   )cloneFZinterpolatesizer	   Zconv2dZsoftmaxviewr=   ZargmaxrA   str)r   vgg_featr$   Z	dict_featlocation	part_namef_sizeZ	part_featZpart_resize_featZsimilarity_scoreZ
select_idx	swap_featZattnZ	attn_featr!   r!   r"   rU   i   s    4

$ 4zDFDNet.swap_featc                 C   sN   | j du rJ| j D ].\}}| D ]\}}||| j| |< q$qd| _ d S )NFT)r<   r?   itemsto)r   r#   kvkkvvr!   r!   r"   put_dict_to_device~   s
    
zDFDNet.put_dict_to_devicec              
   C   s   |  | | |}g }d}t| j| jD ]r\}}| j|  }|| }	|	 }
t| jD ]8\}}|| | d|  	 }| 
|	|
|| |||}
qX||
 q*| |d }| ||d }| ||d }| ||d }| ||d }| |}|S )z
        Now only support testing with batch size = 0.

        Args:
            x (Tensor): Input faces with shape (b, c, 512, 512).
            part_locations (list[Tensor]): Part locations.
        r   r0   r3   r   r   r   )r\   r@   zipr;   r:   r?   rL   rB   r7   intrU   appendrC   rD   rE   rF   rG   rH   )r   r#   Zpart_locationsZvgg_featuresZupdated_vgg_featuresbatchZ	vgg_layerrT   Zdict_featuresrQ   r$   Zpart_idxrS   rR   Zvgg_feat_dilationZupsampled_featr%   r!   r!   r"   r'      s*    


zDFDNet.forward)	r(   r)   r*   r+   r   rU   r\   r'   r,   r!   r!   r   r"   r-   8   s
   	&r-   )numpyr8   r=   Ztorch.nnr   Ztorch.nn.functionalZ
functionalrM   Ztorch.nn.utils.spectral_normr   Zbasicsr.utils.registryr   Zdfdnet_utilr   r   r   r   r	   Zvgg_archr
   Moduler   registerr-   r!   r!   r!   r"   <module>   s   ,