Ë
    þÍ:jè  ã            
       ó„   — d dl mZ d dlZd dlmZ d dlmZ dedededefd	„Zdeded
ededef
d„Z	ddeded
ee   defd„Z
y)é    )ÚOptionalN)ÚTensor)Ú"_check_retrieval_functional_inputsÚtargetÚpredsÚdiscount_cumsumÚreturnc                 óÆ  — t        j                  | dd¬«      \  }}}t        j                  |t         j                  ¬«      }|j	                  d|| j                  |j                  ¬«      «       ||z  }|j                  d¬«      dz
  }t        j                  |t         j                  ¬«      }||d      |d<   ||   j                  «       |dd ||z  j                  «       S )aI  Translated version of sklearns `_tie_average_dcg` function.

    Args:
        target: ground truth about each document relevance.
        preds: estimated probabilities of each document to be relevant.
        discount_cumsum: cumulative sum of the discount.

    Returns:
        The cumulative gain of the tied elements.

    T)Úreturn_inverseÚreturn_counts)Údtyper   ©Údimé   N)
ÚtorchÚuniqueÚ
zeros_likeÚfloat32Úscatter_add_Útor   ÚcumsumÚdiffÚsum)	r   r   r   Ú_ÚinvÚcountsÚrankedÚgroupsÚdiscount_sumss	            ú{/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/retrieval/ndcg.pyÚ_tie_average_dcgr!      sÊ   € ô —\‘\ 5 &¸ÈTÔR�N€A€sˆFÜ×Ñ˜f¬E¯M©MÔ:€FØ
×Ñ˜˜3 §	¡	°·± 	Ó =Ô>Ø�f‰_€FØ�]‰]˜qˆ]Ó! AÑ%€FÜ×$Ñ$ V´5·=±=ÔA€MØ& v¨a¡yÑ1€M�!ÑØ'¨Ñ/×4Ñ4Ó6€M�!�"ÐØ�]Ñ"×'Ñ'Ó)Ð)ó    Útop_kÚignore_tiesc                 ó8  — dt        j                  t        j                  | j                  d   | j                  ¬«      dz   «      z  }d||d |r,|j                  d¬«      }| |   }||z  j                  «       }|S |j                  d¬	«      }t        | ||«      }|S )
ay  Translated version of sklearns `_dcg_sample_scores` function.

    Args:
        target: ground truth about each document relevance.
        preds: estimated probabilities of each document to be relevant.
        top_k: consider only the top k elements
        ignore_ties: If True, ties are ignored. If False, ties are averaged.

    Returns:
        The cumulative gain

    g      ð?éÿÿÿÿ)Údeviceg       @g        NT)Ú
descendingr   )	r   Úlog2ÚarangeÚshaper'   Úargsortr   r   r!   )	r   r   r#   r$   ÚdiscountÚrankingr   Úcumulative_gainr   s	            r    Ú_dcg_sample_scoresr0   -   sž   € ð ”e—j‘j¤§¡¨f¯l©l¸2Ñ.>ÀvÇ}Á}Ô!UÐX[Ñ![Ó\Ñ]€HØ€HˆUˆVÐáØ—-‘-¨4�-Ó0ˆØ˜‘ˆØ# fÑ,×1Ñ1Ó3ˆð Ðð #Ÿ/™/¨b˜/Ó1ˆÜ*¨6°5¸/ÓJˆØÐr"   c                 ó  — t        | |d¬«      \  } }|€| j                  d   n|}t        |t        «      r|dkD  st	        d«      ‚t        || |d¬«      }t        |||d¬«      }|dk(  }d||<   || xx   ||    z  cc<   |j                  «       S )a  Compute `Normalized Discounted Cumulative Gain`_ (for information retrieval).

    ``preds`` and ``target`` should be of the same shape and live on the same device.
    ``target`` must be either `bool` or `integers` and ``preds`` must be ``float``,
    otherwise an error is raised.

    Args:
        preds: estimated probabilities of each document to be relevant.
        target: ground truth about each document relevance.
        top_k: consider only the top k elements (default: ``None``, which considers them all)

    Return:
        A single-value tensor with the nDCG of the predictions ``preds`` w.r.t. the labels ``target``.

    Raises:
        ValueError:
            If ``top_k`` parameter is not `None` or an integer larger than 0

    Example:
        >>> from torchmetrics.functional.retrieval import retrieval_normalized_dcg
        >>> preds = torch.tensor([.1, .2, .3, 4, 70])
        >>> target = torch.tensor([10, 0, 0, 1, 5])
        >>> retrieval_normalized_dcg(preds, target)
        tensor(0.6957)

    T)Úallow_non_binary_targetr&   r   z,`top_k` has to be a positive integer or NoneF)r$   )r   r+   Ú
isinstanceÚintÚ
ValueErrorr0   Úmean)r   r   r#   ÚgainÚnormalized_gainÚall_irrelevants         r    Úretrieval_normalized_dcgr:   G   s¡   € ô6 7°u¸fÐ^bÔc�M€Eˆ6à$˜}ˆE�K‰K˜ŠO°%€Eä�uœcÔ" u¨q¢yÜÐGÓHÐHä˜f e¨UÀÔF€DÜ(¨°¸ÈDÔQ€Oð %¨Ñ)€NØ€DˆÑØˆ.ˆÓ˜_¨n¨_Ñ=Ñ=Óà�9‰9‹;Ðr"   )N)Útypingr   r   r   Útorchmetrics.utilities.checksr   r!   r4   Úboolr0   r:   © r"   r    ú<module>r?      s„   ðõ ã Ý å Lð*˜Vð *¨Fð *ÀVð *ÐPVó *ð.˜vð ¨fð ¸Sð Ètð ÐX^ó ñ4* Fð *°Fð *À8ÈCÁ=ð *Ð\bô *r"   