Ë
    l^(h´  ã                  ó®   — 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Úf1_score)ÚBinaryClassificationEvaluator)ÚSentenceEvaluator)ÚCrossEncoderc                  óX   — e Zd ZdZdddddœ	 	 	 	 	 	 	 	 	 	 	 d	d„Z	 d
	 	 	 	 	 	 	 	 	 dd„Zy)Ú#CrossEncoderClassificationEvaluatora|
  
    Evaluate a CrossEncoder model based on the accuracy of the predicted class vs. the gold labels.
    The evaluator expects a list of sentence pairs and a list of gold labels. If the model has a single output,
    it is assumed to be a binary classification model and the evaluator will calculate accuracy, F1, precision, recall,
    and average precision. If the model has multiple outputs, the evaluator will calculate macro F1, micro F1, and
    weighted F1.

    Args:
        sentence_pairs (List[List[str]]): A list of sentence pairs with each element being a list of two strings.
        labels (List[int]): A list of integers with the gold labels for each sentence pair.
        name (str): Name of the evaluator, useful for the generated model card.
        batch_size (int): Batch size used for the evaluation. Defaults to 32.
        show_progress_bar (bool): Output a progress bar. Defaults to None, which shows the progress bar if the logging level is INFO or DEBUG.
        write_csv (bool): Write results to a CSV file. If a CSV already exists, then values are appended. Defaults to True.

    Example:
        ::

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

            # Load a model
            model = CrossEncoder("cross-encoder/nli-deberta-v3-base")

            # Load a dataset with two text columns and a class label column (https://huggingface.co/datasets/sentence-transformers/all-nli)
            eval_dataset = load_dataset("sentence-transformers/all-nli", "pair-class", split="dev[-1000:]")

            # Create a list of pairs, and map the labels to the labels that the model knows
            pairs = list(zip(eval_dataset["premise"], eval_dataset["hypothesis"]))
            label_mapping = {0: 1, 1: 2, 2: 0}
            labels = [label_mapping[label] for label in eval_dataset["label"]]

            # Initialize the evaluator
            cls_evaluator = CrossEncoderClassificationEvaluator(
                sentence_pairs=pairs,
                labels=labels,
                name="all-nli-dev",
            )
            results = cls_evaluator(model)
            '''
            CrossEncoderClassificationEvaluator: Evaluating the model on all-nli-dev dataset:
            Macro F1:           89.43
            Micro F1:           89.30
            Weighted F1:        89.33
            '''
            print(cls_evaluator.primary_metric)
            # => all-nli-dev_f1_macro
            print(results[cls_evaluator.primary_metric])
            # => 0.8942858180262628
    Ú é    NT)ÚnameÚ
batch_sizeÚshow_progress_barÚ	write_csvc               óV  — t        |«      t        |«      k7  rt        d«      ‚|| _        t        j                  |«      | _        || _        || _        |€4t        j                  «       t        j                  t        j                  fv }|| _        d|rd|z   ndz   dz   | _        || _        y )Nz3sentence_pairs and labels must have the same lengthr   Ú_r   z_results.csv)ÚlenÚ
ValueErrorÚsentence_pairsÚnpÚasarrayÚlabelsr   r   ÚloggerÚgetEffectiveLevelÚloggingÚINFOÚDEBUGr   Úcsv_filer   )Úselfr   r   r   r   r   r   Úkwargss           ú{/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/sentence_transformers/cross_encoder/evaluation/classification.pyÚ__init__z,CrossEncoderClassificationEvaluator.__init__I   s˜   € ô ˆ~Ó¤# f£+Ò-ÜÐRÓSÐSà,ˆÔÜ—j‘j Ó(ˆŒØˆŒ	Ø$ˆŒØÐ$Ü &× 8Ñ 8Ó :¼w¿|¹|ÌWÏ]É]Ð>[Ð [ÐØ!2ˆÔà=ÉtÀÀtÂÐY[Ñ\Ð_mÑmˆŒØ"ˆ�ó    c                ó´  — |dk7  r|dk(  rd|› �}nd|› d|› d�}nd}t         j                  d| j                  › d|› d	�«       |j                  | j                  d
| j
                  ¬«      }|j                  dk(  �rt        j                  || j                  d
«      \  }}t        j                  || j                  d
«      \  }	}
}}t        | j                  |«      }t         j                  d|dz  d›d|d›d�«       t         j                  d|	dz  d›d|d›d�«       t         j                  d|
dz  d›�«       t         j                  d|dz  d›�«       t         j                  d|dz  d›�«       |||	||
||dœ}g d¢| _        d| _        nÉt        j                  |d¬«      }t!        | j                  |d¬«      }t!        | j                  |d¬«      }t!        | j                  |d¬«      }t         j                  d|dz  d›�«       t         j                  d |dz  d›�«       t         j                  d!|dz  d›�«       |||d"œ}g d#¢| _        d$| _        |�Å| j"                  r¹t$        j&                  j)                  || j*                  «      }t$        j&                  j-                  |«      }t/        ||rd%nd&d'¬(«      5 }t1        j2                  |«      }|s|j5                  | j                  «       |j5                  ||g|j7                  «       ¢«       d d d «       | j9                  || j                  «      }| j;                  ||||«       |S # 1 sw Y   Œ;xY w))Néÿÿÿÿz after epoch z
 in epoch z after z stepsr   z=CrossEncoderClassificationEvaluator: Evaluating the model on z datasetú:T)Úconvert_to_numpyr   é   zAccuracy:          éd   z.2fz	(Threshold: z.4fú)zF1:                zPrecision:         zRecall:            zAverage Precision: )ÚaccuracyÚaccuracy_thresholdÚf1Úf1_thresholdÚ	precisionÚrecallÚaverage_precision)	ÚepochÚstepsÚAccuracyÚAccuracy_ThresholdÚF1ÚF1_ThresholdÚ	PrecisionÚRecallÚAverage_Precisionr2   )ÚaxisÚmacro)ÚaverageÚmicroÚweightedzMacro F1:           zMicro F1:           zWeighted F1:        )Úf1_macroÚf1_microÚf1_weighted)r3   r4   ÚMacro_F1ÚMicro_F1ÚWeighted_F1rA   ÚaÚwzutf-8)ÚmodeÚencoding)r   Úinfor   Úpredictr   r   Ú
num_labelsr   Úfind_best_acc_and_thresholdr   Úfind_best_f1_and_thresholdr   Úcsv_headersÚprimary_metricr   Úargmaxr   r   ÚosÚpathÚjoinr   ÚisfileÚopenÚcsvÚwriterÚwriterowÚvaluesÚprefix_name_to_metricsÚ store_metrics_in_model_card_data)r    ÚmodelÚoutput_pathr3   r4   Úout_txtÚpred_scoresÚaccÚacc_thresholdr.   r0   r1   r/   ÚapÚmetricsÚpred_labelsrA   rB   rC   Úcsv_pathÚoutput_file_existsÚfrY   s                          r"   Ú__call__z,CrossEncoderClassificationEvaluator.__call__b   sP  € ð �BŠ;Ø˜Š{Ø)¨%¨Ð1‘à& u g¨W°U°G¸6ÐB‘àˆGä�‰ÐSÐTX×T]ÑT]ÐS^Ð^fÐgnÐfoÐopÐqÔrØ—m‘mØ×Ñ°$È$×J`ÑJ`ð $ó 
ˆð ×Ñ˜qÓ Ü!>×!ZÑ!ZØ˜TŸ[™[¨$ó"ÑˆC�ô 3P×2jÑ2jØ˜TŸ[™[¨$ó3Ñ/ˆB�	˜6 <ô )¨¯©°kÓBˆBä�K‰KÐ-¨c°C©i¸¨_¸NÈ=ÐY\ÐJ]Ð]^Ð_Ô`Ü�K‰KÐ-¨b°3©h°s¨^¸>È,ÐWZÐI[Ð[\Ð]Ô^Ü�K‰KÐ-¨i¸#©o¸cÐ-BÐCÔDÜ�K‰KÐ-¨f°s©l¸3Ð-?Ð@ÔAÜ�K‰KÐ-¨b°3©h°s¨^Ð<Ô=ð  Ø&3ØØ ,Ø&Ø Ø%'ñˆGò
 ˆDÔð #6ˆDÕäŸ)™) K°aÔ8ˆKÜ §¡¨[À'ÔJˆHÜ §¡¨[À'ÔJˆHÜ" 4§;¡;°ÀZÔPˆKä�K‰KÐ.¨x¸#©~¸cÐ.BÐCÔDÜ�K‰KÐ.¨x¸#©~¸cÐ.BÐCÔDÜ�K‰KÐ.¨{¸SÑ/@ÀÐ.EÐFÔGð %Ø$Ø*ñˆGò
  YˆDÔØ",ˆDÔàÐ" t§~¢~Ü—w‘w—|‘| K°·±Ó?ˆHÜ!#§¡§¡°Ó!9ÐÜ�hÑ,>¡SÀCÐRYÔZð CÐ^_ÜŸ™ A›�Ù)Ø—O‘O D×$4Ñ$4Ô5à—‘ ¨Ð A°·±Ó0@Ð AÔB÷Cð ×-Ñ-¨g°t·y±yÓAˆØ×-Ñ-¨e°W¸eÀUÔKØˆ÷Cð Cús   Ê>AMÍM)r   zlist[list[str]]r   z	list[int]r   Ústrr   Úintr   zbool | Noner   Úbool)Nr&   r&   )
r^   r	   r_   rk   r3   rl   r4   rl   Úreturnzdict[str, float])Ú__name__Ú
__module__Ú__qualname__Ú__doc__r#   rj   © r$   r"   r   r      s�   „ ñ2ðr ØØ)-Øñ#à'ð#ð ð#ð
 ð#ð ð#ð 'ð#ð ó#ð4 []ðRØ!ðRØ03ðRØCFðRØTWðRà	ôRr$   r   )Ú
__future__r   rX   r   rS   Útypingr   Únumpyr   Úsklearn.metricsr   r   Ú sentence_transformers.evaluationr   Ú2sentence_transformers.evaluation.SentenceEvaluatorr   Ú0sentence_transformers.cross_encoder.CrossEncoderr	   Ú	getLoggerro   r   r   rs   r$   r"   ú<module>r|      sG   ðÝ "ã 
Û Û 	Ý  ã ß =å JÝ PáÝMà	ˆ×	Ñ	˜8Ó	$€ô`Ð*;õ `r$   