a
    d09                     @   sh   d Z ddlZddlZdddZdddZdd	d
ZdddZdddZG dd dZ	G dd de	Z
dS )aY   Mixup and Cutmix

Papers:
mixup: Beyond Empirical Risk Minimization (https://arxiv.org/abs/1710.09412)

CutMix: Regularization Strategy to Train Strong Classifiers with Localizable Features (https://arxiv.org/abs/1905.04899)

Code Reference:
CutMix: https://github.com/clovaai/CutMix-PyTorch

Hacked together by / Copyright 2019, Ross Wightman
    N      ?        c                 C   s8   |   dd} tj|  d |f|| jdd| |S )N   r   )device)longviewtorchfullsizer   Zscatter_)xnum_classeson_value	off_value r   X/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/data/mixup.pyone_hot   s    r   c                 C   sN   || }d| | }t | |||d}t | d|||d}|| |d|   S )Nr   )r   r   r   )r   flip)targetr   lamZ	smoothingr   r   y1y2r   r   r   mixup_target   s
    r   c                 C   s   t d| }| dd \}}t|| t||  }}t|| t||  }	}
t jjd|	 ||	 |d}t jjd|
 ||
 |d}t ||d  d|}t ||d  d|}t ||d  d|}t ||d  d|}||||fS )a   Standard CutMix bounding-box
    Generates a random square bbox based on lambda value. This impl includes
    support for enforcing a border margin as percent of bbox dimensions.

    Args:
        img_shape (tuple): Image shape as tuple
        lam (float): Cutmix lambda value
        margin (float): Percentage of bbox dimension to enforce as margin (reduce amount of box outside image)
        count (int): Number of bbox to generate
    r   Nr   r      )npsqrtintrandomrandintZclip)	img_shaper   margincountZratioimg_himg_wcut_hcut_wZmargin_yZmargin_xcyZcxylyhxlxhr   r   r   	rand_bbox   s    r-   c                 C   s   t |dksJ | dd \}}tjjt||d  t||d  |d}tjjt||d  t||d  |d}tjjd|| |d}tjjd|| |d}|| }	|| }
||	||
fS )a   Min-Max CutMix bounding-box
    Inspired by Darknet cutmix impl, generates a random rectangular bbox
    based on min/max percent values applied to each dimension of the input image.

    Typical defaults for minmax are usually in the  .2-.3 for min and .8-.9 range for max.

    Args:
        img_shape (tuple): Image shape as tuple
        minmax (tuple or list): Min and max bbox ratios (as percent of image size)
        count (int): Number of bbox to generate
    r   r   Nr   r   r   )lenr   r   r    r   )r!   Zminmaxr#   r$   r%   r&   r'   r)   r+   yuxur   r   r   rand_bbox_minmax6   s    **r1   Tc           
      C   s~   |dur t | ||d\}}}}nt| ||d\}}}}|sB|durn|| ||  }	d|	t| d | d    }||||f|fS )z0 Generate bbox and apply lambda correction.
    N)r#   r   r   r   )r1   r-   float)
r!   r   ratio_minmaxcorrect_lamr#   r)   r/   r+   r0   Z	bbox_arear   r   r   cutmix_bbox_and_lamM   s    r5   c                	   @   sJ   e Zd ZdZdd
dZdd Zdd Zdd Zdd Zdd Z	dd Z
dS )Mixupas   Mixup/Cutmix that applies different params to each element or whole batch

    Args:
        mixup_alpha (float): mixup alpha value, mixup is active if > 0.
        cutmix_alpha (float): cutmix alpha value, cutmix is active if > 0.
        cutmix_minmax (List[float]): cutmix min/max image ratio, cutmix is active and uses this vs alpha if not None.
        prob (float): probability of applying mixup or cutmix per batch or element
        switch_prob (float): probability of switching to cutmix instead of mixup when both are active
        mode (str): how to apply mixup/cutmix params (per 'batch', 'pair' (pair of elements), 'elem' (element)
        correct_lam (bool): apply lambda correction when cutmix bbox clipped by image borders
        label_smoothing (float): apply label smoothing to the mixed target tensor
        num_classes (int): number of classes for target
    r   r   N      ?batchT皙?  c
           
      C   sb   || _ || _|| _| jd ur4t| jdks.J d| _|| _|| _|| _|	| _|| _|| _	d| _
d S )Nr   r   T)mixup_alphacutmix_alphacutmix_minmaxr.   mix_probswitch_problabel_smoothingr   moder4   mixup_enabled)
selfr;   r<   r=   Zprobr?   rA   r4   r@   r   r   r   r   __init__h   s    
zMixup.__init__c              	   C   s  t j|t jd}t j|t jd}| jr| jdkr| jdkrt j	|| j
k }t |t jj| j| j|dt jj| j| j|d}n`| jdkrt jj| j| j|d}n>| jdkrt j|t jd}t jj| j| j|d}ndsJ dt t j	|| jk |t j|}||fS )Ndtyper   r   FROne of mixup_alpha > 0., cutmix_alpha > 0., cutmix_minmax not None should be true.)r   onesfloat32zerosboolrB   r;   r<   r   randr?   wherebetar>   astype)rC   
batch_sizer   
use_cutmixlam_mixr   r   r   _params_per_elemy   s$    

$zMixup._params_per_elemc                 C   s   d}d}| j rtj | jk r| jdkrl| jdkrltj | jk }|rXtj| j| jntj| j| j}nL| jdkrtj| j| j}n.| jdkrd}tj| j| j}ndsJ dt	|}||fS )Nr   Fr   TrG   )
rB   r   r   rL   r>   r;   r<   r?   rN   r2   )rC   r   rQ   rR   r   r   r   _params_per_batch   s     

zMixup._params_per_batchc                 C   s   t |}| |\}}| }t|D ]}|| d }|| }|dkr&|| rt|| j|| j| jd\\}	}
}}}|| d d |	|
||f || d d |	|
||f< |||< q&|| | || d|   ||< q&tj	||j
|jddS )Nr   r   r3   r4   r   rF   )r.   rS   cloneranger5   shaper=   r4   r	   tensorr   rF   	unsqueezerC   r   rP   	lam_batchrQ   Zx_origijr   r)   r*   r+   r,   r   r   r   	_mix_elem   s    8
"zMixup._mix_elemc                 C   sl  t |}| |d \}}| }t|d D ]}|| d }|| }|dkr.|| rt|| j|| j| jd\\}	}
}}}|| d d |	|
||f || d d |	|
||f< || d d |	|
||f || d d |	|
||f< |||< q.|| | || d|   ||< || | || d|   ||< q.t	||d d d f}t
j||j|jddS )Nr   r   r   rU   r   rV   )r.   rS   rW   rX   r5   rY   r=   r4   r   concatenater	   rZ   r   rF   r[   r\   r   r   r   	_mix_pair   s$    88
 "zMixup._mix_pairc           	      C   s   |   \}}|dkrdS |rt|j|| j| jd\\}}}}}|dd d d d ||||f |d d d d ||||f< n$|dd| }||| |S )Nr   rU   r   )rT   r5   rY   r=   r4   r   Zmul_Zadd_)	rC   r   r   rQ   r)   r*   r+   r,   Z	x_flippedr   r   r   
_mix_batch   s    DzMixup._mix_batchc                 C   sh   t |d dksJ d| jdkr.| |}n | jdkrD| |}n
| |}t|| j|| j}||fS )Nr   r   )Batch size should be even when using thiselempair)r.   rA   r`   rb   rc   r   r   r@   )rC   r   r   r   r   r   r   __call__   s    


zMixup.__call__)	r   r   Nr   r7   r8   Tr9   r:   )__name__
__module____qualname____doc__rD   rS   rT   r`   rb   rc   rg   r   r   r   r   r6   Z   s     
r6   c                   @   s4   e Zd ZdZdddZdd Zdd Zdd
dZd	S )FastCollateMixupz Fast Collate w/ Mixup/Cutmix that applies different params to each element or whole batch

    A Mixup impl that's performed while collating the batches.
    Fc              	   C   sh  t |}|r|d n|}t ||ks(J | |\}}t|D ]}|| d }	|| }
|| d }|
dkr|| r|s| }t|j|
| j| jd\\}}}}}
||	 d d d ||||f |d d ||||f< |
||< n:|t	j
|
 ||	 d t	j
d|
   }t	j||d ||  t|t	j7  < q>|rXt	|t	|f}t|dS )Nr   r   r   r   rU   out)r.   rS   rX   copyr5   rY   r=   r4   rO   r   rI   rintr	   
from_numpyuint8ra   rH   rZ   r[   )rC   outputr8   halfrP   Znum_elemr]   rQ   r^   r_   r   mixedr)   r*   r+   r,   r   r   r   _mix_elem_collate   s.    
8
, z"FastCollateMixup._mix_elem_collatec              	   C   s  t |}| |d \}}t|d D ]}|| d }|| }|| d }	|| d }
d|  krldksrn J |dk r|| rt|j|| j| jd\\}}}}}|	d d ||||f  }|
d d ||||f |	d d ||||f< ||
d d ||||f< |||< nh|	t	j
| |
t	j
d|   }|
t	j
| |	t	j
d|   }
|}	t	j|
|
d t	j|	|	d ||  t|	t	j7  < ||  t|
t	j7  < q&t	||d d d f}t|dS )Nr   r   r   r   rU   rm   r   )r.   rS   rX   r5   rY   r=   r4   ro   rO   r   rI   rp   r	   rq   rr   ra   rZ   r[   )rC   rs   r8   rP   r]   rQ   r^   r_   r   Zmixed_iZmixed_jr)   r*   r+   r,   Zpatch_iZ
mixed_tempr   r   r   _mix_pair_collate   s4    

0
$$ z"FastCollateMixup._mix_pair_collatec              	   C   s
  t |}|  \}}|r:t|j|| j| jd\\}}}}	}t|D ]}
||
 d }||
 d }|dkr|r| }|| d d d ||||	f |d d ||||	f< n:|t	j
| || d t	j
d|   }t	j||d ||
  t|t	j7  < qB|S )NrU   r   r   r   rm   )r.   rT   r5   rY   r=   r4   rX   ro   rO   r   rI   rp   r	   rq   rr   )rC   rs   r8   rP   r   rQ   r)   r*   r+   r,   r^   r_   ru   r   r   r   _mix_batch_collate  s"    :, z#FastCollateMixup._mix_batch_collateNc                 C   s   t |}|d dksJ dd| jv }|r2|d }tj|g|d d jR tjd}| jdksh| jdkrz| j|||d}n$| jdkr| ||}n| ||}tj	d	d
 |D tj
d}t|| j|| j}|d | }||fS )Nr   r   rd   rt   rE   re   )rt   rf   c                 S   s   g | ]}|d  qS )r   r   ).0br   r   r   
<listcomp>8      z-FastCollateMixup.__call__.<locals>.<listcomp>)r.   rA   r	   rJ   rY   rr   rv   rw   rx   rZ   int64r   r   r@   )rC   r8   _rP   rt   rs   r   r   r   r   r   rg   +  s    
"
zFastCollateMixup.__call__)F)N)rh   ri   rj   rk   rv   rw   rx   rg   r   r   r   r   rl      s
   
rl   )r   r   )r   r   )r   N)N)NTN)rk   numpyr   r	   r   r   r-   r1   r5   r6   rl   r   r   r   r   <module>   s   




 