Ë
    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mZm	Z	 d dl
Zd dl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)Únullcontext)ÚTYPE_CHECKINGÚCallable)Úaverage_precision_scoreÚ
ndcg_score)ÚSentenceEvaluator)Úcos_sim)ÚSentenceTransformerc            	      ó”   ‡ — e Zd ZdZdddedddddf		 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dˆ fd„Z	 d	 	 	 	 	 	 	 	 	 dd	„Zd
„ Zd„ Zd„ Z	d„ Z
ˆ xZS )ÚRerankingEvaluatoraC  
    This class evaluates a SentenceTransformer 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 is compute to measure the quality of the ranking.

    Args:
        samples (list): A list of dictionaries, where each dictionary represents a sample and has the following keys:
            - 'query': The search query.
            - 'positive': A list of positive (relevant) documents.
            - 'negative': A list of negative (irrelevant) documents.
        at_k (int, optional): Only consider the top k most similar documents to each query for the evaluation. Defaults to 10.
        name (str, optional): Name of the evaluator. Defaults to "".
        write_csv (bool, optional): Write results to CSV file. Defaults to True.
        similarity_fct (Callable[[torch.Tensor, torch.Tensor], torch.Tensor], optional): Similarity function between sentence embeddings. By default, cosine similarity. Defaults to cos_sim.
        batch_size (int, optional): Batch size to compute sentence embeddings. Defaults to 64.
        show_progress_bar (bool, optional): Show progress bar when computing embeddings. Defaults to False.
        use_batched_encoding (bool, optional): Whether or not to encode queries and documents in batches for greater speed, or 1-by-1 to save memory. Defaults to True.
        truncate_dim (Optional[int], optional): The dimension to truncate sentence embeddings to. `None` uses the model's current truncation dimension. Defaults to None.
        mrr_at_k (Optional[int], optional): Deprecated parameter. Please use `at_k` instead. Defaults to None.

    Example:
        ::

            from sentence_transformers import SentenceTransformer
            from sentence_transformers.evaluation import RerankingEvaluator
            from datasets import load_dataset

            # Load a model
            model = SentenceTransformer("all-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],
                    "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 = RerankingEvaluator(
                samples=samples,
                name="ms-marco-dev",
            )
            results = reranking_evaluator(model)
            '''
            RerankingEvaluator: Evaluating the model on the ms-marco-dev dataset:
            Queries: 9706      Positives: Min 1.0, Mean 1.1, Max 5.0   Negatives: Min 1.0, Mean 7.1, Max 9.0
            MAP: 56.07
            MRR@10: 56.70
            NDCG@10: 67.08
            '''
            print(reranking_evaluator.primary_metric)
            # => ms-marco-dev_ndcg@10
            print(results[reranking_evaluator.primary_metric])
            # => 0.6708042171399308
    é
   Ú Té@   FNc                ó²  •— t         ‰| �  «        || _        || _        |
�!t        j                  d|
› d�«       |
| _        n|| _        || _        || _        || _	        || _
        |	| _        t        | j                  t        «      r(t        | j                  j                  «       «      | _        | j                  D �cg c](  }t!        |d   «      dkD  sŒt!        |d   «      dkD  sŒ'|‘Œ* c}| _        d|rd|z   ndz   d	| j                  › d
�z   | _        dddd| j                  › �d| j                  › �g| _        || _        d| j                  › �| _        y c c}w )Nz?The `mrr_at_k` parameter has been deprecated; please use `at_k=z
` instead.Úpositiver   Únegativer   Ú_r   z
_results_@z.csvÚepochÚstepsÚMAPúMRR@úNDCG@úndcg@)ÚsuperÚ__init__ÚsamplesÚnameÚloggerÚwarningÚat_kÚsimilarity_fctÚ
batch_sizeÚshow_progress_barÚuse_batched_encodingÚtruncate_dimÚ
isinstanceÚdictÚlistÚvaluesÚlenÚcsv_fileÚcsv_headersÚ	write_csvÚprimary_metric)Úselfr   r!   r   r.   r"   r#   r$   r%   r&   Úmrr_at_kÚsampleÚ	__class__s               €úq/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/sentence_transformers/evaluation/RerankingEvaluator.pyr   zRerankingEvaluator.__init__V   sZ  ø€ ô 	‰ÑÔØˆŒØˆŒ	àÐÜ�N‰NÐ\Ð]eÐ\fÐfpÐqÔrØ ˆD�IàˆDŒIà,ˆÔØ$ˆŒØ!2ˆÔØ$8ˆÔ!Ø(ˆÔä�d—l‘l¤DÔ)Ü §¡× 3Ñ 3Ó 5Ó6ˆDŒLð "&§¡ö
Ø´°V¸JÑ5GÓ1HÈ1Ó1LÔQTÐU[Ð\fÑUgÓQhÐklÓQlŠFò
ˆŒð -¹d°°d²
ÈÑKÐPZÐ[_×[dÑ[dÐZeÐeiÐNjÑjˆŒàØØØ�4—9‘9�+ÐØ�D—I‘I�;Ðð
ˆÔð #ˆŒØ % d§i¡i [Ð1ˆÕùò
s   Â;EÃEÃ$Ec                ó(  — |dk7  r|dk(  rd|› �}nd|› d|› d�}nd}| j                   �|d| j                   › d	�z  }t        j                  d
| j                  › d|› d�«       | j	                  |«      }|d   }|d   }|d   }	| j
                  D �
cg c]  }
t        |
d   «      ‘Œ }}
| j
                  D �
cg c]  }
t        |
d   «      ‘Œ }}
t        j                  dt        | j
                  «      › dt        j                  |«      d›dt        j                  |«      d›dt        j                  |«      d›dt        j                  |«      d›dt        j                  |«      d›dt        j                  |«      d›�«       t        j                  d|dz  d›�«       t        j                  d| j                  › d|dz  d›�«       t        j                  d| j                  › d|	dz  d›�«       |�¹| j                  r­t        j                  j                  || j                   «      }t        j                  j#                  |«      }t%        |d|rdndd ¬!«      5 }t'        j(                  |«      }|s|j+                  | j,                  «       |j+                  |||||	g«       ddd«       d|d"| j                  › �|d#| j                  › �|	i}| j/                  || j                  «      }| j1                  ||||«       |S c c}
w c c}
w # 1 sw Y   ŒgxY w)$a  
        Evaluates the model on the dataset and returns the evaluation metrics.

        Args:
            model (SentenceTransformer): The SentenceTransformer model to evaluate.
            output_path (str, optional): The output path to write the results. Defaults to None.
            epoch (int, optional): The current epoch number. Defaults to -1.
            steps (int, optional): The current step number. Defaults to -1.

        Returns:
            Dict[str, float]: A dictionary containing the evaluation metrics.
        éÿÿÿÿz after epoch z
 in epoch z after z stepsr   Nz (truncated to ú)z0RerankingEvaluator: Evaluating the model on the z datasetú:ÚmapÚmrrÚndcgr   r   z	Queries: z 	 Positives: Min z.1fz, Mean z, Max z 	 Negatives: Min zMAP: éd   z.2fr   z: r   ÚaÚwzutf-8)ÚnewlineÚmodeÚencodingzmrr@r   )r&   r   Úinfor   Úcompute_metricesr   r+   ÚnpÚminÚmeanÚmaxr!   r.   ÚosÚpathÚjoinr,   ÚisfileÚopenÚcsvÚwriterÚwriterowr-   Úprefix_name_to_metricsÚ store_metrics_in_model_card_data)r0   ÚmodelÚoutput_pathr   r   Úout_txtÚscoresÚmean_apÚmean_mrrÚ	mean_ndcgr2   Únum_positivesÚnum_negativesÚcsv_pathÚoutput_file_existsÚfrN   Úmetricss                     r4   Ú__call__zRerankingEvaluator.__call__†   s�  € ð �BŠ;Ø˜Š{Ø)¨%¨Ð1‘à& u g¨W°U°G¸6ÐB‘àˆGØ×ÑÐ(Ø˜¨×):Ñ):Ð(;¸1Ð=Ñ=ˆGä�‰ÐFÀtÇyÁyÀkÐQYÐZaÐYbÐbcÐdÔeà×&Ñ& uÓ-ˆØ˜‘-ˆØ˜%‘=ˆØ˜6‘Nˆ	ð @D¿|¹|ÖL°Vœ˜V JÑ/Õ0ÐLˆÐLØ?C¿|¹|ÖL°Vœ˜V JÑ/Õ0ÐLˆÐLä�‰Øœ˜DŸL™LÓ)Ð*Ð*=¼b¿f¹fÀ]Ó>SÐTWÐ=XÐX_Ô`b×`gÑ`gÐhuÓ`vÐwzÐ_{ð  |Bô  CE÷  CIñ  CIð  JWó  CXð  Y\ð  B]ð  ]pô  qs÷  qwñ  qwð  xEó  qFð  GJð  pKð  KRô  SU÷  SZñ  SZð  [hó  Sið  jmð  Rnð  ntô  uw÷  u{ñ  u{ð  |Ió  uJð  KNð  tOð  Pô	
ô 	�‰�e˜G c™M¨#Ð.Ð/Ô0Ü�‰�d˜4Ÿ9™9˜+ R¨°3©°sÐ';Ð<Ô=Ü�‰�e˜DŸI™I˜; b¨°S©¸Ð(=Ð>Ô?ð Ð" t§~¢~Ü—w‘w—|‘| K°·±Ó?ˆHÜ!#§¡§¡°Ó!9ÐÜ�h¨Ñ8J±ÐPSÐ^eÔfð NÐjkÜŸ™ A›�Ù)Ø—O‘O D×$4Ñ$4Ô5à—‘ ¨¨w¸À)Ð LÔM÷Nð �7Ø�4—9‘9�+Ð Ø�D—I‘I�;Ð ð
ˆð
 ×-Ñ-¨g°t·y±yÓAˆØ×-Ñ-¨e°W¸eÀUÔKØˆùò9 MùÚL÷Nð Nús   ÂK>Â5LÉA	LÌLc                ó^   — | j                   r| j                  |«      S | j                  |«      S )a  
        Computes the evaluation metrics for the given model.

        Args:
            model (SentenceTransformer): The SentenceTransformer model to compute metrics for.

        Returns:
            Dict[str, float]: A dictionary containing the evaluation metrics.
        )r%   Úcompute_metrices_batchedÚcompute_metrices_individual)r0   rR   s     r4   rC   z#RerankingEvaluator.compute_metricesÅ   s7   € ð ×(Ò(ð ×)Ñ)¨%Ó0ð	
ð ×1Ñ1°%Ó8ð	
ó    c                ó$  — g }g }g }| j                   €
t        «       n|j                  | j                   «      5  |j                  | j                  D �cg c]  }|d   ‘Œ	 c}d| j
                  | j                  ¬«      }g }| j                  D ]*  }|j                  |d   «       |j                  |d   «       Œ, |j                  |d| j
                  | j                  ¬«      }ddd«       d\  }	}
| j                  D �]=  }|	   }|	dz  }	t        |d   «      }t        |d   «      }|
|
|z   |z    }|
||z   z  }
|d	k(  s|d	k(  rŒH| j                  ||«      }t        |j                  «      dkD  r|d	   }t        j                  | «      }|j                  «       j                  «       }dg|z  d	g|z  z   }d	}t        |d	| j                    «      D ]  \  }}||   sŒd|dz   z  } n |j#                  |«       |j#                  t%        |g|g| j                   ¬
«      «       |j#                  t'        ||«      «       �Œ@ t)        j*                  |«      }t)        j*                  |«      }t)        j*                  |«      }|||dœS c c}w # 1 sw Y   �Œ¦xY w)aE  
        Computes the evaluation metrics in a batched way, by batching all queries and all documents together.

        Args:
            model (SentenceTransformer): The SentenceTransformer model to compute metrics for.

        Returns:
            Dict[str, float]: A dictionary containing the evaluation metrics.
        NÚqueryT©Úconvert_to_tensorr#   r$   r   r   )r   r   é   r   ©Úk©r9   r:   r;   )r&   r   Útruncate_sentence_embeddingsÚencoder   r#   r$   Úextendr+   r"   ÚshapeÚtorchÚargsortÚcpuÚtolistÚ	enumerater!   Úappendr   r   rD   rF   )r0   rR   Úall_mrr_scoresÚall_ndcg_scoresÚall_ap_scoresr2   Úall_query_embsÚall_docsÚall_docs_embsÚ	query_idxÚdocs_idxÚinstanceÚ	query_embÚnum_posÚnum_negÚdocs_embÚpred_scoresÚpred_scores_argsortÚis_relevantÚ	mrr_scoreÚrankÚindexrV   rW   rX   s                            r4   ra   z+RerankingEvaluator.compute_metrices_batchedÕ   s«  € ð ˆØˆØˆà"×/Ñ/Ð7Œ[Œ]¸U×=_Ñ=_Ð`d×`qÑ`qÓ=rñ 	Ø"Ÿ\™\Ø/3¯|©|Ö< V�˜“Ò<Ø"&ØŸ?™?Ø"&×"8Ñ"8ð	 *ó ˆNð ˆHàŸ,™,ò 4�Ø—‘  zÑ 2Ô3Ø—‘  zÑ 2Õ3ð4ð "ŸL™LØ¨D¸T¿_¹_Ð`d×`vÑ`vð )ó ˆM÷	ð& #Ñˆ	�8ØŸ™ó  	TˆHØ& yÑ1ˆIØ˜‰NˆIä˜( :Ñ.Ó/ˆGÜ˜( :Ñ.Ó/ˆGØ$ X°¸7Ñ0BÀWÑ0LÐMˆHØ˜ 'Ñ)Ñ)ˆHà˜!Š|˜w¨!š|Øà×-Ñ-¨i¸ÓBˆKÜ�;×$Ñ$Ó%¨Ò)Ø)¨!™n�ä"'§-¡-°°Ó"=ÐØ%Ÿ/™/Ó+×2Ñ2Ó4ˆKð ˜# ™-¨1¨#°©-Ñ7ˆKØˆIÜ(Ð)<¸QÀÇÁÐ)KÓLò ‘��eØ˜uÓ%Ø ! T¨A¡X¡�IÙðð ×!Ñ! )Ô,ð ×"Ñ"¤:¨{¨m¸k¸]ÈdÏiÉiÔ#XÔYð × Ñ Ô!8¸ÀkÓ!RÖSðA 	TôD —'‘'˜-Ó(ˆÜ—7‘7˜>Ó*ˆÜ—G‘G˜OÓ,ˆ	à x¸ÑCÐCùòq =÷	ñ 	ús   ¸JÁJ 
ÁBJÊ JÊJc                ó¢  — g }g }g }t        j                   | j                  | j                   d¬«      D �]Ç  }|d   }t        |d   «      }t        |d   «      }t	        |«      dk(  st	        |«      dk(  rŒB||z   }	dgt	        |«      z  dgt	        |«      z  z   }
| j
                  €
t        «       n|j                  | j
                  «      5  |j                  |gd	| j                  d
¬«      }|j                  |	d	| j                  d
¬«      }ddd«       | j                  «      }t	        |j                  «      dkD  r|d   }t        j                  | «      }|j                  «       j                  «       }d}t!        |d| j"                   «      D ]  \  }}|
|   sŒd|dz   z  } n |j%                  |«       |j%                  t'        |
g|g| j"                  ¬«      «       |j%                  t)        |
|«      «       �ŒÊ t+        j,                  |«      }t+        j,                  |«      }t+        j,                  |«      }|||dœS # 1 sw Y   �Œ;xY w)aO  
        Computes the evaluation metrics individually by embedding every (query, positive, negative) tuple individually.

        Args:
            model (SentenceTransformer): The SentenceTransformer model to compute metrics for.

        Returns:
            Dict[str, float]: A dictionary containing the evaluation metrics.
        ÚSamples)ÚdisableÚdescre   r   r   r   rh   NTFrf   ri   rk   )Útqdmr   r$   r)   r+   r&   r   rl   rm   r#   r"   ro   rp   rq   rr   rs   rt   r!   ru   r   r   rD   rF   )r0   rR   rv   rw   rx   r~   re   r   r   Údocsr…   r   r‚   rƒ   r„   r†   r‡   rˆ   rV   rW   rX   s                        r4   rb   z.RerankingEvaluator.compute_metrices_individual  sH  € ð ˆØˆØˆäŸ	™	 $§,¡,¸D×<RÑ<RÐ8RÐYbÔcó &	TˆHØ˜WÑ%ˆEÜ˜H ZÑ0Ó1ˆHÜ˜H ZÑ0Ó1ˆHä�8‹} Ò!¤S¨£]°aÒ%7Øà˜hÑ&ˆDØ˜#¤ H£Ñ-°°´c¸(³mÑ0CÑCˆKà"&×"3Ñ"3Ð";””À×AcÑAcÐdh×duÑduÓAvñ Ø!ŸL™LØ�G¨tÀÇÁÐchð )ó �	ð !Ÿ<™<Ø¨D¸T¿_¹_Ð`eð (ó �÷	ð ×-Ñ-¨i¸ÓBˆKÜ�;×$Ñ$Ó%¨Ò)Ø)¨!™n�ä"'§-¡-°°Ó"=ÐØ%Ÿ/™/Ó+×2Ñ2Ó4ˆKð ˆIÜ(Ð)<¸QÀÇÁÐ)KÓLò ‘��eØ˜uÓ%Ø ! T¨A¡X¡�IÙðð ×!Ñ! )Ô,ð ×"Ñ"¤:¨{¨m¸k¸]ÈdÏiÉiÔ#XÔYð × Ñ Ô!8¸ÀkÓ!RÖSðM&	TôP —'‘'˜-Ó(ˆÜ—7‘7˜>Ó*ˆÜ—G‘G˜OÓ,ˆ	à x¸ÑCÐC÷Cñ ús   ÃA IÉI	c                óX   — d| j                   i}| j                  �| j                  |d<   |S )Nr!   r&   )r!   r&   )r0   Úconfig_dicts     r4   Úget_config_dictz"RerankingEvaluator.get_config_dict[  s2   € Ø˜tŸy™yÐ)ˆØ×ÑÐ(Ø*.×*;Ñ*;ˆK˜Ñ'ØÐrc   )r   z list[dict[str, str | list[str]]]r!   Úintr   Ústrr.   Úboolr"   z4Callable[[torch.Tensor, torch.Tensor], torch.Tensor]r#   r’   r$   r”   r%   r”   r&   ú
int | Noner1   r•   )Nr6   r6   )
rR   r   rS   r“   r   r’   r   r’   Úreturnzdict[str, float])Ú__name__Ú
__module__Ú__qualname__Ú__doc__r
   r   r_   rC   ra   rb   r‘   Ú__classcell__)r3   s   @r4   r   r      sÝ   ø„ ñ<ðB ØØØOVØØ"'Ø%)Ø#'Ø#ð.2à1ð.2ð ð.2ð ð	.2ð
 ð.2ð Mð.2ð ð.2ð  ð.2ð #ð.2ð !ð.2ð õ.2ðb bdð=Ø(ð=Ø7:ð=ØJMð=Ø[^ð=à	ó=ò~
ò HDòT:Döxrc   r   )Ú
__future__r   rM   ÚloggingrH   Ú
contextlibr   Útypingr   r   ÚnumpyrD   rp   r�   Úsklearn.metricsr   r   Ú2sentence_transformers.evaluation.SentenceEvaluatorr	   Úsentence_transformers.utilr
   Ú)sentence_transformers.SentenceTransformerr   Ú	getLoggerr—   r   r   © rc   r4   ú<module>r§      sP   ðÝ "ã 
Û Û 	Ý "ß *ã Û Û ß ?å PÝ .áÝMà	ˆ×	Ñ	˜8Ó	$€ôHÐ*õ Hrc   