Ë
    þÍ:j{m  ã            #       óN  — d dl Z d dlmZ d dlmZ d dlmZmZmZm	Z	 d dl
Z
d dl
mZ d dlmZ d dlmZ d dlmZ d d	lmZmZmZmZ d d
lmZ d dlmZ er
erd dlmZmZ esdgZed   Z e G d„ de«      «       Z! G d„ d«      Z"dedede#de$de$defd„Z%ddde&e'e$f   fd„Z(dede$de$de$def
d „Z)d!d"d#e&e'ef   d$e*de#d%e&e'e$f   defd&„Z+ e
jX                  «       d!d"d'ed$e*de#d%e&e'e$f   d(e#defd)„«       Z-d*e	e'ee'   f   d+e	e'ee'   f   ddd,e$de.eeeef   f
d-„Z/	 d:d!d"d.ed/ed$e*de#d0e"d%e&e'e$f   d(e#defd1„Z0	 	 	 	 	 	 	 	 	 	 	 	 d;d*e	e'ee'   f   d+e	e'ee'   f   d2e	e'e jb                  f   d$e*d3e de#d4ee*   d5ee*   d6ee	e'e
jd                  f      d,ee$   de$d7e$d(e#d8e#de	ee.eef   f   fd9„Z3y)<é    N)ÚSequence)Úunique)ÚTYPE_CHECKINGÚListÚOptionalÚUnion)ÚTensor)Ú
functional)Ú
DataLoader)ÚLiteral)ÚTokenizedDatasetÚ_get_progress_barÚ_input_data_collatorÚ_load_tokenizer_and_model)ÚEnumStr)Ú_TRANSFORMERS_GREATER_EQUAL_4_4)ÚPreTrainedModelÚPreTrainedTokenizerBaseÚinfolm)	Úkl_divergenceÚalpha_divergenceÚbeta_divergenceÚab_divergenceÚrenyi_divergenceÚl1_distanceÚl2_distanceÚl_infinity_distanceÚfisher_rao_distancec                   óJ   — e Zd ZdZedefd„«       ZdZdZdZ	dZ
dZd	Zd
ZdZdZy)Ú_IMEnumz8A helper Enum class for storing the information measure.Úreturnc                   ó   — y)NzInformation measure© r#   ó    úx/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/text/infolm.pyÚ_namez_IMEnum._name:   s   € à$r$   r   r   r   r   r   r   r   r   r   N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚstaticmethodÚstrr&   ÚKL_DIVERGENCEÚALPHA_DIVERGENCEÚBETA_DIVERGENCEÚAB_DIVERGENCEÚRENYI_DIVERGENCEÚL1_DISTANCEÚL2_DISTANCEÚL_INFINITY_DISTANCEÚFISHER_RAO_DISTANCEr#   r$   r%   r    r    6   sQ   „ áBàð%�3ò %ó ð%ð $€MØ)ÐØ'€OØ#€MØ)ÐØ€KØ€KØ/ÐØ/Ñr$   r    c            	       ó4  — e Zd ZdZ	 	 ddedee   dee   ddfd„Zded	edefd
„Z	e
ded	edefd„«       Zded	edefd„Zded	edefd„Zded	edefd„Zded	edefd„Ze
ded	edefd„«       Ze
ded	edefd„«       Ze
ded	edefd„«       Ze
ded	edefd„«       Zy)Ú_InformationMeasureu  A wrapper class used for the calculation of different information measures.

    This metric can be used to measure the information between the discrete reference distributions of predicted and
    reference sentences. The class also handles input validation for `alpha` and `beta` parameters.

    Args:
        information_measure:
            A name of information measure to be used. Please use one of: ['kl_divergence', 'alpha_divergence',
            'beta_divergence', 'ab_divergence', 'renyi_divergence', 'l1_distance', 'l2_distance', 'l_infinity_distance',
            'fisher_rao_distance']
        alpha:
            Alpha parameter of the divergence used for alpha, AB and RÃ©nyi divergence measures.
        beta:
            Beta parameter of the divergence used for beta and AB divergence measures.

    Raises:
        ValueError:
            If information measure is one from alpha, AB or RÃ©nyi divergence and parameter `alpha` is `None`.
        ValueError:
            If information measure is one from beta or divergence and parameter `beta` is `None`.
        ValueError:
            If information measure is alpha divergence and parameter `alpha` equals 0 or 1.
        ValueError:
            If information measure is beta divergence and parameter `beta` equals 0 or -1.
        ValueError:
            If information measure is AB divergence and parameter `alpha`, `beta` or `alpha + beta` equal 0.
        ValueError:
            If information measure is RÃ©nyi divergence and parameter `alpha` equals 1.

    NÚinformation_measureÚalphaÚbetar!   c                 óÐ  — t         j                  |«      | _        t         j                  t         j                  t         j
                  f}| j                  |v rt        |t        «      st        d|› d�«      ‚| j                  t         j                  t         j                  fv rt        |t        «      st        d|› d�«      ‚| j                  t         j                  k(  r#t        |t        «      r|dv rt        d|› d�«      ‚| j                  t         j                  k(  r#t        |t        «      r|dv rt        d|› d�«      ‚| j                  t         j                  k(  r1|� |�t        d„ ||fD «       «      s
d	||||z   fv rt        d
|› d�«      ‚| j                  t         j
                  k(  r$t        |t        «      r|dk(  rt        d|› d�«      ‚|xs d	| _        |xs d	| _        y )Nz0Parameter `alpha` is expected to be defined for ú.z/Parameter `beta` is expected to be defined for )r   é   zFParameter `alpha` is expected to be float differened from 0 and 1 for )r   éÿÿÿÿzFParameter `beta` is expected to be float differened from 0 and -1 for c              3   ó>   K  — | ]  }t        |t        «       –— Œ y ­w)N)Ú
isinstanceÚfloat)Ú.0Úps     r%   ú	<genexpr>z/_InformationMeasure.__init__.<locals>.<genexpr>€   s   è ø€ ÒD°œ
 1¤eÓ,Ô,ÑDùs   ‚r   zRParameters `alpha`, `beta` and their sum are expected to be differened from 0 for r=   z@Parameter `alpha` is expected to be float differened from 1 for )r    Úfrom_strr8   r.   r0   r1   r@   rA   Ú
ValueErrorr/   Úanyr9   r:   )Úselfr8   r9   r:   Ú_bad_measuress        r%   Ú__init__z_InformationMeasure.__init__i   sñ  € ô $+×#3Ñ#3Ð4GÓ#HˆÔ Ü ×1Ñ1´7×3HÑ3HÌ'×JbÑJbÐcˆØ×#Ñ# }Ñ4¼ZÈÌuÔ=UÜÐOÐPcÐOdÐdeÐfÓgÐgØ×#Ñ#¬×(?Ñ(?Ä×AVÑAVÐ'WÑWÔ`jÐkoÔqvÔ`wÜÐNÐObÐNcÐcdÐeÓfÐfØ×#Ñ#¤w×'?Ñ'?Ò?ÌÐTYÔ[`ÔIaÐejÐntÑetÜØXÐYlÐXmÐmnÐoóð ð ×#Ñ#¤w×'>Ñ'>Ò>Ì
ÐSWÔY^ÔH_ÐcgÐkrÑcrÜØXÐYlÐXmÐmnÐoóð ð ×#Ñ#¤w×'<Ñ'<Ò<ØˆMØˆ|ÜÑD°u¸d°mÔDÔDÈÈeÐUYÐ[`ÐcgÑ[gÐMhÑHhäØdØ&Ð' qð*óð ð ×#Ñ#¤w×'?Ñ'?Ò?ÌÐTYÔ[`ÔIaÐejÐnoÒeoÜÐ_Ð`sÐ_tÐtuÐvÓwÐwð ’Z˜aˆŒ
Ø’I˜Aˆ�	r$   Úpreds_distributionÚtarget_distributionc                 ó€   — t        | d| j                  j                  › �«      }t        j                   |||«      «      S )NÚ_calculate_)Úgetattrr8   ÚvalueÚtorchÚ
nan_to_num)rH   rK   rL   Úinformation_measure_functions       r%   Ú__call__z_InformationMeasure.__call__�   s>   € Ü'.¨t°{À4×C[ÑC[×CaÑCaÐBbÐ5cÓ'dÐ$Ü×ÑÑ <Ð=OÐQdÓ eÓfÐfr$   c                 ób   — t        j                  |t        j                  | |z  «      z  d¬«      S )aú  Calculate Kullback-Leibler divergence between discrete distributions of predicted and reference sentences.

        Args:
            preds_distribution:
                Discrete reference distribution of predicted sentences over the vocabulary.
            target_distribution:
                Discrete reference distribution of reference sentences over the vocabulary.

        Return:
            Kullback-Leibler divergence between discrete distributions of predicted and reference sentences.

        r>   ©Údim)rQ   ÚsumÚlog©rK   rL   s     r%   Ú_calculate_kl_divergencez,_InformationMeasure._calculate_kl_divergence‘   s,   € ô �y‰yÐ,¬u¯y©yÐ9KÐNaÑ9aÓ/bÑbÐhjÔkÐkr$   c                 ó´   — | j                   | j                   dz
  z  }dt        j                  || j                   z  |d| j                   z
  z  z  d¬«      z
  |z  S )aä  Calculate alpha divergence between discrete distributions of predicted and reference sentences.

        Args:
            preds_distribution:
                Discrete reference distribution of predicted sentences over the vocabulary.
            target_distribution:
                Discrete reference distribution of reference sentences over the vocabulary.

        Return:
            Alpha divergence between discrete distributions of predicted and reference sentences.

        r=   r>   rV   )r9   rQ   rX   )rH   rK   rL   Ú_alpha_denoms       r%   Ú_calculate_alpha_divergencez/_InformationMeasure._calculate_alpha_divergence¡   s]   € ð —z‘z T§Z¡Z°!¡^Ñ4ˆà”—	‘	Ð-¨t¯z©zÑ9Ð<NÐSTÐW[×WaÑWaÑSaÑ<bÑbÐhjÔkÑkØñð 	r$   c                 óŒ  — t        j                  t        j                  || j                  | j                  z   z  d¬«      «      }|| j                  | j                  | j                  z   z  z  }t        j                  t        j                  || j                  | j                  z   z  d¬«      «      }|| j                  | j                  | j                  z   z  z  }t        j                  t        j                  || j                  z  || j                  z  z  d¬«      «      }|| j                  | j                  z  z  }||z   |z
  S )aÞ  Calculate AB divergence between discrete distributions of predicted and reference sentences.

        Args:
            preds_distribution:
                Discrete reference distribution of predicted sentences over the vocabulary.
            target_distribution:
                Discrete reference distribution of reference sentences over the vocabulary.

        Return:
            AB divergence between discrete distributions of predicted and reference sentences.

        r>   rV   )rQ   rY   rX   r:   r9   )rH   rK   rL   ÚaÚbÚcs         r%   Ú_calculate_ab_divergencez,_InformationMeasure._calculate_ab_divergence³   sù   € ô �I‰I”e—i‘iÐ 3¸¿	¹	ÀDÇJÁJÑ8NÑ OÐUWÔXÓYˆØ	ˆT�Y‰Y˜$Ÿ)™) d§j¡jÑ0Ñ1Ñ1ˆÜ�I‰I”e—i‘iÐ 2°t·y±yÀ4Ç:Á:Ñ7MÑ NÐTVÔWÓXˆØ	ˆT�Z‰Z˜4Ÿ9™9 t§z¡zÑ1Ñ2Ñ2ˆÜ�I‰I”e—i‘iÐ 3°T·Z±ZÑ ?ÐBTÐVZ×V_ÑV_ÑB_Ñ _ÐegÔhÓiˆØ	ˆT�Z‰Z˜$Ÿ)™)Ñ#Ñ#ˆà�1‰u�q‰yÐr$   c                 ó4   — d| _         | j                  ||«      S )aâ  Calculate beta divergence between discrete distributions of predicted and reference sentences.

        Args:
            preds_distribution:
                Discrete reference distribution of predicted sentences over the vocabulary.
            target_distribution:
                Discrete reference distribution of reference sentences over the vocabulary.

        Return:
            Beta divergence between discrete distributions of predicted and reference sentences.

        g      ð?)r9   rc   ©rH   rK   rL   s      r%   Ú_calculate_beta_divergencez._InformationMeasure._calculate_beta_divergenceÉ   s    € ð ˆŒ
Ø×,Ñ,Ð-?ÐATÓUÐUr$   c                 ó¶   — t        j                  t        j                  || j                  z  |d| j                  z
  z  z  d¬«      «      | j                  dz
  z  S )uæ  Calculate RÃ©nyi divergence between discrete distributions of predicted and reference sentences.

        Args:
            preds_distribution:
                Discrete reference distribution of predicted sentences over the vocabulary.
            target_distribution:
                Discrete reference distribution of reference sentences over the vocabulary.

        Return:
            RÃ©nyi divergence between discrete distributions of predicted and reference sentences.

        r=   r>   rV   )rQ   rY   rX   r9   re   s      r%   Ú_calculate_renyi_divergencez/_InformationMeasure._calculate_renyi_divergenceÙ   sS   € ô �I‰I”e—i‘iÐ 3°T·Z±ZÑ ?ÐBTÐYZÐ]a×]gÑ]gÑYgÑBhÑ hÐnpÔqÓrØ�Z‰Z˜!‰^ñð 	r$   c                 ó8   — t        j                  || z
  dd¬«      S )aÚ  Calculate L1 distance between discrete distributions of predicted and reference sentences.

        Args:
            preds_distribution:
                Discrete reference distribution of predicted sentences over the vocabulary.
            target_distribution:
                Discrete reference distribution of reference sentences over the vocabulary.

        Return:
            L1 distance between discrete distributions of predicted and reference sentences.

        r=   r>   ©rC   rW   ©rQ   ÚnormrZ   s     r%   Ú_calculate_l1_distancez*_InformationMeasure._calculate_l1_distanceê   ó   € ô �z‰zÐ-Ð0BÑBÀaÈRÔPÐPr$   c                 ó8   — t        j                  || z
  dd¬«      S )aÚ  Calculate L2 distance between discrete distributions of predicted and reference sentences.

        Args:
            preds_distribution:
                Discrete reference distribution of predicted sentences over the vocabulary.
            target_distribution:
                Discrete reference distribution of reference sentences over the vocabulary.

        Return:
            L2 distance between discrete distributions of predicted and reference sentences.

        é   r>   rj   rk   rZ   s     r%   Ú_calculate_l2_distancez*_InformationMeasure._calculate_l2_distanceú   rn   r$   c                 óJ   — t        j                  || z
  t        d«      d¬«      S )aê  Calculate L-infinity distance between discrete distributions of predicted and reference sentences.

        Args:
            preds_distribution:
                Discrete reference distribution of predicted sentences over the vocabulary.
            target_distribution:
                Discrete reference distribution of reference sentences over the vocabulary.

        Return:
            L-infinity distance between discrete distributions of predicted and reference sentences.

        Úinfr>   rj   )rQ   rl   rA   rZ   s     r%   Ú_calculate_l_infinity_distancez2_InformationMeasure._calculate_l_infinity_distance
  s#   € ô �z‰zÐ-Ð0BÑBÄeÈEÃlÐXZÔ[Ð[r$   c           	      ó¦   — dt        j                  t        j                  t        j                  | |z  «      j	                  d«      dd«      «      z  S )aê  Calculate Fisher-Rao distance between discrete distributions of predicted and reference sentences.

        Args:
            preds_distribution:
                Discrete reference distribution of predicted sentences over the vocabulary.
            target_distribution:
                Discrete reference distribution of reference sentences over the vocabulary.

        Return:
            Fisher-Rao distance between discrete distributions of predicted and reference sentences.

        rp   r>   r   r=   )rQ   ÚacosÚclampÚsqrtrX   rZ   s     r%   Ú_calculate_fisher_rao_distancez2_InformationMeasure._calculate_fisher_rao_distance  sC   € ð ”5—:‘:œeŸk™k¬%¯*©*Ð5GÐJ]Ñ5]Ó*^×*bÑ*bÐceÓ*fÐhiÐklÓmÓnÑnÐnr$   )NN)r'   r(   r)   r*   Ú$_ALLOWED_INFORMATION_MEASURE_LITERALr   rA   rJ   r	   rT   r+   r[   r^   rc   rf   rh   rm   rq   rt   ry   r#   r$   r%   r7   r7   I   sÇ  „ ñðD "&Ø $ñ	"àAð"ð ˜‰ð"ð �u‰oð	"ð
 
ó"ðHg¨6ð gÈð gÐSYó gð ðl°Vð lÐRXð lÐ]cò ló ðlð¸fð Ð[að Ðfló ð$¸6ð ÐX^ð Ðció ð,V¸Vð VÐZ`ð VÐekó Vð ¸fð Ð[að Ðfló ð" ðQ°6ð QÐPVð QÐ[aò Qó ðQð ðQ°6ð QÐPVð QÐ[aò Qó ðQð ð\¸6ð \ÐX^ð \Ðciò \ó ð\ð ðo¸6ð oÐX^ð oÐciò oó ñor$   r7   Ú	input_idsÚattention_maskÚidfÚ
batch_sizeÚnum_workersr!   c                 ó8   — t        | ||«      }t        |||¬«      S )aH  Prepare dataloader.

    Args:
        input_ids:
            Indices of input sequence tokens in the vocabulary.
        attention_mask:
            Mask to avoid performing attention on padding token indices.
        idf:
            A bool indicating whether normalization using inverse document frequencies should be used.
        batch_size:
            A batch size used for model processing.
        num_workers:
            A number of workers to use for a dataloader.

    Return:
        An instance of ``torch.utils.data.DataLoader`` used for iterating over examples.

    )r~   r   )r   r   )r{   r|   r}   r~   r   Údatasets         r%   Ú_get_dataloaderr‚   +  s!   € ô* ˜y¨.¸#Ó>€GÜ�g¨*À+ÔNÐNr$   Ú	tokenizerr   c                 ó`   — | j                   | j                  | j                  | j                  dœS )a  Build a dictionary of model/tokenizer special tokens.

    Args:
        tokenizer:
            Initialized tokenizer from HuggingFace's `transformers package.

    Return:
        A dictionary containing: mask_token_id, pad_token_id, sep_token_id and cls_token_id.

    ©Úmask_token_idÚpad_token_idÚsep_token_idÚcls_token_idr…   )rƒ   s    r%   Ú_get_special_tokens_maprŠ   D  s2   € ð #×0Ñ0Ø!×.Ñ.Ø!×.Ñ.Ø!×.Ñ.ñ	ð r$   r‡   rˆ   r‰   c                 ór   — | j                  |«      | j                  |«      z  | j                  |«      z  }| S )a(  Generate a token mask for differentiating all special tokens in the input batch.

    There are 0s for special tokens and 1s otherwise.

    Args:
        input_ids:
            Indices of input sequence tokens in the vocabulary.
        pad_token_id:
            An id of ``<PAD>`` tokens that are used to make arrays of tokens the same size for batching purpose
        cls_token_id:
            An id of ``<CLS>`` token that represents the class of the input. (It might be ``<BOS>`` token for some
            models.)
        sep_token_id:
            An id of ``<SEP>`` token that separates two different sentences in the same input. (It might be ``<EOS>``
            token for some models.)

    Return:
        Tensor mask of 0s and 1s that masks all special tokens in the ``input_ids`` tensor.

    )Úeq)r{   r‡   rˆ   r‰   Ú
token_masks        r%   Ú_get_token_maskrŽ   W  s7   € ð* —‘˜lÓ+¨i¯l©l¸<Ó.HÑHÈ9Ï<É<ÐXdÓKeÑe€JØˆ;Ðr$   Úmodelr   ÚbatchÚtemperatureÚspecial_tokens_mapc                 ó¬  — |d   j                   d   }g }t        |d   |d   |d   |d   «      }t        |«      D ]Ç  }|d   j                  «       }	|d   |	dd…|f<    | |	|d   «      j                  }
|
dd…|dd…f   }
t        j                  |
|z  d	¬
«      }|r7||d   dd…|f   j                  d«      j                  |j                  «      z  }|j                  |j                  d«      j                  «       «       ~	~
~ŒÉ t        j                  |d¬
«      }t        j                  d|j                  |j                  «      |«      }|rU||d   j                  |j                  «      z  }|j                  d¬
«      |j                  d¬
«      j                  d«      z  S |j                  d¬
«      |j                  d¬
«      j                  d«      z  S )a_  Calculate a discrete probability distribution for a batch of examples. See `InfoLM`_ for details.

    Args:
        model:
            Initialized model from HuggingFace's `transformers package.
        batch:
            An input batch dictionary containing ``input_ids`` and ``attention_mask``.
        temperature:
            A temperature for calibrating language modelling. For more information, please reference `InfoLM`_ paper.
        max_length:
            A maximum length of input sequences. Sequences longer than `max_length` are to be trimmed.
        idf:
            An indication of whether normalization using inverse document frequencies should be used.
        special_tokens_map:
            A dictionary mapping tokenizer special tokens into the corresponding integer values.

    Return:
        A discrete probability distribution.

    r{   r=   r‡   rˆ   r‰   r†   Nr|   r>   rV   Úinput_ids_idfzbsv, bs -> bsv)ÚshaperŽ   ÚrangeÚcloneÚlogitsÚFÚsoftmaxÚ	unsqueezeÚtoÚdeviceÚappendÚcpurQ   ÚcatÚeinsumrX   )r�   r�   r‘   r}   r’   Úseq_lenÚprob_distribution_batch_listr�   Úmask_idxr{   Úlogits_distributionÚprob_distributionÚprob_distribution_batchÚmasked_input_ids_idfs                 r%   Ú_get_batch_distributionr©   p  sï  € ð6 �KÑ ×&Ñ& qÑ)€GØ13Ð Ü ØˆkÑØ˜>Ñ*Ø˜>Ñ*Ø˜>Ñ*ó	€Jô ˜'“Nò >ˆØ˜+Ñ&×,Ñ,Ó.ˆ	Ø!3°OÑ!Dˆ	’!�X�+ÑÙ# I¨uÐ5EÑ/FÓG×NÑNÐà1²!°Xºq°.ÑAÐÜŸI™IÐ&9¸KÑ&GÈRÔPÐÙØ  Ñ!7º¸8¸Ñ!D×!NÑ!NÈqÓ!Q×!TÑ!TÐUf×UmÑUmÓ!nÑnÐØ$×+Ñ+Ð,=×,GÑ,GÈÓ,J×,NÑ,NÓ,PÔQàÐ*Ñ,=ð>ô $Ÿi™iÐ(DÈ!ÔLÐÜ#Ÿl™lÐ+;Ð=T×=WÑ=WÐXb×XiÑXiÓ=jÐlvÓwÐÙ
Ø)¨E°/Ñ,B×,EÑ,EÀj×FWÑFWÓ,XÑXÐØ&×*Ñ*¨qÐ*Ó1Ð4H×4LÑ4LÐQRÐ4LÓ4S×4]Ñ4]Ð^_Ó4`Ñ`Ð`à"×&Ñ&¨1Ð&Ó-°
·±À1°Ó0E×0OÑ0OÐPQÓ0RÑRÐRr$   Ú
dataloaderÚverbosec           
      óÂ   — | j                   }g }t        ||«      D ],  }t        ||«      }|j                  t	        | ||||«      «       Œ. t        j                  |d¬«      S )aá  Calculate a discrete probability distribution according to the methodology described in `InfoLM`_.

    Args:
        model:
            Initialized model from HuggingFace's `transformers package.
        dataloader:
            An instance of `torch.utils.data.DataLoader` used for iterating over examples.
        temperature:
            A temperature for calibrating language modelling. For more information, please reference `InfoLM`_ paper.
        max_length:
            A maximum length of input sequences. Sequences longer than `max_length` are to be trimmed.
        idf:
            An indication of whether normalization using inverse document frequencies should be used.
        special_tokens_map:
            A dictionary mapping tokenizer special tokens into the corresponding integer values.
        verbose:
            An indication of whether a progress bar to be displayed during the embeddings calculation.

    Return:
        A discrete probability distribution.

    r   rV   )r�   r   r   rž   r©   rQ   r    )	r�   rª   r‘   r}   r’   r«   r�   r¦   r�   s	            r%   Ú_get_data_distributionr­   ©  si   € ð> �\‰\€FØ&(Ðä" :¨wÓ7ò nˆÜ$ U¨FÓ3ˆØ× Ñ Ô!8¸ÀÀ{ÐTWÐYkÓ!lÕmðnô �9‰9Ð&¨AÔ.Ð.r$   ÚpredsÚtargetÚ
max_lengthc                 ó  — t        | t        t        f«      st        | «      } t        |t        t        f«      st        |«      } || d|dd¬«      } ||d|dd¬«      }|j                  |j                  |j                  |j                  fS )a8  Update the metric state by a tokenization of ``preds`` and ``target`` sentencens.

    Args:
        preds:
            An iterable of hypothesis corpus.
        target:
            An iterable of reference corpus.
        tokenizer:
            Initialized tokenizer from HuggingFace's `transformers package.
        max_length:
            A maximum length of input sequences. Sequences longer than `max_length` are to be trimmed.

    Return:
        Tokenizerd ``preds`` and ``target`` sentences represented with ``input_ids`` and ``attention_mask`` tensors.

    r°   TÚpt)Úpaddingr°   Ú
truncationÚreturn_tensors)r@   r,   Úlistr{   r|   )r®   r¯   rƒ   r°   Úpreds_inputÚtarget_inputs         r%   Ú_infolm_updater¹   Ò  s‚   € ô. �eœc¤4˜[Ô)Ü�U“ˆÜ�fœs¤D˜kÔ*Ü�f“ˆá˜E¨<ÀJÐ[_ÐptÔu€KÙ˜V¨\ÀjÐ]aÐrvÔw€Là× Ñ  +×"<Ñ"<¸l×>TÑ>TÐVb×VqÑVqÐqÐqr$   Úpreds_dataloaderÚtarget_dataloaderÚinformation_measure_clsc                 ó¸   — t        | |||||«      }t        | |||||«      }	||j                  j                     }|	|j                  j                     }	 |||	«      S )al  Calculate selected information measure using the pre-trained language model.

    Args:
        model:
            Initialized model from HuggingFace's `transformers package.
        preds_dataloader:
            Loader iterating over tokenizer predicted sentences.
        target_dataloader:
            Loader iterating over tokenizer reference sentences.
        temperature:
            A temperature for calibrating language modelling. For more information, please reference `InfoLM`_ paper.
        idf:
            An indication of whether normalization using inverse document frequencies should be used.
        information_measure_cls:
            Information measure class containing all parameters necessary for calculating information measure values
            using ``preds_distribution`` and ``target_distribution``.
        special_tokens_map:
            A dictionary mapping tokenizer special tokens into the corresponding integer values.
        verbose:
            An indication of whether a progress bar to be displayed during the embeddings calculation.

    Return:
        A corpus-level InfoLM score.

    )r­   r�   Úsorting_indices)
r�   rº   r»   r‘   r}   r¼   r’   r«   rK   rL   s
             r%   Ú_infolm_computer¿   ô  sy   € ôF 0°Ð7GÈÐVYÐ[mÐovÓwÐÜ0ØÐ  +¨sÐ4FÈóÐð ,Ð,<×,DÑ,D×,TÑ,TÑUÐØ-Ð.?×.GÑ.G×.WÑ.WÑXÐá"Ð#5Ð7JÓKÐKr$   Úmodel_name_or_pathr8   r9   r:   r�   Únum_threadsÚreturn_sentence_level_scorec           
      óR  — t        ||«      \  }}t        |||«      }|	xs |j                  j                  }	t	        |«      }t        | |||	«      \  }}}}t        ||||
|«      }t        ||||
|«      }t        ||||||||«      }|r|j                  «       |fS |j                  «       S )uo  Calculate `InfoLM`_ [1].

    InfoML corresponds to distance/divergence between predicted and reference sentence discrete distribution using
    one of the following information measures:

        - `KL divergence`_
        - `alpha divergence`_
        - `beta divergence`_
        - `AB divergence`_
        - `RÃ©nyi divergence`_
        - L1 distance
        - L2 distance
        - L-infinity distance
        - `Fisher-Rao distance`_

    `InfoLM`_ is a family of untrained embedding-based metrics which addresses some famous flaws of standard
    string-based metrics thanks to the usage of pre-trained masked language models. This family of metrics is mainly
    designed for summarization and data-to-text tasks.

    If you want to use IDF scaling over the whole dataset, please use the class metric.

    The implementation of this metric is fully based HuggingFace `transformers`' package.

    Args:
        preds:
            An iterable of hypothesis corpus.
        target:
            An iterable of reference corpus.
        model_name_or_path:
            A name or a model path used to load `transformers` pretrained model.
        temperature:
            A temperature for calibrating language modelling. For more information, please reference `InfoLM`_ paper.
        information_measure:
            A name of information measure to be used. Please use one of: ['kl_divergence', 'alpha_divergence',
            'beta_divergence', 'ab_divergence', 'renyi_divergence', 'l1_distance', 'l2_distance', 'l_infinity_distance',
            'fisher_rao_distance']
        idf:
            An indication of whether normalization using inverse document frequencies should be used.
        alpha:
            Alpha parameter of the divergence used for alpha, AB and RÃ©nyi divergence measures.
        beta:
            Beta parameter of the divergence used for beta and AB divergence measures.
        device:
            A device to be used for calculation.
        max_length:
            A maximum length of input sequences. Sequences longer than `max_length` are to be trimmed.
        batch_size:
            A batch size used for model processing.
        num_threads:
            A number of threads to use for a dataloader.
        verbose:
            An indication of whether a progress bar to be displayed during the embeddings calculation.
        return_sentence_level_score:
            An indication whether a sentence-level InfoLM score to be returned.

    Returns:
        A corpus-level InfoLM score.
        (Optionally) A list of sentence-level InfoLM scores if `return_sentence_level_score=True`.

    Example:
        >>> from torchmetrics.functional.text.infolm import infolm
        >>> preds = ['he read the book because he was interested in world history']
        >>> target = ['he was interested in world history because he read the book']
        >>> infolm(preds, target, model_name_or_path='google/bert_uncased_L-2_H-128_A-2', idf=False)
        tensor(-0.1784)

    References:
        [1] InfoLM: A New Metric to Evaluate Summarization & Data2Text Generation by Pierre Colombo, ChloÃ© Clavel and
        Pablo Piantanida `InfoLM`_

    )	r   r7   Úconfigr°   rŠ   r¹   r‚   r¿   Úmean)r®   r¯   rÀ   r‘   r8   r}   r9   r:   r�   r°   r~   rÁ   r«   rÂ   rƒ   r�   r¼   r’   Úpreds_input_idsÚpreds_attention_maskÚtarget_input_idsÚtarget_attention_maskrº   r»   Úinfo_lm_scores                            r%   r   r   "  sà   € ôn 1Ð1CÀVÓLÑ€IˆuÜ1Ð2EÀuÈdÓSÐØÒ6˜uŸ|™|×6Ñ6€JÜ0°Ó;ÐäUcØˆv�y *óVÑR€OÐ)Ð+;Ð=Rô ' Ð8LÈcÐS]Ð_jÓkÐÜ'Ð(8Ð:OÐQTÐV`ÐbmÓnÐä#ØØØØØØØØó	€Mñ #Ø×!Ñ!Ó# ]Ð2Ð2à×ÑÓÐr$   )T)zbert-base-uncasedg      Ð?r   TNNNNé@   r   TF)4ÚosÚcollections.abcr   Úenumr   Útypingr   r   r   r   rQ   r	   Útorch.nnr
   r™   Útorch.utils.datar   Útyping_extensionsr   Ú4torchmetrics.functional.text.helper_embedding_metricr   r   r   r   Útorchmetrics.utilities.enumsr   Útorchmetrics.utilities.importsr   Útransformersr   r   Ú__doctest_skip__rz   r    r7   ÚboolÚintr‚   Údictr,   rŠ   rŽ   rA   r©   Úno_gradr­   Útupler¹   r¿   ÚPathLiker�   r   r#   r$   r%   ú<module>rÞ      sÇ  ðó 
Ý $Ý ß 7Ó 7ã Ý Ý $Ý 'Ý %÷ó õ 1Ý JáÑ4ßEá&Ø �zÐð (/ðñ
(Ð $ð ô0ˆgó 0ó ð0÷$_oñ _oðDOØðOØ'-ðOØ48ðOØFIðOØX[ðOàóOð2Ð'@ð ÀTÈ#ÈsÈ(Á^ó ð&˜vð °Sð Èð Ð[^ð Ðció ð26SØð6Sà��V�Ñð6Sð ð6Sð 
ð	6Sð
 ˜S #˜X™ð6Sð ó6Sðr €‡�ƒð%/Øð%/àð%/ð ð%/ð 
ð	%/ð
 ˜S #˜X™ð%/ð ð%/ð ò%/ó ð%/ðPrØ��h˜s‘mÐ#Ñ$ðrà�#�x ‘}Ð$Ñ%ðrð )ðrð ð	rð
 ˆ6�6˜6 6Ð)Ñ*órðT ñ+LØð+Là ð+Lð "ð+Lð ð	+Lð
 
ð+Lð 1ð+Lð ˜S #˜X™ð+Lð ð+Lð ó+Lðb 3FØØ@OØØ!Ø Ø15Ø $ØØØØ(-ñp Ø��h˜s‘mÐ#Ñ$ðp à�#�x ‘}Ð$Ñ%ðp ð ˜c 2§;¡;Ð.Ñ/ðp ð ð	p ð
 >ðp ð 
ðp ð �E‰?ðp ð �5‰/ðp ð �U˜3 §¡Ð,Ñ-Ñ.ðp ð ˜‘ðp ð ðp ð ðp ð ðp ð "&ðp ð ˆ6�5˜ ˜Ñ(Ð(Ñ)ôp r$   