a
    d                     @   sL   d Z ddlmZ ddlZddlmZ ddlm  mZ G dd dej	Z
dS )zY Binary Cross Entropy w/ a few extras

Hacked together by / Copyright 2021 Ross Wightman
    )OptionalNc                       sV   e Zd ZdZdee eej eeej d fddZ	ejejejdd	d
Z
  ZS )BinaryCrossEntropyz BCE with optional one-hot from dense targets, label smoothing, thresholding
    NOTE for experiments comparing CE to BCE /w label smoothing, may remove
    皙?Nmean)target_thresholdweight	reduction
pos_weightc                    sV   t t|   d|  kr"dk s(n J || _|| _|| _| d| | d| d S )Ng              ?r   r	   )superr   __init__	smoothingr   r   Zregister_buffer)selfr   r   r   r   r	   	__class__ g/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/loss/binary_cross_entropy.pyr      s    zBinaryCrossEntropy.__init__)xtargetreturnc                 C   s   |j d |j d ksJ |j |j kr|j d }| j| }d| j | }| dd}tj| d |f||j|jd	d||}| j
d ur|| j
j|jd}tj||| j| j| jdS )Nr   r
      )devicedtype)r   )r	   r   )shaper   longviewtorchfullsizer   r   Zscatter_r   gttoFZ binary_cross_entropy_with_logitsr   r	   r   )r   r   r   Znum_classesZ	off_valueZon_valuer   r   r   forward   s*    


zBinaryCrossEntropy.forward)r   NNr   N)__name__
__module____qualname____doc__r   floatr   ZTensorstrr   r#   __classcell__r   r   r   r   r      s     
r   )r'   typingr   r   Ztorch.nnnnZtorch.nn.functionalZ
functionalr"   Moduler   r   r   r   r   <module>   s
   