a
    d;                     @   sH   d dl Z d dlmZ d dlm  mZ ddlmZ G dd dejZ	dS )    N   )LabelSmoothingCrossEntropyc                       s*   e Zd ZdZd	 fdd	Zdd Z  ZS )
JsdCrossEntropyaL   Jensen-Shannon Divergence + Cross-Entropy Loss

    Based on impl here: https://github.com/google-research/augmix/blob/master/imagenet.py
    From paper: 'AugMix: A Simple Data Processing Method to Improve Robustness and Uncertainty -
    https://arxiv.org/abs/1912.02781

    Hacked together by / Copyright 2020 Ross Wightman
          皙?c                    sB   t    || _|| _|d ur2|dkr2t|| _ntj | _d S )Nr   )	super__init__
num_splitsalphar   cross_entropy_losstorchnnZCrossEntropyLoss)selfr
   r   Z	smoothing	__class__ V/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/loss/jsd.pyr	      s    
zJsdCrossEntropy.__init__c                    s   |j d | j }|| j |j d ks(J t||}| |d |d | }dd |D }tt|jdddd  || j	t
 fdd|D  t| 7 }|S )Nr   c                 S   s   g | ]}t j|d dqS )r   )Zdim)FZsoftmax).0Zlogitsr   r   r   
<listcomp>!       z,JsdCrossEntropy.__call__.<locals>.<listcomp>)ZaxisgHz>r   c                    s   g | ]}t j |d dqS )Z	batchmean)Z	reduction)r   Zkl_div)r   Zp_splitZlogp_mixturer   r   r   %   s   )shaper
   r   splitr   clampstackmeanlogr   sumlen)r   outputtargetZ
split_sizeZlogits_splitZlossZprobsr   r   r   __call__   s     zJsdCrossEntropy.__call__)r   r   r   )__name__
__module____qualname____doc__r	   r#   __classcell__r   r   r   r   r      s   	r   )
r   Ztorch.nnr   Ztorch.nn.functionalZ
functionalr   Zcross_entropyr   Moduler   r   r   r   r   <module>   s   