Ë
    l^(h 9  ã                  ó®   — d dl mZ d dlZd dlZd dlZd dlmZ d dlZd dl	m
Z
mZ d dlmZ d dlmZ erd dlmZ  ej"                  e«      Z G d„ d	e«      Zy)
é    )ÚannotationsN)ÚTYPE_CHECKING)Úaverage_precision_scoreÚ
ndcg_score)Útqdm)ÚSentenceEvaluator)ÚCrossEncoderc                  ó|   ‡ — e Zd ZdZ	 	 	 	 	 	 	 d	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dˆ fd„Z	 d	 	 	 	 	 	 	 	 	 d	d„Zd„ Zd„ Zˆ xZS )
ÚCrossEncoderRerankingEvaluatora  
    This class evaluates a CrossEncoder model for the task of re-ranking.

    Given a query and a list of documents, it computes the score [query, doc_i] for all possible
    documents and sorts them in decreasing order. Then, MRR@10, NDCG@10 and MAP are computed to measure the quality of the ranking.

    The evaluator expects a list of samples. Each sample is a dictionary with the mandatory "query" and "positive" keys,
    and either a "negative" or a "documents" key. The "query" is the search query, the "positive" is a list of relevant
    documents, and the "negative" is a list of irrelevant documents. Alternatively, the "documents" key can be used to
    provide a list of all documents, including the positive ones. In this case, the evaluator will assume that the list
    is already ranked by similarity, with the most similar documents first, and will report both the reranking performance
    as well as the performance before reranking. This can be useful to measure the improvement of the reranking on
    top of a first-stage retrieval (e.g. a SentenceTransformer model).

    Note that the maximum score is 1.0 by default, because all positive documents are included in the ranking. This
    can be toggled off by using samples with ``documents`` instead of ``negative``, i.e. ranked lists of all documents
    including the positive ones, together with ``always_rerank_positives=False``. ``always_rerank_positives=False`` only
    works when using ``documents`` instead of ``negative``.

    Args:
        samples (list): A list of dictionaries, where each dictionary represents a sample and has the following keys:
            - 'query' (mandatory): The search query.
            - 'positive' (mandatory): A list of positive (relevant) documents.
            - 'negative' (optional): A list of negative (irrelevant) documents. Mutually exclusive with 'documents'.
            - 'documents' (optional): A list of all documents, including the positive ones. This list is assumed to be
                ranked by similarity, with the most similar documents first. Mutually exclusive with 'negative'.
        at_k (int, optional): Only consider the top k most similar documents to each query for the evaluation. Defaults to 10.
        always_rerank_positives (bool): If True, always evaluate with all positives included. If False, only include
            the positives that are already in the documents list. Always set to True if your ``samples`` contain ``negative``
            instead of ``documents``. When using ``documents``, setting this to True will result in a more useful evaluation
            signal, but setting it to False will result in a more realistic evaluation. Defaults to True.
        name (str, optional): Name of the evaluator, used for logging, saving in a CSV, and the model card. Defaults to "".
        batch_size (int): Batch size to compute sentence embeddings. Defaults to 64.
        show_progress_bar (bool): Show progress bar when computing embeddings. Defaults to False.
        write_csv (bool): Write results to CSV file. Defaults to True.
        mrr_at_k (Optional[int], optional): Deprecated parameter. Please use `at_k` instead. Defaults to None.

    Example:
        ::

            from sentence_transformers import CrossEncoder
            from sentence_transformers.cross_encoder.evaluation import CrossEncoderRerankingEvaluator
            from datasets import load_dataset

            # Load a model
            model = CrossEncoder("cross-encoder/ms-marco-MiniLM-L6-v2")

            # Load a dataset with queries, positives, and negatives
            eval_dataset = load_dataset("microsoft/ms_marco", "v1.1", split="validation")

            samples = [
                {
                    "query": sample["query"],
                    "positive": [text for is_selected, text in zip(sample["passages"]["is_selected"], sample["passages"]["passage_text"]) if is_selected],
                    "documents": sample["passages"]["passage_text"],
                    # or
                    # "negative": [text for is_selected, text in zip(sample["passages"]["is_selected"], sample["passages"]["passage_text"]) if not is_selected],
                }
                for sample in eval_dataset
            ]

            # Initialize the evaluator
            reranking_evaluator = CrossEncoderRerankingEvaluator(
                samples=samples,
                name="ms-marco-dev",
                show_progress_bar=True,
            )
            results = reranking_evaluator(model)
            '''
            CrossEncoderRerankingEvaluator: Evaluating the model on the ms-marco-dev dataset:
            Queries: 10047    Positives: Min 0.0, Mean 1.1, Max 5.0   Negatives: Min 1.0, Mean 7.1, Max 10.0
                     Base  -> Reranked
            MAP:     34.03 -> 62.36
            MRR@10:  34.67 -> 62.96
            NDCG@10: 49.05 -> 71.05
            '''
            print(reranking_evaluator.primary_metric)
            # => ms-marco-dev_ndcg@10
            print(results[reranking_evaluator.primary_metric])
            # => 0.7104656857184184
    c	                ó  •— t         ‰	| �  «        || _        |�!t        j	                  d|› d�«       || _        n|| _        || _        || _        || _        || _	        t        | j                  t        «      r(t        | j                  j                  «       «      | _        d|rd|z   ndz   d| j
                  › d�z   | _        dd	d
d| j
                  › �d| j
                  › �g| _        || _        d| j
                  › �| _        y )Nz?The `mrr_at_k` parameter has been deprecated; please use `at_k=z
` instead.r   Ú_Ú z
_results_@z.csvÚepochÚstepsÚMAPúMRR@úNDCG@úndcg@)ÚsuperÚ__init__ÚsamplesÚloggerÚwarningÚat_kÚalways_rerank_positivesÚnameÚ
batch_sizeÚshow_progress_barÚ
isinstanceÚdictÚlistÚvaluesÚcsv_fileÚcsv_headersÚ	write_csvÚprimary_metric)
Úselfr   r   r   r   r   r   r%   Úmrr_at_kÚ	__class__s
            €úv/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/sentence_transformers/cross_encoder/evaluation/reranking.pyr   z'CrossEncoderRerankingEvaluator.__init__g   sý   ø€ ô 	‰ÑÔØˆŒØÐÜ�N‰NÐ\Ð]eÐ\fÐfpÐqÔrØ ˆD�IàˆDŒIØ'>ˆÔ$àˆŒ	Ø$ˆŒØ!2ˆÔä�d—l‘l¤DÔ)Ü §¡× 3Ñ 3Ó 5Ó6ˆDŒLà8É$¸CÀ$ºJÐTVÑWÐ\fÐgk×gpÑgpÐfqÐquÐZvÑvˆŒØ# W¨e°t¸D¿I¹I¸;Ð5GÈ5ÐQU×QZÑQZÐP[ÐI\Ð]ˆÔØ"ˆŒØ % d§i¡i [Ð1ˆÕó    c                ó  — |dk7  r|dk(  rd|› �}nd|› d|› d�}nd}t         j                  d| j                  › d|› d	�«       g }g }g }g }	g }
g }d
}g }g }t        | j                  d| j
                   d¬«      D �]	  }d|vrt        d«      ‚d|vrt        d«      ‚d|v rd|v sd|vrd|vrt        d«      ‚|d   }|d   }t        |t        «      r|g}|j                  dd «      }|j                  dd «      }|�r,|D �cg c]  }t        ||v «      ‘Œ }}t        |«      d
k(  rd\  }}}n]|dgt        |«      t        |«      z
  z  z  }t        j                  t        t        |«      d
d«      «      }| j!                  ||«      \  }}}|j#                  |«       |j#                  |«       |j#                  |«       | j$                  rD||D �cg c]	  }||vsŒ|‘Œ c}z   }dgt        |«      z  d
gt        |«      t        |«      z
  z  z   }nA|}|D �cg c]  }t        ||v «      ‘Œ }}n$||z   }dgt        |«      z  d
gt        |«      z  z   }|dz  }|j#                  t        |«      «       |j#                  t        |«      t        |«      z
  «       t        |«      d
k(  r5|	j#                  d
«       |
j#                  d
«       |j#                  d
«       �ŒY|D �cg c]  }||g‘Œ }}|j'                  |dd¬«      }t        |«      t        |«      z
  x}r*t        j(                  |t        j*                  |«      g«      }| j!                  ||«      \  } }!}"|	j#                  | «       |
j#                  |!«       |j#                  |"«       �Œ t        j,                  |	«      }#t        j,                  |
«      }$t        j,                  |«      }%d|%d| j.                  › �|#d| j.                  › �|$i}&t         j                  d|› dt        j0                  |«      d›dt        j,                  |«      d›d t        j2                  |«      d›d!t        j0                  |«      d›dt        j,                  |«      d›d t        j2                  |«      d›�«       �rÞt        j,                  |«      }'t        j,                  |«      }(t        j,                  |«      })d"|)d#| j.                  › �|'d$| j.                  › �|(i}*t         j                  d%t        t        | j.                  «      «      z  › d&�«       t         j                  d'd%t        t        | j.                  «      «      z  › d(|)d)z  d*›d+|%d)z  d*›�«       t         j                  d,| j.                  › d-|'d)z  d*›d+|#d)z  d*›�«       t         j                  d.| j.                  › d/|(d)z  d*›d+|$d)z  d*›�«       d|%d0›d1|%|)z
  d2›d3�d| j.                  › �|#d0›d1|#|'z
  d2›d3�d| j.                  › �|$d0›d1|$|(z
  d2›d3�i}+| j5                  |+| j                  «      }+| j7                  ||+||«       |&j9                  |*«       | j5                  |&| j                  «      }&nÀt         j                  d'd%t        t        | j.                  «      «      z  › d(|%d)z  d*›�«       t         j                  d,| j.                  › d-|#d)z  d*›�«       t         j                  d.| j.                  › d/|$d)z  d*›�«       | j5                  |&| j                  «      }&| j7                  ||&||«       |�º| j:                  r®t<        j>                  jA                  || jB                  «      },t<        j>                  jE                  |,«      }-tG        |,|-rd4nd5d6¬7«      5 }.tI        jJ                  |.«      }/|-s|/jM                  | jN                  «       |/jM                  |||%|#|$g«       d d d «       |&S |&S c c}w c c}w c c}w c c}w # 1 sw Y   |&S xY w)8Néÿÿÿÿz after epoch z
 in epoch z after z stepsr   z<CrossEncoderRerankingEvaluator: Evaluating the model on the z datasetú:r   zEvaluating samplesF)ÚdescÚdisableÚleaveÚqueryzECrossEncoderRerankingEvaluator requires a 'query' key in each sample.ÚpositivezHCrossEncoderRerankingEvaluator requires a 'positive' key in each sample.ÚnegativeÚ	documentszaCrossEncoderRerankingEvaluator requires exactly one of 'negative' and 'documents' in each sample.)r   r   r   é   T)Úconvert_to_numpyr   Úmapzmrr@r   z	Queries: z	Positives: Min z.1fz, Mean z, Max z	Negatives: Min Úbase_mapz	base_mrr@z
base_ndcg@ú z       Base  -> RerankedzMAP:z   éd   z.2fz -> r   z:  r   z: z.4fz (z+.4fú)ÚaÚwzutf-8)ÚmodeÚencoding)(r   Úinfor   r   r   r   Ú
ValueErrorr   ÚstrÚgetÚintÚsumÚlenÚnpÚarrayÚrangeÚcompute_metricsÚappendr   ÚpredictÚconcatenateÚzerosÚmeanr   ÚminÚmaxÚprefix_name_to_metricsÚ store_metrics_in_model_card_dataÚupdater%   ÚosÚpathÚjoinr#   ÚisfileÚopenÚcsvÚwriterÚwriterowr$   )0r'   ÚmodelÚoutput_pathr   r   Úout_txtÚbase_mrr_scoresÚbase_ndcg_scoresÚbase_ap_scoresÚall_mrr_scoresÚall_ndcg_scoresÚall_ap_scoresÚnum_queriesÚnum_positivesÚnum_negativesÚinstancer2   r3   r4   r5   ÚsampleÚbase_is_relevantÚbase_mrrÚ	base_ndcgÚbase_apÚbase_pred_scoresÚdocÚdocsÚis_relevantÚmodel_inputÚpred_scoresÚnum_ignored_positivesÚmrrÚndcgÚapÚmean_mrrÚ	mean_ndcgÚmean_apÚmetricsÚmean_base_mrrÚmean_base_ndcgÚmean_base_apÚbase_metricsÚmodel_card_metricsÚcsv_pathÚoutput_file_existsÚfr\   s0                                                   r*   Ú__call__z'CrossEncoderRerankingEvaluator.__call__‡   s.  € ð �BŠ;Ø˜Š{Ø)¨%¨Ð1‘à& u g¨W°U°G¸6ÐB‘àˆGä�‰ÐRÐSW×S\ÑS\ÐR]Ð]eÐfmÐenÐnoÐpÔqàˆØÐØˆØˆØˆØˆØˆØˆØˆÜ˜TŸ\™\Ð0DÐRV×RhÑRhÐNhÐpuÔvó A	%ˆHØ˜hÑ&Ü Ð!hÓiÐiØ Ñ)Ü Ð!kÓlÐlØ˜hÑ&¨;¸(Ñ+BØ (Ñ*¨{À(Ñ/Jä Øwóð ð ˜WÑ%ˆEØ 
Ñ+ˆHÜ˜(¤CÔ(Ø$˜:�à—|‘| J°Ó5ˆHØ Ÿ™ [°$Ó7ˆIâØJSÖ#TÀ¤C¨°(Ð(:Õ$;Ð#TÐ Ð#TÜÐ'Ó(¨AÒ-Ø3:Ñ0�H˜i©ð %¨¨¬s°8«}¼sÐCSÓ?TÑ/TÑ(UÑUÐ$Ü')§x¡x´´cÐ:JÓ6KÈQÐPRÓ0SÓ'TÐ$Ø37×3GÑ3GÐHXÐZjÓ3kÑ0�H˜i¨Ø×&Ñ& xÔ0Ø ×'Ñ'¨	Ô2Ø×%Ñ% gÔ.à×/Ò/Ø#°iÖ&W¨sÀ3ÈhÒCV¢sÒ&WÑW�DØ#$ #¬¨H«Ñ"5¸¸¼sÀ4»yÌ3ÈxË=Ñ?XÑ8YÑ"Y‘Kà$�DØIRÖ"S¸v¤3 v°Ð'9Õ#:Ð"S�KÑ"Sà (Ñ*�Ø ˜c¤C¨£MÑ1°Q°C¼#¸h»-Ñ4GÑG�à˜1ÑˆKà× Ñ ¤ X£Ô/Ø× Ñ ¤ [Ó!1´C¸Ó4DÑ!DÔEä�;Ó 1Ò$Ø×%Ñ% aÔ(Ø×&Ñ& qÔ)Ø×$Ñ$ QÔ'Ùà37Ö8¨C˜E 3š<Ð8ˆKÐ8ØŸ-™-¨ÀdÐ^c˜-ÓdˆKô ),¨KÓ(8¼3¸{Ó;KÑ(KÐKÐ$ÐKÜ Ÿn™n¨k¼2¿8¹8ÐDYÓ;ZÐ-[Ó\�à ×0Ñ0°¸kÓJ‰MˆC��rà×!Ñ! #Ô&Ø×"Ñ" 4Ô(Ø× Ñ  Ö$ðCA	%ôF —7‘7˜>Ó*ˆÜ—G‘G˜OÓ,ˆ	Ü—'‘'˜-Ó(ˆà�7Ø�4—9‘9�+Ð Ø�D—I‘I�;Ð ð
ˆô 	�‰Ø˜�}ð %Ü Ÿf™f ]Ó3°CÐ8¸ÄÇÁÈÓ@VÐWZÐ?[Ð[aÔbd×bhÑbhÐivÓbwÐx{Ða|ð }Ü Ÿf™f ]Ó3°CÐ8¸ÄÇÁÈÓ@VÐWZÐ?[Ð[aÔbd×bhÑbhÐivÓbwÐx{Ða|ð~ô	
ò
 ÜŸG™G OÓ4ˆMÜŸW™WÐ%5Ó6ˆNÜŸ7™7 >Ó2ˆLà˜LØ˜DŸI™I˜;Ð'¨Ø˜TŸY™Y˜KÐ(¨.ðˆLô
 �K‰K˜3¤¤S¨¯©£^Ó!4Ñ4Ð5Ð5MÐNÔOÜ�K‰K˜$˜s¤S¬¨T¯Y©Y«Ó%8Ñ8Ð9¸¸\ÈCÑ=OÐPSÐ<TÐTXÐY`ÐcfÑYfÐgjÐXkÐlÔmÜ�K‰K˜$˜tŸy™y˜k¨¨]¸SÑ-@ÀÐ,EÀTÈ(ÐUXÉ.ÐY\ÐI]Ð^Ô_Ü�K‰K˜% §	¡	˜{¨"¨^¸cÑ-AÀ#Ð,FÀdÈ9ÐWZÉ?Ð[^ÐJ_Ð`Ôað ˜' #˜ b¨°<Ñ)?ÀÐ(EÀQÐGØ�t—y‘y�kÐ" x° n°B°xÀ-Ñ7OÐPTÐ6UÐUVÐ$WØ˜Ÿ	™	�{Ð#¨	°# °b¸À^Ñ9SÐTXÐ8YÐYZÐ%[ð"Ðð
 "&×!<Ñ!<Ð=OÐQU×QZÑQZÓ![ÐØ×1Ñ1°%Ð9KÈUÐTYÔZà�N‰N˜<Ô(Ø×1Ñ1°'¸4¿9¹9ÓE‰Gä�K‰K˜$˜s¤S¬¨T¯Y©Y«Ó%8Ñ8Ð9¸¸WÀs¹]È3Ð<OÐPÔQÜ�K‰K˜$˜tŸy™y˜k¨¨X¸©^¸CÐ,@ÐAÔBÜ�K‰K˜% §	¡	˜{¨"¨Y¸©_¸SÐ,AÐBÔCà×1Ñ1°'¸4¿9¹9ÓEˆGØ×1Ñ1°%¸À%ÈÔOàÐ" t§~¢~Ü—w‘w—|‘| K°·±Ó?ˆHÜ!#§¡§¡°Ó!9ÐÜ�hÑ,>¡SÀCÐRYÔZð NÐ^_ÜŸ™ A›�Ù)Ø—O‘O D×$4Ñ$4Ô5à—‘ ¨¨w¸À)Ð LÔM÷Nð ˆˆwˆùòM $Uùò 'Xùò #Tùò  9÷|Nð ˆús+   Ã>_(Ç
	_-Ç_-È_2Ë_7ÞA	_<ß<`c                óð   — t        j                  |«      d d d…   }d}t        |d| j                   «      D ]  \  }}||   sŒd|dz   z  } n t	        |g|g| j                  ¬«      }t        ||«      }|||fS )Nr-   r   r6   )Úk)rH   ÚargsortÚ	enumerater   r   r   )	r'   Úy_trueÚy_predÚrankingrw   ÚrankÚindexrx   ry   s	            r*   rK   z.CrossEncoderRerankingEvaluator.compute_metrics  sˆ   € Ü—*‘*˜VÓ$¡T r TÑ*ˆàˆÜ$ W¨Q°·±Ð%;Ó<ò 	‰KˆD�%Ø�e‹}Ø˜4 !™8‘n�Ùð	ô
 ˜6˜( V H°·	±	Ô:ˆÜ$ V¨VÓ4ˆØ�D˜"ˆ}Ðr+   c                óz   — d| j                   i}| j                  r d| j                  d   v r| j                  |d<   |S )Nr   r5   r   r   )r   r   r   )r'   Úconfig_dicts     r*   Úget_config_dictz.CrossEncoderRerankingEvaluator.get_config_dict'  sA   € à�D—I‘Ið
ˆð �<Š<˜K¨4¯<©<¸©?Ñ:Ø59×5QÑ5QˆKÐ1Ñ2ØÐr+   )é
   Tr   é@   FTN)r   z list[dict[str, str | list[str]]]r   rE   r   Úboolr   rC   r   rE   r   r•   r%   r•   r(   z
int | None)Nr-   r-   )
r^   r	   r_   rC   r   rE   r   rE   Úreturnzdict[str, float])	Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r†   rK   r’   Ú__classcell__)r)   s   @r*   r   r      sº   ø„ ñPðj Ø(,ØØØ"'ØØ#ð2à1ð2ð ð2ð "&ð	2ð
 ð2ð ð2ð  ð2ð ð2ð õ2ðB []ðQØ!ðQØ03ðQØCFðQØTWðQà	óQòför+   r   )Ú
__future__r   r[   ÚloggingrV   Útypingr   ÚnumpyrH   Úsklearn.metricsr   r   r   Ú2sentence_transformers.evaluation.SentenceEvaluatorr   Ú0sentence_transformers.cross_encoder.CrossEncoderr	   Ú	getLoggerr—   r   r   © r+   r*   ú<module>r¥      sG   ðÝ "ã 
Û Û 	Ý  ã ß ?Ý å PáÝMà	ˆ×	Ñ	˜8Ó	$€ôYÐ%6õ Yr+   