Ë
    þÍ:j¸  ã                   ó–   — d dl mZ d dlZd dlmZ dededdfd„Zddededee   deeef   fd	„Zd
ededefd„Zddededee   defd„Z	y)é    )ÚOptionalN)ÚTensorÚpredsÚtargetÚreturnc                 óV  — t        | j                  «      dk7  r"t        dt        | j                  «      › d�«      ‚t        |j                  «      dk7  r"t        dt        |j                  «      › d�«      ‚| j                  dd |j                  k7  r)t        d| j                  dd › d|j                  › d�«      ‚| j                  «       st	        d	| j
                  › d�«      ‚|j
                  t        j                  k7  r*t	        d
t        j                  › d|j
                  › d�«      ‚y)a3  Check shape and type consistency of input vectors.

    Args:
        preds:
            Logits or a unnormalized score assigned to each token in a sequence with shape [batch_size, seq_len,
            vocab_size]. Scores will be normalized internally using softmax.
        target:
            Ground truth values with a shape [batch_size, seq_len].

    Raises:
        ValueError:
            If ``preds`` tensor has no 3 dimensions.
        ValueError:
            If ``target`` tensor has no 2 dimensions.
        ValueError:
            If the first two dimensions of ``preds`` and ``target`` do not equal.
        TypeError:
            If ``preds`` dtype is not one of ``(torch.float16, torch.float32, torch.float64)``
        TypeError:
            If ``target`` is not of a type LongTensor (torch.int64)

    é   zbInput tensor `preds` is expected to have 3 dimensions, [batch_size, seq_len, vocab_size], but got ú.é   zWInput tensor `target` is expected to have 2 dimensions, [batch_size, seq_len], but got NzvInput tensors `preds` and `target` are expected to have equaling first two dimensions, [batch_size, seq_len], but got z and zFInput tensor `preds` is expected to be of floating point type but got z2Input tensor `target` is expected to be of a type z	 but got )ÚlenÚshapeÚ
ValueErrorÚis_floating_pointÚ	TypeErrorÚdtypeÚtorchÚint64)r   r   s     ú|/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/text/perplexity.pyÚ!_check_shape_and_type_consistencyr      s/  € ô. ˆ5�;‰;Ó˜1ÒÜðÜ˜EŸK™KÓ(Ð)¨ð,ó
ð 	
ô ˆ6�<‰<Ó˜AÒÜðÜ˜FŸL™LÓ)Ð*¨!ð-ó
ð 	
ð ‡{�{�2�A€˜&Ÿ,™,Ò&Üð/Ø/4¯{©{¸2¸A¨Ð.?¸uÀVÇ\Á\ÀNÐRSðUó
ð 	
ð ×"Ñ"Ô$ÜÐ`Ðaf×alÑalÐ`mÐmnÐoÓpÐpØ‡|�|”u—{‘{Ò"ÜÐLÌUÏ[É[ÈMÐYbÐci×coÑcoÐbpÐpqÐrÓsÐsð #ó    Úignore_indexc                 ób  — t        | |«       t        j                  j                  j	                  | j                  d| j                  d   «      d¬«      }|j                  d«      }|�F|j                  |«      }|j                  ||k7  t        j                  d|j                  ¬«      «      }n%t        j                  |t        j                  ¬«      }|t        j                  |j                  «       «      |f   |   }|j                  «       j!                  «        }|j!                  «       }||fS )a]  Compute intermediate statistics for Perplexity.

    Args:
        preds:
            Logits or a unnormalized score assigned to each token in a sequence with shape [batch_size, seq_len,
            vocab_size]. Scores will be normalized internally using softmax.
        target:
            Ground truth values with a shape [batch_size, seq_len].
        ignore_index:
            Integer specifying a target class to ignore. If given, this class index does not contribute
            to the returned score.

    Returns:
        Log probabilities, summed over all samples
        Number of samples

    éÿÿÿÿé   )Údimr   )Údevice)r   )r   r   ÚnnÚ
functionalÚsoftmaxÚreshaper   ÚneÚwhereÚtensorr   Ú	ones_likeÚboolÚarangeÚnumelÚlogÚsum)r   r   r   ÚprobsÚmaskÚtotal_log_probsÚcounts          r   Ú_perplexity_updater.   A   sé   € ô$ & e¨VÔ4ä�H‰H×Ñ×'Ñ'¨¯©°b¸%¿+¹+Àb¹/Ó(JÐPQÐ'ÓR€EØ�^‰^˜BÓ€FàÐØ�y‰y˜Ó&ˆØ—‘˜f¨Ñ4´e·l±lÀ1ÈVÏ]É]Ô6[Ó\‰ä�‰˜v¬U¯Z©ZÔ8ˆà”%—,‘,˜vŸ|™|›~Ó.°Ð6Ñ7¸Ñ=€EØ—y‘y“{—‘Ó(Ð(€OØ�H‰H‹J€Eà˜EÐ!Ð!r   Útotalr-   c                 ó2   — t        j                  | |z  «      S )z£Compute the Perplexity.

    Args:
        total: Log probabilities, summed over all samples
        count: Number of samples
    Returns:
        Perplexity

    )r   Úexp)r/   r-   s     r   Ú_perplexity_computer2   e   s   € ô �9‰9�U˜U‘]Ó#Ð#r   c                 ó:   — t        | ||«      \  }}t        ||«      S )aË  Perplexity measures how well a language model predicts a text sample.

    This metric is calculated as the average number of bits per word a model needs to represent the sample.

    Args:
        preds:
            Logits or a unnormalized score assigned to each token in a sequence with shape [batch_size, seq_len,
            vocab_size], which is the output of a language model. Scores will be normalized internally using softmax.
        target:
            Ground truth values with a shape [batch_size, seq_len].
        ignore_index:
            Integer specifying a target class to ignore. If given, this class index does not contribute
            to the returned score.

    Returns:
        Perplexity value

    Examples:
        >>> from torch import rand, randint
        >>> preds = rand(2, 8, 5)
        >>> target = randint(5, (2, 8))
        >>> target[0, 6:] = -100
        >>> perplexity(preds, target, ignore_index=-100)
        tensor(5.8540)

    )r.   r2   )r   r   r   r/   r-   s        r   Ú
perplexityr4   r   s#   € ô6 & e¨V°\ÓB�L€Eˆ5Ü˜u eÓ,Ð,r   )N)
Útypingr   r   r   r   ÚintÚtupler.   r2   r4   © r   r   ú<module>r9      s¢   ðõ ã Ý ð)t¨Vð )t¸Vð )tÈó )tñX!"˜fð !"¨fð !"ÀHÈSÁMð !"Ð]bÐciÐkqÐcqÑ]ró !"ðH
$˜vð 
$¨fð 
$¸ó 
$ñ-�fð - fð -¸HÀS¹Mð -ÐU[ô -r   