a
    þd$  ã                   @   s0  d dl mZmZ d dlmZ d dlmZ d dlmZ d dl	m
Z
 d dlmZ d dlmZmZ d dlmZ d d	lmZ d d
lmZ d dlmZ G dd„ deƒZG dd„ deƒZG dd„ deƒZG dd„ deƒZG dd„ deƒZG dd„ de
ƒZG dd„ deƒZG dd„ deƒZG dd„ deƒZ G dd„ deƒZ!d S )!é    )ÚAnyÚOptional)ÚRetrievalMAP)ÚRetrievalFallOut)ÚRetrievalHitRate)ÚRetrievalNormalizedDCG)ÚRetrievalPrecision)ÚRetrievalPrecisionRecallCurveÚRetrievalRecallAtFixedPrecision)ÚRetrievalRPrecision)ÚRetrievalRecall)ÚRetrievalMRR)Ú_deprecated_root_import_classc                       s8   e Zd ZdZdeee ee eddœ‡ fdd„Z‡  Z	S )Ú_RetrievalFallOuta`  Wrapper for deprecated import.

    >>> from torch import tensor
    >>> indexes = tensor([0, 0, 0, 1, 1, 1, 1])
    >>> preds = tensor([0.2, 0.3, 0.5, 0.1, 0.3, 0.5, 0.2])
    >>> target = tensor([False, False, True, False, True, False, True])
    >>> fo = _RetrievalFallOut(top_k=2)
    >>> fo(preds, target, indexes=indexes)
    tensor(0.5000)

    ÚposN©Úempty_target_actionÚignore_indexÚtop_kÚkwargsÚreturnc                    s(   t ddƒ tƒ jf |||dœ|¤Ž d S )Nr   Ú	retrieval©r   r   r   ©r   ÚsuperÚ__init__©Úselfr   r   r   r   ©Ú	__class__© úk/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torchmetrics/retrieval/_deprecated.pyr      s    
z_RetrievalFallOut.__init__)r   NN©
Ú__name__Ú
__module__Ú__qualname__Ú__doc__Ústrr   Úintr   r   Ú__classcell__r    r    r   r!   r      s      üúr   c                       s8   e Zd ZdZdeee ee eddœ‡ fdd„Z‡  Z	S )Ú_RetrievalHitRateab  Wrapper for deprecated import.

    >>> from torch import tensor
    >>> indexes = tensor([0, 0, 0, 1, 1, 1, 1])
    >>> preds = tensor([0.2, 0.3, 0.5, 0.1, 0.3, 0.5, 0.2])
    >>> target = tensor([True, False, False, False, True, False, True])
    >>> hr2 = _RetrievalHitRate(top_k=2)
    >>> hr2(preds, target, indexes=indexes)
    tensor(0.5000)

    ÚnegNr   c                    s(   t ddƒ tƒ jf |||dœ|¤Ž d S )Nr   r   r   r   r   r   r    r!   r   4   s    
z_RetrievalHitRate.__init__)r+   NNr"   r    r    r   r!   r*   '   s      üúr*   c                       s8   e Zd ZdZdeee ee eddœ‡ fdd„Z‡  Z	S )Ú_RetrievalMAPaY  Wrapper for deprecated import.

    >>> from torch import tensor
    >>> indexes = tensor([0, 0, 0, 1, 1, 1, 1])
    >>> preds = tensor([0.2, 0.3, 0.5, 0.1, 0.3, 0.5, 0.2])
    >>> target = tensor([False, False, True, False, True, False, True])
    >>> rmap = _RetrievalMAP()
    >>> rmap(preds, target, indexes=indexes)
    tensor(0.7917)

    r+   Nr   c                    s(   t ddƒ tƒ jf |||dœ|¤Ž d S )Nr   r   r   r   r   r   r    r!   r   L   s    
z_RetrievalMAP.__init__)r+   NNr"   r    r    r   r!   r,   ?   s      üúr,   c                       s8   e Zd ZdZdeee ee eddœ‡ fdd„Z‡  Z	S )Ú_RetrievalRecalla_  Wrapper for deprecated import.

    >>> from torch import tensor
    >>> indexes = tensor([0, 0, 0, 1, 1, 1, 1])
    >>> preds = tensor([0.2, 0.3, 0.5, 0.1, 0.3, 0.5, 0.2])
    >>> target = tensor([False, False, True, False, True, False, True])
    >>> r2 = _RetrievalRecall(top_k=2)
    >>> r2(preds, target, indexes=indexes)
    tensor(0.7500)

    r+   Nr   c                    s(   t ddƒ tƒ jf |||dœ|¤Ž d S )Nr   r   r   r   r   r   r    r!   r   d   s    
z_RetrievalRecall.__init__)r+   NNr"   r    r    r   r!   r-   W   s      üúr-   c                       s2   e Zd ZdZdeee eddœ‡ fdd„Z‡  Z	S )Ú_RetrievalRPrecisiona\  Wrapper for deprecated import.

    >>> from torch import tensor
    >>> indexes = tensor([0, 0, 0, 1, 1, 1, 1])
    >>> preds = tensor([0.2, 0.3, 0.5, 0.1, 0.3, 0.5, 0.2])
    >>> target = tensor([False, False, True, False, True, False, True])
    >>> p2 = _RetrievalRPrecision()
    >>> p2(preds, target, indexes=indexes)
    tensor(0.7500)

    r+   N©r   r   r   r   c                    s&   t ddƒ tƒ jf ||dœ|¤Ž d S )Nr   r   ©r   r   r   ©r   r   r   r   r   r    r!   r   |   s    
z_RetrievalRPrecision.__init__)r+   Nr"   r    r    r   r!   r.   o   s     ýûr.   c                       s8   e Zd ZdZdeee ee eddœ‡ fdd„Z‡  Z	S )Ú_RetrievalNormalizedDCGac  Wrapper for deprecated import.

    >>> from torch import tensor
    >>> indexes = tensor([0, 0, 0, 1, 1, 1, 1])
    >>> preds = tensor([0.2, 0.3, 0.5, 0.1, 0.3, 0.5, 0.2])
    >>> target = tensor([False, False, True, False, True, False, True])
    >>> ndcg = _RetrievalNormalizedDCG()
    >>> ndcg(preds, target, indexes=indexes)
    tensor(0.8467)

    r+   Nr   c                    s(   t ddƒ tƒ jf |||dœ|¤Ž d S )Nr   r   r   r   r   r   r    r!   r   “   s    
z _RetrievalNormalizedDCG.__init__)r+   NNr"   r    r    r   r!   r2   †   s      üúr2   c                       s:   e Zd ZdZdeee ee eeddœ‡ fdd„Z	‡  Z
S )	Ú_RetrievalPrecisionab  Wrapper for deprecated import.

    >>> from torch import tensor
    >>> indexes = tensor([0, 0, 0, 1, 1, 1, 1])
    >>> preds = tensor([0.2, 0.3, 0.5, 0.1, 0.3, 0.5, 0.2])
    >>> target = tensor([False, False, True, False, True, False, True])
    >>> p2 = _RetrievalPrecision(top_k=2)
    >>> p2(preds, target, indexes=indexes)
    tensor(0.5000)

    r+   NF)r   r   r   Ú
adaptive_kr   r   c                    s*   t ddƒ tƒ jf ||||dœ|¤Ž d S )NÚ r   )r   r   r   r4   r   )r   r   r   r   r4   r   r   r    r!   r   «   s    
üûz_RetrievalPrecision.__init__)r+   NNF)r#   r$   r%   r&   r'   r   r(   Úboolr   r   r)   r    r    r   r!   r3   ž   s       ûùr3   c                       s:   e Zd ZdZdee eeee eddœ‡ fdd„Z	‡  Z
S )	Ú_RetrievalPrecisionRecallCurvea  Wrapper for deprecated import.

    >>> from torch import tensor
    >>> indexes = tensor([0, 0, 0, 0, 1, 1, 1])
    >>> preds = tensor([0.4, 0.01, 0.5, 0.6, 0.2, 0.3, 0.5])
    >>> target = tensor([True, False, False, True, True, False, True])
    >>> r = _RetrievalPrecisionRecallCurve(max_k=4)
    >>> precisions, recalls, top_k = r(preds, target, indexes=indexes)
    >>> precisions
    tensor([1.0000, 0.5000, 0.6667, 0.5000])
    >>> recalls
    tensor([0.5000, 0.5000, 1.0000, 1.0000])
    >>> top_k
    tensor([1, 2, 3, 4])

    NFr+   )Úmax_kr4   r   r   r   r   c                    s*   t ddƒ tƒ jf ||||dœ|¤Ž d S )Nr5   r   )r8   r4   r   r   r   )r   r8   r4   r   r   r   r   r    r!   r   Ï   s    
üûz'_RetrievalPrecisionRecallCurve.__init__)NFr+   N)r#   r$   r%   r&   r   r(   r6   r'   r   r   r)   r    r    r   r!   r7   ½   s       ûùr7   c                	       s<   e Zd ZdZd	eee eeee e	ddœ‡ fdd„Z
‡  ZS )
Ú _RetrievalRecallAtFixedPrecisiona„  Wrapper for deprecated import.

    >>> from torch import tensor
    >>> indexes = tensor([0, 0, 0, 0, 1, 1, 1])
    >>> preds = tensor([0.4, 0.01, 0.5, 0.6, 0.2, 0.3, 0.5])
    >>> target = tensor([True, False, False, True, True, False, True])
    >>> r = _RetrievalRecallAtFixedPrecision(min_precision=0.8)
    >>> r(preds, target, indexes=indexes)
    (tensor(0.5000), tensor(1))

    ç        NFr+   )Úmin_precisionr8   r4   r   r   r   r   c                    s,   t ddƒ tƒ jf |||||dœ|¤Ž d S )Nr
   r   )r;   r8   r4   r   r   r   )r   r;   r8   r4   r   r   r   r   r    r!   r   î   s    	
ûúz)_RetrievalRecallAtFixedPrecision.__init__)r:   NFr+   N)r#   r$   r%   r&   Úfloatr   r(   r6   r'   r   r   r)   r    r    r   r!   r9   á   s        úør9   c                       s2   e Zd ZdZdeee eddœ‡ fdd„Z‡  Z	S )Ú_RetrievalMRRaW  Wrapper for deprecated import.

    >>> from torch import tensor
    >>> indexes = tensor([0, 0, 0, 1, 1, 1, 1])
    >>> preds = tensor([0.2, 0.3, 0.5, 0.1, 0.3, 0.5, 0.2])
    >>> target = tensor([False, False, True, False, True, False, True])
    >>> mrr = _RetrievalMRR()
    >>> mrr(preds, target, indexes=indexes)
    tensor(0.7500)

    r+   Nr/   c                    s&   t ddƒ tƒ jf ||dœ|¤Ž d S )Nr5   r   r0   r   r1   r   r    r!   r     s    
z_RetrievalMRR.__init__)r+   Nr"   r    r    r   r!   r=     s     ýûr=   N)"Útypingr   r   Z(torchmetrics.retrieval.average_precisionr   Ztorchmetrics.retrieval.fall_outr   Ztorchmetrics.retrieval.hit_rater   Ztorchmetrics.retrieval.ndcgr   Z torchmetrics.retrieval.precisionr   Z-torchmetrics.retrieval.precision_recall_curver	   r
   Z"torchmetrics.retrieval.r_precisionr   Ztorchmetrics.retrieval.recallr   Z&torchmetrics.retrieval.reciprocal_rankr   Ztorchmetrics.utilities.printsr   r   r*   r,   r-   r.   r2   r3   r7   r9   r=   r    r    r    r!   Ú<module>   s(   $!