a
    dD                     @   s  d dl Z d dlZd dlmZ d dlmZmZmZmZm	Z	 d dl
mZmZmZ d dlmZ d dlmZ eeeddd	Zd)eeeeeeedddZeedddZeedddZee edddZe	eee f ee	eee f  e	ed ed f ee	eee f ee	eee f  f dddZd*ee	eee f eeeeeddd Zd+e	eee f ee	eee f  ed! eeeeeee  ee d"	d#d$Zd,e	eee f ee	eee f  ed! eeeeee	eeeef f d&	d'd(ZdS )-    N)inf)ListOptionalSequenceTupleUnion)Tensorstacktensor)Literal)_validate_inputs)
preds_wordtarget_wordreturnc                 C   s   t | |kS )a.  Distance measure used for substitutions/identity operation.

    Code adapted from https://github.com/rwth-i6/ExtendedEditDistance/blob/master/EED.py.

    Args:
        preds_word: hypothesis word string
        target_word: reference word string

    Return:
        0 for match, 1 for no match

    )int)r   r    r   i/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torchmetrics/functional/text/eed.py_distance_between_wordsd   s    r          @333333?皙?      ?)hyprefalpharhodeletion	insertionr   c              
      sj  dgt | d  }dgt | d  }d|d< tgt | d  }tdt |d D ]}	tdt | d D ]d}
|
dkrt||
d  | ||
d  t| |
d  ||	d   ||
 | ||
< qf||
 d ||
< qf|t|}||  d7  < ||	d  dkr|||    fdd|D }|}tgt | d  }qP|td	d
 |D  }td|d | tt ||  S )a8  Compute extended edit distance score for two lists of strings: hyp and ref.

    Code adapted from: https://github.com/rwth-i6/ExtendedEditDistance/blob/master/EED.py.

    Args:
        hyp: A hypothesis string
        ref: A reference string
        alpha: optimal jump penalty, penalty for jumps between characters
        rho: coverage cost, penalty for repetition of characters
        deletion: penalty for deletion of character
        insertion: penalty for insertion or substitution of character

    Return:
        Extended edit distance score as float
       r           r    c                    s   g | ]}t | qS r   )min.0xZjumpr   r   
<listcomp>       z!_eed_function.<locals>.<listcomp>c                 s   s   | ]}|d kr|ndV  qdS )r   r   Nr   r#   r   r   r   	<genexpr>   r(   z _eed_function.<locals>.<genexpr>)lenr   ranger"   r   indexsumfloat)r   r   r   r   r   r   Znumber_of_visitsrowZnext_rowwiZ	min_indexZcoverager   r&   r   _eed_functiont   s,    $

r2   )sentencer   c                 C   s   t | tstdt|  d|  } g d}|D ]\}}| ||} q2g d}|D ]\}}t||| } qTg d}|D ]\}}| ||} qxd|  d S )zPreprocess english sentences.

    Copied from https://github.com/rwth-i6/ExtendedEditDistance/blob/master/util.py.

    Raises:
        ValueError: If input sentence is not of a type `str`.

    6Only strings allowed during preprocessing step, found  instead)).z .)!z !)?z ?),z ,))z\s+r!   )z(\d) ([.,]) (\d)z\1\2\3)z#(Dr|Jr|Prof|Rev|Gen|Mr|Mt|Mrs|Ms) .z\1.))ze . g .ze.g.)zi . e .zi.e.)zU . S .zU.S.r!   )
isinstancestr
ValueErrortyperstripreplaceresub)r3   Zrules_interpunctionpatternreplacementZrules_rer   r   r   _preprocess_en   s    	
rD   c                 C   s2   t | tstdt|  d|  } td| S )zPreprocess japanese sentences.

    Copy from https://github.com/rwth-i6/ExtendedEditDistance/blob/master/util.py.

    Raises:
        ValueError: If input sentence is not of a type `str`.

    r4   r5   NFKC)r:   r;   r<   r=   r>   unicodedata	normalize)r3   r   r   r   _preprocess_ja   s    	
rH   )sentence_level_scoresr   c                 C   s(   t | dkrtdS t| tt |  S )zReduction for extended edit distance.

    Args:
        sentence_level_scores: list of sentence-level scores as floats

    Return:
        average of scores as a tensor

    r   r    )r*   r
   r-   )rI   r   r   r   _eed_compute   s    
rJ   enja)predstargetlanguager   c                    sf   t | |d\}} |dkrt n|dkr,t ntd|  fdd| D }  fdd|D }| |fS )au  Preprocess strings according to language requirements.

    Args:
        preds: An iterable of hypothesis corpus.
        target: An iterable of iterables of reference corpus.
        language: Language used in sentences. Only supports English (en) and Japanese (ja) for now. Defaults to en

    Return:
        Tuple of lists that contain the cleaned strings for target and preds

    Raises:
        ValueError: If a different language than ``'en'`` or ``'ja'`` is used
        ValueError: If length of target not equal to length of preds
        ValueError: If objects in reference and hypothesis corpus are not strings

    )Zhypothesis_corpusZ
ref_corpusrK   rL   z?Expected argument `language` to either be `en` or `ja` but got c                    s   g | ]} |qS r   r   )r$   predZpreprocess_functionr   r   r'     r(   z)_preprocess_sentences.<locals>.<listcomp>c                    s   g | ]} fd d|D qS )c                    s   g | ]} |qS r   r   )r$   r   rQ   r   r   r'     r(   z4_preprocess_sentences.<locals>.<listcomp>.<listcomp>r   )r$   	referencerQ   r   r   r'     r(   )r   rD   rH   r<   )rM   rN   rO   r   rQ   r   _preprocess_sentences   s    rS   )r   target_wordsr   r   r   r   r   c           	      C   s4   t }|D ]"}t| |||||}||k r|}qt|S )a  Compute scores for ExtendedEditDistance.

    Args:
        target_words: An iterable of reference words
        preds_word: A hypothesis word
        alpha: An optimal jump penalty, penalty for jumps between characters
        rho: coverage cost, penalty for repetition of characters
        deletion: penalty for deletion of character
        insertion: penalty for insertion or substitution of character

    Return:
        best_score: best (lowest) sentence-level score as a Tensor

    )r   r2   r
   )	r   rT   r   r   r   r   Z
best_scorerR   scorer   r   r   _compute_sentence_statistics"  s    rV   )rK   rL   )	rM   rN   rO   r   r   r   r   sentence_eedr   c                 C   sl   t | ||\} }|du rg }dt| t|d fv r8|S t| |D ]$\}}	t||	||||}
||
 qB|S )a  Compute scores for ExtendedEditDistance.

    Args:
        preds: An iterable of hypothesis corpus
        target: An iterable of iterables of reference corpus
        language: Language used in sentences. Only supports English (en) and Japanese (ja) for now. Defaults to en
        alpha: optimal jump penalty, penalty for jumps between characters
        rho: coverage cost, penalty for repetition of characters
        deletion: penalty for deletion of character
        insertion: penalty for insertion or substitution of character
        sentence_eed: list of sentence-level scores

    Return:
        individual sentence scores as a list of Tensors

    Nr   )rS   r*   ziprV   append)rM   rN   rO   r   r   r   r   rW   Z
hypothesisrT   rU   r   r   r   _eed_updateB  s    rZ   F)	rM   rN   rO   return_sentence_level_scorer   r   r   r   r   c                 C   s|   t g d||||gD ]4\}}	t|	tr:t|	tr|	dk rtd| dqt| ||||||}
t|
}|rx|t|
fS |S )uX  Compute extended edit distance score (`ExtendedEditDistance`_) [1] for strings or list of strings.

    The metric utilises the Levenshtein distance and extends it by adding a jump operation.

    Args:
        preds: An iterable of hypothesis corpus.
        target: An iterable of iterables of reference corpus.
        language: Language used in sentences. Only supports English (en) and Japanese (ja) for now. Defaults to en
        return_sentence_level_score: An indication of whether sentence-level EED score is to be returned.
        alpha: optimal jump penalty, penalty for jumps between characters
        rho: coverage cost, penalty for repetition of characters
        deletion: penalty for deletion of character
        insertion: penalty for insertion or substitution of character

    Return:
        Extended edit distance score as a tensor

    Example:
        >>> from torchmetrics.functional.text import extended_edit_distance
        >>> preds = ["this is the prediction", "here is an other sample"]
        >>> target = ["this is the reference", "here is another one"]
        >>> extended_edit_distance(preds=preds, target=target)
        tensor(0.3078)

    References:
        [1] P. Stanchev, W. Wang, and H. Ney, “EED: Extended Edit Distance Measure for Machine Translation”,
        submitted to WMT 2019. `ExtendedEditDistance`_

    )r   r   r   r   r   zParameter `z)` is expected to be a non-negative float.)rX   r:   r.   r<   rZ   rJ   r	   )rM   rN   rO   r[   r   r   r   r   
param_nameparamrI   Zaverager   r   r   extended_edit_distancel  s    (r^   )r   r   r   r   )r   r   r   r   )rK   r   r   r   r   N)rK   Fr   r   r   r   )r@   rF   mathr   typingr   r   r   r   r   Ztorchr   r	   r
   Ztyping_extensionsr   Z#torchmetrics.functional.text.helperr   r;   r   r   r.   r2   rD   rH   rJ   rS   rV   rZ   boolr^   r   r   r   r   <module>Y   s       :-&)    #      
-      