Ë
    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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)ÚSentenceEvaluator)ÚSentenceTransformerc                  ón   ‡ — e Zd ZdZ	 	 	 	 	 	 d	 	 	 	 	 	 	 	 	 	 	 	 	 dˆ fd„Zdd	d„Zed
d„«       Zd„ Zˆ xZ	S )ÚMSEEvaluatora
  
    Computes the mean squared error (x100) between the computed sentence embedding
    and some target sentence embedding.

    The MSE is computed between ||teacher.encode(source_sentences) - student.encode(target_sentences)||.

    For multilingual knowledge distillation (https://arxiv.org/abs/2004.09813), source_sentences are in English
    and target_sentences are in a different language like German, Chinese, Spanish...

    Args:
        source_sentences (List[str]): Source sentences to embed with the teacher model.
        target_sentences (List[str]): Target sentences to embed with the student model.
        teacher_model (SentenceTransformer, optional): The teacher model to compute the source sentence embeddings.
        show_progress_bar (bool, optional): Show progress bar when computing embeddings. Defaults to False.
        batch_size (int, optional): Batch size to compute sentence embeddings. Defaults to 32.
        name (str, optional): Name of the evaluator. Defaults to "".
        write_csv (bool, optional): Write results to CSV file. Defaults to True.
        truncate_dim (int, optional): The dimension to truncate sentence embeddings to. `None` uses the model's current truncation
            dimension. Defaults to None.

    Example:
        ::

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

            # Load a model
            student_model = SentenceTransformer('paraphrase-multilingual-mpnet-base-v2')
            teacher_model = SentenceTransformer('all-mpnet-base-v2')

            # Load any dataset with some texts
            dataset = load_dataset("sentence-transformers/stsb", split="validation")
            sentences = dataset["sentence1"] + dataset["sentence2"]

            # Given queries, a corpus and a mapping with relevant documents, the InformationRetrievalEvaluator computes different IR metrics.
            mse_evaluator = MSEEvaluator(
                source_sentences=sentences,
                target_sentences=sentences,
                teacher_model=teacher_model,
                name="stsb-dev",
            )
            results = mse_evaluator(student_model)
            '''
            MSE evaluation (lower = better) on the stsb-dev dataset:
            MSE (*100):  0.805045
            '''
            print(mse_evaluator.primary_metric)
            # => "stsb-dev_negative_mse"
            print(results[mse_evaluator.primary_metric])
            # => -0.8050452917814255
    c	                óp  •— t         ‰	| �  «        || _        | j                  €
t        «       n|j	                  | j                  «      5  |j                  |||d¬«      | _        d d d «       || _        || _        || _	        || _
        d|z   dz   | _        g d¢| _        || _        d| _        y # 1 sw Y   ŒJxY w)NT©Úshow_progress_barÚ
batch_sizeÚconvert_to_numpyÚmse_evaluation_z_results.csv)ÚepochÚstepsÚMSEÚnegative_mse)ÚsuperÚ__init__Útruncate_dimr   Útruncate_sentence_embeddingsÚencodeÚsource_embeddingsÚtarget_sentencesr   r   ÚnameÚcsv_fileÚcsv_headersÚ	write_csvÚprimary_metric)
ÚselfÚsource_sentencesr   Úteacher_modelr   r   r   r   r   Ú	__class__s
            €úk/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/sentence_transformers/evaluation/MSEEvaluator.pyr   zMSEEvaluator.__init__G   sÄ   ø€ ô 	‰ÑÔØ(ˆÔð × Ñ Ð(ô ŒMà×;Ñ;¸D×<MÑ<MÓNñ	ð
 &3×%9Ñ%9Ø Ð4EÐR\Ðosð &:ó &ˆDÔ"÷	ð !1ˆÔØ!2ˆÔØ$ˆŒØˆŒ	à)¨DÑ0°>ÑAˆŒÚ4ˆÔØ"ˆŒØ,ˆÕ÷#	ð 	ús   ÁB,Â,B5c                ó,  — |dk7  r|dk(  rd|› �}nd|› d|› d�}nd}| j                   �|d| j                   › d�z  }| j                   €
t        «       n|j                  | j                   «      5  |j                  | j                  | j
                  | j                  d	¬
«      }d d d «       | j                  z
  dz  j                  «       }|dz  }t        j                  d| j                  › d|› d�«       t        j                  d|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| i}| j/                  || j                  «      }| j1                  ||||«       |S # 1 sw Y   �Œ^xY w# 1 sw Y   ŒMxY w)Néÿÿÿÿz after epoch z
 in epoch z after z stepsÚ z (truncated to ú)Tr   é   éd   z'MSE evaluation (lower = better) on the z datasetú:zMSE (*100):	Ú4fÚaÚwzutf-8)ÚnewlineÚmodeÚencodingr   )r   r   r   r   r   r   r   r   ÚmeanÚloggerÚinfor   r   ÚosÚpathÚjoinr   ÚisfileÚopenÚcsvÚwriterÚwriterowr   Úprefix_name_to_metricsÚ store_metrics_in_model_card_data)r    ÚmodelÚoutput_pathr   r   Úout_txtÚtarget_embeddingsÚmseÚcsv_pathÚoutput_file_existsÚfr;   Úmetricss                r$   Ú__call__zMSEEvaluator.__call__g   sü  € Ø�BŠ;Ø˜Š{Ø)¨%¨Ð1‘à& u g¨W°U°G¸6ÐB‘àˆGØ×ÑÐ(Ø˜¨×):Ñ):Ð(;¸1Ð=Ñ=ˆGà"×/Ñ/Ð7Œ[Œ]¸U×=_Ñ=_Ð`d×`qÑ`qÓ=rñ 	Ø %§¡Ø×%Ñ%Ø"&×"8Ñ"8ØŸ?™?Ø!%ð	 !-ó !Ð÷	ð ×&Ñ&Ð):Ñ:¸qÑ@×FÑFÓHˆØˆs‰
ˆä�‰Ð=¸d¿i¹i¸[ÈÐQXÐPYÐYZÐ[Ô\Ü�‰�m C¨ 8Ð,Ô-àÐ" t§~¢~Ü—w‘w—|‘| K°·±Ó?ˆHÜ!#§¡§¡°Ó!9ÐÜ�h¨Ñ8J±ÐPSÐ^eÔfð 5ÐjkÜŸ™ A›�Ù)Ø—O‘O D×$4Ñ$4Ô5à—‘ ¨¨sÐ 3Ô4÷5ð " C 4Ð(ˆØ×-Ñ-¨g°t·y±yÓAˆØ×-Ñ-¨e°W¸eÀUÔKØˆ÷9	ñ 	ú÷"5ð 5ús   Á-4G=Å7AH
Ç=HÈ
Hc                 ó   — y)NzKnowledge Distillation© )r    s    r$   ÚdescriptionzMSEEvaluator.description�   s   € à'ó    c                ó@   — i }| j                   �| j                   |d<   |S )Nr   )r   )r    Úconfig_dicts     r$   Úget_config_dictzMSEEvaluator.get_config_dict”   s)   € ØˆØ×ÑÐ(Ø*.×*;Ñ*;ˆK˜Ñ'ØÐrL   )NFé    r'   TN)r!   ú	list[str]r   rQ   r   Úboolr   Úintr   Ústrr   rR   r   z
int | None)Nr&   r&   )r?   r   r@   rT   Úreturnzdict[str, float])rU   rT   )
Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   rH   ÚpropertyrK   rO   Ú__classcell__)r#   s   @r$   r	   r	      s†   ø„ ñ3ðr Ø"'ØØØØ#'ð-à#ð-ð $ð-ð
  ð-ð ð-ð ð-ð ð-ð !õ-ô@'ðR ò(ó ð(örL   r	   )Ú
__future__r   r:   Úloggingr5   Ú
contextlibr   Útypingr   Ú2sentence_transformers.evaluation.SentenceEvaluatorr   Ú)sentence_transformers.SentenceTransformerr   Ú	getLoggerrV   r3   r	   rJ   rL   r$   ú<module>rc      sA   ðÝ "ã 
Û Û 	Ý "Ý  å PáÝMà	ˆ×	Ñ	˜8Ó	$€ôGÐ$õ GrL   