a
    þd@  ã                   @   s’   d dl Z d dlZd dlmZ d dlmZ d dlmZ d dlmZ e 	¡ G dd„ dej
ƒƒZe 	¡ G dd	„ d	eƒƒZd
d„ Zddd„Zddd„ZdS )é    N)Úautograd)Únn)Ú
functional)ÚLOSS_REGISTRYc                       sD   e Zd ZdZd‡ fdd„	Zdd„ Zdd	„ Zd
d„ Zddd„Z‡  Z	S )ÚGANLossaœ  Define GAN loss.

    Args:
        gan_type (str): Support 'vanilla', 'lsgan', 'wgan', 'hinge'.
        real_label_val (float): The value for real label. Default: 1.0.
        fake_label_val (float): The value for fake label. Default: 0.0.
        loss_weight (float): Loss weight. Default: 1.0.
            Note that loss_weight is only for generators; and it is always 1.0
            for discriminators.
    ç      ð?ç        c                    s¦   t t| ƒ ¡  || _|| _|| _|| _| jdkr<t ¡ | _	nf| jdkrRt 
¡ | _	nP| jdkrf| j| _	n<| jdkrz| j| _	n(| jdkrt ¡ | _	ntd| j› dƒ‚d S )NZvanillaZlsganÚwganÚwgan_softplusÚhingez	GAN type z is not implemented.)Úsuperr   Ú__init__Úgan_typeÚloss_weightÚreal_label_valÚfake_label_valr   ZBCEWithLogitsLossÚlossZMSELossÚ
_wgan_lossÚ_wgan_softplus_lossZReLUÚNotImplementedError©Úselfr   r   r   r   ©Ú	__class__© ú`/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/basicsr/losses/gan_loss.pyr      s     






zGANLoss.__init__c                 C   s   |r|  ¡  S |  ¡ S )z¦wgan loss.

        Args:
            input (Tensor): Input tensor.
            target (bool): Target label.

        Returns:
            Tensor: wgan loss.
        )Úmean©r   ÚinputÚtargetr   r   r   r   +   s    
zGANLoss._wgan_lossc                 C   s"   |rt  | ¡ ¡ S t  |¡ ¡ S )aw  wgan loss with soft plus. softplus is a smooth approximation to the
        ReLU function.

        In StyleGAN2, it is called:
            Logistic loss for discriminator;
            Non-saturating loss for generator.

        Args:
            input (Tensor): Input tensor.
            target (bool): Target label.

        Returns:
            Tensor: wgan loss.
        )ÚFZsoftplusr   r   r   r   r   r   7   s    zGANLoss._wgan_softplus_lossc                 C   s0   | j dv r|S |r| jn| j}| | ¡ ¡| S )a  Get target label.

        Args:
            input (Tensor): Input tensor.
            target_is_real (bool): Whether the target is real or fake.

        Returns:
            (bool | Tensor): Target tensor. Return bool for wgan, otherwise,
                return Tensor.
        )r	   r
   )r   r   r   Znew_onesÚsize)r   r   Útarget_is_realZ
target_valr   r   r   Úget_target_labelH   s    
zGANLoss.get_target_labelFc                 C   sf   |   ||¡}| jdkrH|r<|r$| n|}|  d| ¡ ¡ }qT| ¡  }n|  ||¡}|r\|S || j S )ad  
        Args:
            input (Tensor): The input for the loss module, i.e., the network
                prediction.
            target_is_real (bool): Whether the targe is real or fake.
            is_disc (bool): Whether the loss for discriminators or not.
                Default: False.

        Returns:
            Tensor: GAN loss value.
        r   é   )r#   r   r   r   r   )r   r   r"   Úis_discZtarget_labelr   r   r   r   ÚforwardY   s    
zGANLoss.forward)r   r   r   )F)
Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r   r#   r&   Ú__classcell__r   r   r   r   r   
   s   r   c                       s0   e Zd ZdZd	‡ fdd„	Zd
‡ fdd„	Z‡  ZS )ÚMultiScaleGANLossz9
    MultiScaleGANLoss accepts a list of predictions
    r   r   c                    s   t t| ƒ ||||¡ d S )N)r   r,   r   r   r   r   r   r   y   s    zMultiScaleGANLoss.__init__Fc                    sf   t |tƒrRd}|D ]2}t |tƒr(|d }tƒ  |||¡ ¡ }||7 }q|t|ƒ S tƒ  |||¡S dS )zR
        The input is a list of tensors, or a list of (a list of tensors)
        r   éÿÿÿÿN)Ú
isinstanceÚlistr   r&   r   Úlen)r   r   r"   r%   r   Zpred_iZloss_tensorr   r   r   r&   |   s    


zMultiScaleGANLoss.forward)r   r   r   )F)r'   r(   r)   r*   r   r&   r+   r   r   r   r   r,   s   s   r,   c                 C   s>   t j|  ¡ |ddd }| d¡ |jd d¡ d¡ ¡ }|S )a  R1 regularization for discriminator. The core idea is to
        penalize the gradient on real data alone: when the
        generator distribution produces the true data distribution
        and the discriminator is equal to 0 on the data manifold, the
        gradient penalty ensures that the discriminator cannot create
        a non-zero gradient orthogonal to the data manifold without
        suffering a loss in the GAN game.

        Ref:
        Eq. 9 in Which training methods for GANs do actually converge.
        T©ÚoutputsÚinputsÚcreate_graphr   é   r-   r$   )r   ÚgradÚsumÚpowÚviewÚshaper   )Z	real_predZreal_imgZ	grad_realZgrad_penaltyr   r   r   Ú
r1_penalty   s    "r;   ç{®Gáz„?c           	      C   s˜   t  | ¡t | jd | jd  ¡ }tj| |  ¡ |ddd }t  | d¡ d¡ 	d¡¡}||| 	¡ |   }||  d¡ 	¡ }|| 
¡  	¡ | 
¡ fS )Nr5   é   Tr1   r   r$   )ÚtorchZ
randn_likeÚmathÚsqrtr:   r   r6   r7   r8   r   Údetach)	Zfake_imgZlatentsZmean_path_lengthZdecayZnoiser6   Zpath_lengthsZ	path_meanZpath_penaltyr   r   r   Úg_path_regularize    s    $rB   c           
      C   s®   |  d¡}| t |ddd¡¡}|| d| |  }tj|dd}| |ƒ}tj||t |¡ddddd }|durz|| }|jddd	d d  	¡ }	|durª|	t 	|¡ }	|	S )
aG  Calculate gradient penalty for wgan-gp.

    Args:
        discriminator (nn.Module): Network for the discriminator.
        real_data (Tensor): Real input data.
        fake_data (Tensor): Fake input data.
        weight (Tensor): Weight tensor. Default: None.

    Returns:
        Tensor: A tensor for gradient penalty.
    r   r$   r   T)Zrequires_grad)r2   r3   Zgrad_outputsr4   Zretain_graphZonly_inputsNr5   )Zdim)
r!   Z
new_tensorr>   Zrandr   ÚVariabler6   Z	ones_likeZnormr   )
ZdiscriminatorZ	real_dataZ	fake_dataZweightZ
batch_sizeÚalphaZinterpolatesZdisc_interpolatesZ	gradientsZgradients_penaltyr   r   r   Úgradient_penalty_loss¬   s*    
úúrE   )r<   )N)r?   r>   r   r   Ztorch.nnr   r    Zbasicsr.utils.registryr   ÚregisterÚModuler   r,   r;   rB   rE   r   r   r   r   Ú<module>   s   h
