Ë
    þÍ:jï
  ã                   ó–   — d dl mZmZ d dlmZ d dlmZ dedeeef   fd„Z	dedeeef   d	eed
      defd„Z
dded	eed
      defd„Zy)é    )ÚOptionalÚUnion)ÚTensor)ÚLiteralÚimgÚreturnc                 ód  — | j                   dk7  rt        d| j                  › �«      ‚| ddd…dd…f   | ddd…dd…f   z
  }| ddd…dd…f   | ddd…dd…f   z
  }|j                  «       j	                  g d¢«      }|j                  «       j	                  g d¢«      }||z   }|| j                  d   fS )	z4Compute total variation statistics on current batch.é   z1Expected input `img` to be an 4D tensor, but got .é   Néÿÿÿÿ)r   é   é   r   )ÚndimÚRuntimeErrorÚshapeÚabsÚsum)r   Údiff1Údiff2Úres1Úres2Úscores         úu/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/image/tv.pyÚ_total_variation_updater      s²   € à
‡x�x�1‚}ÜÐNÈsÏyÉyÈkÐZÓ[Ð[Ø��Q‘Rš�
‰O˜c # s¨ sªA +Ñ.Ñ.€EØ�’Q˜™�
‰O˜c #¢q¨#¨2¨# +Ñ.Ñ.€Eà�9‰9‹;�?‰?š9Ó%€DØ�9‰9‹;�?‰?š9Ó%€DØ�4‰K€EØ�#—)‘)˜A‘,ÐÐó    r   Únum_elementsÚ	reduction)Úmeanr   Únonec                 ó„   — |dk(  r| j                  «       |z  S |dk(  r| j                  «       S |�|dk(  r| S t        d«      ‚)z$Compute final total variation score.r   r   r   zHExpected argument `reduction` to either be 'sum', 'mean', 'none' or None)r   Ú
ValueError)r   r   r   s      r   Ú_total_variation_computer"   !   sO   € ð �FÒØ�y‰y‹{˜\Ñ)Ð)Ø�EÒØ�y‰y‹{ÐØÐ˜I¨Ò/ØˆÜ
Ð_Ó
`Ð`r   c                 ó8   — t        | «      \  }}t        |||«      S )a9  Compute total variation loss.

    Args:
        img: A `Tensor` of shape `(N, C, H, W)` consisting of images
        reduction: a method to reduce metric score over samples.

            - ``'mean'``: takes the mean over samples
            - ``'sum'``: takes the sum over samples
            - ``None`` or ``'none'``: return the score per sample

    Returns:
        A loss scalar value containing the total variation

    Raises:
        ValueError:
            If ``reduction`` is not one of ``'sum'``, ``'mean'``, ``'none'`` or ``None``
        RuntimeError:
            If ``img`` is not 4D tensor

    Example:
        >>> from torch import rand
        >>> from torchmetrics.functional.image import total_variation
        >>> img = rand(5, 3, 28, 28)
        >>> total_variation(img)
        tensor(7546.8018)

    )r   r"   )r   r   r   r   s       r   Útotal_variationr$   .   s"   € ô< 2°#Ó6Ñ€Eˆ<Ü# E¨<¸ÓCÐCr   N)r   )Útypingr   r   Útorchr   Útyping_extensionsr   ÚtupleÚintr   r"   r$   © r   r   ú<module>r+      s”   ð÷ #å Ý %ð
 ð 
¨E°&¸#°+Ñ,>ó 
ð
aØð
aØ!& s¨F {Ñ!3ð
aØ@HÈÐQfÑIgÑ@hð
aàó
añD˜ð D¨H°WÐ=RÑ5SÑ,Tð DÐagô Dr   