Ë
    þÍ:j�&  ã                   óŽ  — d dl Z d dlZd dlmZ d dlmZmZmZ d dlm	Z	m
Z
 d dlmZ eeef   Zeeee   f   Zeeeeeeeee   ee   f   f   f   f   Zeeee   f   Zeeeeeeef   f      eeeeee   ee   f   f   f   Zdgdgdœd	d
dddœZdedefd„Zdedee   fd„Zdedede	fd„Zdedede	fd„Zdeeege	f   dedee   de	fd„Zdededeeeef   eeeeeeeeeef      f      f      f   fd„Zdeeef   deeeeeeeeeef      f      f      dee	e	e	f   fd„Zd e	d!e	d"e	deee	f   fd#„Zdededeee	f   fd$„Z y)%é    N)ÚCounter)ÚAnyÚCallableÚUnion)ÚTensorÚtensor)Úrank_zero_warné   zThis is a test text)Úanswer_startÚtextzThis is a test context.Ú1zIs this a test?z
train test)ÚanswersÚcontextÚidÚquestionÚtitleÚsÚreturnc           	      ó¶   — dt         dt         fd„}dt         dt         fd„}dt         dt         fd„}dt         dt         fd„} | | | || «      «      «      «      S )zALower text and remove punctuation, articles and extra whitespace.r   r   c                 ó0   — t        j                  dd| «      S )Nz\b(a|an|the)\bú )ÚreÚsub©r   s    úw/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/text/squad.pyÚremove_articlesz(_normalize_text.<locals>.remove_articles,   s   € Ü�v‰vÐ'¨¨dÓ3Ð3ó    c                 ó@   — dj                  | j                  «       «      S )Nr   )ÚjoinÚsplitr   s    r   Úwhite_space_fixz(_normalize_text.<locals>.white_space_fix/   s   € Ø�x‰x˜Ÿ
™
›Ó%Ð%r   c                 ój   ‡— t        t        j                  «      Šdj                  ˆfd„| D «       «      S )NÚ c              3   ó,   •K  — | ]  }|‰vsŒ|–— Œ y ­w©N© )Ú.0ÚchÚexcludes     €r   ú	<genexpr>z7_normalize_text.<locals>.remove_punc.<locals>.<genexpr>4   s   øè ø€ Ò>˜b¨B°gÒ,=”rÑ>ùs   ƒ	�)ÚsetÚstringÚpunctuationr   )r   r)   s    @r   Úremove_puncz$_normalize_text.<locals>.remove_punc2   s(   ø€ Ü”f×(Ñ(Ó)ˆØ�w‰wÓ> DÔ>Ó>Ð>r   c                 ó"   — | j                  «       S r%   )Úlowerr   s    r   r0   z_normalize_text.<locals>.lower6   s   € Ø�z‰z‹|Ðr   )Ústr)r   r   r!   r.   r0   s        r   Ú_normalize_textr2   )   si   € ð4œcð 4¤có 4ð&œcð &¤có &ð?œ#ð ?¤#ó ?ð”Cð œCó ñ ™?©;±u¸Q³xÓ+@ÓAÓBÐBr   c                 ó<   — | sg S t        | «      j                  «       S )z&Split a sentence into separate tokens.)r2   r    )r   s    r   Ú_get_tokensr4   <   s   € áˆ2Ð6œO¨AÓ.×4Ñ4Ó6Ð6r   Úpredicted_answerÚtarget_answerc                 óª  — t        |«      }t        | «      }t        |«      t        |«      z  }t        t        |j	                  «       «      «      }t        |«      dk(  st        |«      dk(  rt        t        ||k(  «      «      S |dk(  rt        d«      S d|z  t        t        |«      «      z  }d|z  t        t        |«      «      z  }d|z  |z  ||z   z  S )z#Compute F1 Score for two sentences.r   ç        g      ð?é   )r4   r   r   ÚsumÚvaluesÚlenÚint)r5   r6   Útarget_tokensÚpredicted_tokensÚcommonÚnum_sameÚ	precisionÚrecalls           r   Ú_compute_f1_scorerD   A   sÍ   € ä Ó.€MÜ"Ð#3Ó4ÐÜ�]Ó#¤gÐ.>Ó&?Ñ?€FÜ”c˜&Ÿ-™-›/Ó*Ó+€HÜ
ˆ=Ó˜QÒ¤#Ð&6Ó"7¸1Ò"<ä”c˜-Ð+;Ñ;Ó<Ó=Ð=Ø�1‚}Ü�c‹{ÐØ�h‘¤¬Ð,<Ó(=Ó!>Ñ>€IØ�8‰^œf¤S¨Ó%7Ó8Ñ8€FØ�	‰M˜FÑ" y°6Ñ'9Ñ:Ð:r   Ú
predictionÚground_truthc                 óT   — t        t        t        | «      t        |«      k(  «      «      S )z&Compute Exact Match for two sentences.)r   r=   r2   )rE   rF   s     r   Ú_compute_exact_match_scorerH   Q   s!   € ä”#”o jÓ1´_À\Ó5RÑRÓSÓTÐTr   Ú	metric_fnÚground_truthsc                 ó0   ‡ ‡— t        ˆ ˆfd„|D «       «      S )zJCalculate maximum score for a predicted answer with all reference answers.c              3   ó0   •K  — | ]  } ‰‰|«      –— Œ y ­wr%   r&   )r'   ÚtruthrI   rE   s     €€r   r*   z1_metric_max_over_ground_truths.<locals>.<genexpr>Z   s   øè ø€ ÒG°‰y˜ U×+ÑGùs   ƒ)Úmax)rI   rE   rJ   s   `` r   Ú_metric_max_over_ground_truthsrO   V   s   ù€ ô ÔG¸ÔGÓGÐGr   ÚpredsÚtargetsc                 óÀ  — t        | t        «      r| g} t        |t        «      r|g}| D ]%  }|j                  «       }d|vsd|vsŒt        d«      ‚ |D ]G  }|j                  «       }d|vsd|vrt        dt        › �«      ‚|d   }d|vsŒ7t        dt        › �«      ‚ | D �ci c]  }|d   |d   “Œ }}d„ }	d	d
|D �cg c]
  } |	|«      ‘Œ c}igig}
||
fS c c}w c c}w )zOCheck for types and convert the input to necessary format to compute the input.Úprediction_textr   z¦Expected keys in a single prediction are 'prediction_text' and 'id'.Please make sure that 'prediction_text' maps to the answer string and 'id' maps to the key string.r   z«Expected keys in a single target are 'answers' and 'id'.Please make sure that 'answers' maps to a `SQuAD` format dictionary and 'id' maps to the key string.
SQuAD Format: r   zzExpected keys in a 'answers' are 'text'.Please make sure that 'answer' maps to a `SQuAD` format dictionary.
SQuAD Format: c                 óH   — | d   d   D �cg c]  }d|i‘Œ c}| d   dœS c c}w )Nr   r   r   )r   r   r&   )ÚtgtÚtxts     r   ú<lambda>z$_squad_input_check.<locals>.<lambda>ƒ   s.   € À3ÀyÁ>ÐRXÑCYÖ)Z¸C¨6°3ª-Ò)ZÐbeÐfjÑbkÑl€ ùÒ)Zs   ‹Ú
paragraphsÚqas)Ú
isinstanceÚdictÚkeysÚKeyErrorÚSQuAD_FORMAT)rP   rQ   ÚpredÚ	pred_keysÚtargetÚtarget_keysr   rE   Ú
preds_dictÚ
_fn_answerÚtargets_dicts              r   Ú_squad_input_checkrf   ]   sC  € ô �%œÔØ�ˆä�'œ4Ô Ø�)ˆàò ˆØ—I‘I“Kˆ	Ø IÑ-°¸YÒ1FÜðuóð ðð ò ˆØ—k‘k“mˆØ˜KÑ'¨4°{Ñ+BÜð!ô  �.ð"óð ð ;AÀÑ:KˆØ˜Ò Üð!ô  �.ð"óð ðð& UZÖZÀj�*˜TÑ" JÐ/@Ñ$AÑAÐZ€JÐZÙl€JØ! UÈgÖ,VÀF©Z¸Õ-?Ò,VÐ$WÐ#XÐYÐZ€LØ�|Ð#Ð#ùò [ùâ,Vs   Â!CÂ<Cra   c           	      óV  — t        d«      }t        d«      }t        d«      }|D ]z  }|d   D ]p  }|d   D ]f  }|dz  }|d   | vrt        d|d   › d�«       Œ"|d	   D �cg c]  }|d
   ‘Œ	 }	}| |d      }
|t        t        |
|	«      z  }|t        t        |
|	«      z  }Œh Œr Œ| |||fS c c}w )au  Compute F1 Score and Exact Match for a collection of predictions and references.

    Args:
        preds: A dictionary mapping an `id` to the predicted `answer`.
        target:
            A list of dictionary mapping `paragraphs` to list of dictionary mapping `qas` to a list of dictionary
            containing `id` and list of all possible `answers`.

    Return:
        Tuple containing F1 score, Exact match score and total number of examples.

    Example:
        >>> from torchmetrics.functional.text.squad import _squad_update
        >>> preds = [{"prediction_text": "1976", "id": "56e10a3be3433e1400422b22"}]
        >>> target = [{"answers": {"answer_start": [97], "text": ["1976"]}, "id": "56e10a3be3433e1400422b22"}]
        >>> preds_dict = {pred["id"]: pred["prediction_text"] for pred in preds}
        >>> targets_dict = [
        ...     dict(paragraphs=[dict(qas=[dict(answers=[
        ...         {"text": txt} for txt in tgt["answers"]["text"]], id=tgt["id"]) for tgt in target
        ...     ])])
        ... ]
        >>> _squad_update(preds_dict, targets_dict)
        (tensor(1.), tensor(1.), tensor(1))

    r8   r   rX   rY   r
   r   zUnanswered question z will receive score 0.r   r   )r   r	   rO   rH   rD   )rP   ra   Úf1Úexact_matchÚtotalÚarticleÚ	paragraphÚqaÚxrJ   r_   s              r   Ú_squad_updatero   ˆ   sô   € ô: 
�‹€BÜ˜“+€KÜ�1‹I€EØò 
]ˆØ  Ñ.ò 		]ˆIØ Ñ&ò ]�Ø˜‘
�Ø�d‘8 5Ñ(Ü"Ð%9¸"¸T¹(¸ÐCYÐ#ZÔ[ØØ46°y±MÖ B¨q  6£Ð B�Ð BØ˜R ™X‘�ØÔ=Ô>XÐZ^Ð`mÓnÑn�ØÔ4Ô5FÈÈmÓ\Ñ\‘ñ]ñ		]ð
]ð ˆ{˜EÐ!Ð!ùò !Cs   ÁB&
rh   ri   rj   c                 ó,   — d|z  |z  }d| z  |z  } || dœS )z•Aggregate the F1 Score and Exact match for the batch.

    Return:
        Dictionary containing the F1 score, Exact match score for the batch.

    g      Y@)ri   rh   r&   )rh   ri   rj   s      r   Ú_squad_computerq   ·   s,   € ð ˜+Ñ%¨Ñ-€KØ	�‰�eÑ	€BØ&¨bÑ1Ð1r   c                 óZ   — t        | |«      \  }}t        ||«      \  }}}t        |||«      S )aã  Calculate `SQuAD Metric`_ .

    Args:
        preds: A Dictionary or List of Dictionary-s that map `id` and `prediction_text` to the respective values.

            Example prediction:

            .. code-block:: python

                {"prediction_text": "TorchMetrics is awesome", "id": "123"}

        target: A Dictionary or List of Dictionary-s that contain the `answers` and `id` in the SQuAD Format.

            Example target:

            .. code-block:: python

                {
                    'answers': [{'answer_start': [1], 'text': ['This is a test answer']}],
                    'id': '1',
                }

            Reference SQuAD Format:

            .. code-block:: python

                {
                    'answers': {'answer_start': [1], 'text': ['This is a test text']},
                    'context': 'This is a test context.',
                    'id': '1',
                    'question': 'Is this a test?',
                    'title': 'train test'
                }


    Return:
        Dictionary containing the F1 score, Exact match score for the batch.

    Example:
        >>> from torchmetrics.functional.text.squad import squad
        >>> preds = [{"prediction_text": "1976", "id": "56e10a3be3433e1400422b22"}]
        >>> target = [{"answers": {"answer_start": [97], "text": ["1976"]},"id": "56e10a3be3433e1400422b22"}]
        >>> squad(preds, target)
        {'exact_match': tensor(100.), 'f1': tensor(100.)}

    Raises:
        KeyError:
            If the required keys are missing in either predictions or targets.

    References:
        [1] SQuAD: 100,000+ Questions for Machine Comprehension of Text by Pranav Rajpurkar, Jian Zhang, Konstantin
        Lopyrev, Percy Liang `SQuAD Metric`_ .

    )rf   ro   rq   )rP   ra   rc   Útarget_dictrh   ri   rj   s          r   Úsquadrt   Ã   s8   € ôn 1°¸Ó?Ñ€J�Ü*¨:°{ÓCÑ€Bˆ�UÜ˜"˜k¨5Ó1Ð1r   )!r   r,   Úcollectionsr   Útypingr   r   r   Útorchr   r   Útorchmetrics.utilitiesr	   r[   r1   ÚSINGLE_PRED_TYPEÚlistÚ
PREDS_TYPEr=   ÚSINGLE_TARGET_TYPEÚTARGETS_TYPEÚUPDATE_METHOD_SINGLE_PRED_TYPEr^   r2   r4   rD   rH   rO   Útuplerf   ro   rq   rt   r&   r   r   ú<module>r€      s¦  ðó" 
Û Ý ß 'Ñ 'ç  å 1à˜˜S˜‘>Ð ØÐ# TÐ*:Ñ%;Ð;Ñ<€
Ø˜#˜u S¨$¨s°E¸$¸s¹)ÀTÈ#ÁYÐ:NÑ4OÐ/OÑ*PÐ%PÑQÐQÑRÐ ØÐ'¨Ð.@Ñ)AÐAÑB€Ø!& t¨D°°e¸CÀ¸H±oÐ1EÑ,FÑ'GÈÈdÐSVÐX]Ð^bÐcfÑ^gÐimÐnqÑirÐ^rÑXsÐSsÑNtÐ'tÑ!uÐ ð "# Ð.CÐ-DÑEØ(Ø
Ø!Øñ€ðC�sð C˜só Cð&7�3ð 7˜4 ™9ó 7ð
;¨ð ;¸Cð ;ÀFó ;ð U¨3ð U¸cð UÀfó Uð
HØ˜˜c˜
 FÐ*Ñ+ðHØ9<ðHØMQÐRUÉYðHàóHð($Øð($Ø ,ð($à
ˆ4��S�‰>˜4  S¨$¨t°C¸¸dÀ3ÈÀ8¹nÑ9MÐ4MÑ/NÑ*OÐ%OÑ PÑQÐQÑRó($ðV,"Ø��S�‰>ð,"à��c˜4  S¨$¨t°C¸°H©~Ñ*>Ð%>Ñ ?Ñ@Ð@ÑAÑBð,"ð ˆ6�6˜6Ð!Ñ"ó,"ð^	2�vð 	2¨Fð 	2¸6ð 	2ÀdÈ3ÐPVÈ;ÑFWó 	2ð92�ð 92 \ð 92°d¸3À¸;Ñ6Gô 92r   